submission 418024
shiyegao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 852 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-418024?include=source"interfacepython
Compatibility
measured onNVIDIA H100
declared hardwareNVIDIA H100
architecturessm_90
dtypesfp32
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:d953e31dd740063673c4bdf89849475ede426a98a0f9214f66a111b8b28697cd
license declaredunknown
license concludedunknown
authorsshiyegao
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ half tile_l[32][33];vector-width = float4
float4 v = reinterpret_cast<const float4*>(xrow)[c4];Kernel source
submission.py852 lines
from __future__ import annotations
from typing import Any, Dict, Tuple
__PRECISION_NOTE__ = "fp16_gemm_fp32_accum"
_EXT = None
_KERNEL_CACHE: dict[tuple, dict[str, Any]] = {}
_WEIGHT_CACHE: dict[tuple, dict[str, Any]] = {}
def _load_ext():
global _EXT
if _EXT is not None:
return _EXT
import hashlib
import os
from torch.utils.cpp_extension import load_inline
this_dir = os.path.dirname(os.path.abspath(__file__))
build_dir = os.path.join(this_dir, ".torch_ext_build")
os.makedirs(build_dir, exist_ok=True)
os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "9.0")
cpp_src = r"""
#include <torch/extension.h>
torch::Tensor trimul_forward(
torch::Tensor x,
torch::Tensor mask_h,
torch::Tensor norm_w,
torch::Tensor norm_b,
torch::Tensor w_stack,
torch::Tensor out_norm_w,
torch::Tensor out_norm_b,
torch::Tensor w_to_out,
torch::Tensor work_u8,
torch::Tensor x_norm_h,
torch::Tensor y0_h,
torch::Tensor left_packed,
torch::Tensor right_packed,
torch::Tensor out_gate_packed,
torch::Tensor out_packed,
torch::Tensor out2_h,
torch::Tensor y_out_f);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("forward", &trimul_forward, "trimul outgoing forward (CUDA)");
}
"""
cuda_src = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cublasLt.h>
#ifndef CHECK_CUDA
#define CHECK_CUDA(x) TORCH_CHECK((x).is_cuda(), #x " must be a CUDA tensor")
#endif
#ifndef CHECK_CONTIGUOUS
#define CHECK_CONTIGUOUS(x) TORCH_CHECK((x).is_contiguous(), #x " must be contiguous")
#endif
#ifndef CHECK_DTYPE
#define CHECK_DTYPE(x, dt) TORCH_CHECK((x).dtype() == (dt), #x " dtype mismatch")
#endif
static __device__ __forceinline__ float _warp_sum(float v) {
unsigned mask = 0xffffffffu;
v += __shfl_down_sync(mask, v, 16);
v += __shfl_down_sync(mask, v, 8);
v += __shfl_down_sync(mask, v, 4);
v += __shfl_down_sync(mask, v, 2);
v += __shfl_down_sync(mask, v, 1);
return __shfl_sync(mask, v, 0);
}
static __device__ __forceinline__ float _sigmoid(float x) {
return 1.0f / (1.0f + __expf(-x));
}
__global__ void ln_fwd_fp16(
const float* __restrict__ x,
half* __restrict__ y,
const float* __restrict__ w,
const float* __restrict__ b,
int rows,
int dim) {
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
int row = (blockIdx.x * (blockDim.x >> 5)) + warp;
if (row >= rows) return;
const float* xrow = x + (long long)row * dim;
half* yrow = y + (long long)row * dim;
float sum = 0.0f;
float sum2 = 0.0f;
for (int c4 = lane; c4 < (dim >> 2); c4 += 32) {
float4 v = reinterpret_cast<const float4*>(xrow)[c4];
sum += v.x + v.y + v.z + v.w;
sum2 += v.x * v.x + v.y * v.y + v.z * v.z + v.w * v.w;
}
sum = _warp_sum(sum);
sum2 = _warp_sum(sum2);
float mean = sum / (float)dim;
float var = fmaxf(sum2 / (float)dim - mean * mean, 0.0f);
float inv = rsqrtf(var + 1e-5f);
for (int c4 = lane; c4 < (dim >> 2); c4 += 32) {
float4 v = reinterpret_cast<const float4*>(xrow)[c4];
int c = c4 << 2;
float4 ww = make_float4(w[c + 0], w[c + 1], w[c + 2], w[c + 3]);
float4 bb = make_float4(b[c + 0], b[c + 1], b[c + 2], b[c + 3]);
float4 o;
o.x = (v.x - mean) * inv * ww.x + bb.x;
o.y = (v.y - mean) * inv * ww.y + bb.y;
o.z = (v.z - mean) * inv * ww.z + bb.z;
o.w = (v.w - mean) * inv * ww.w + bb.w;
reinterpret_cast<half2*>(yrow)[c4 * 2 + 0] = __floats2half2_rn(o.x, o.y);
reinterpret_cast<half2*>(yrow)[c4 * 2 + 1] = __floats2half2_rn(o.z, o.w);
}
}
__global__ void pack_lr_og(
const half* __restrict__ y0,
const half* __restrict__ mask,
half* __restrict__ left_p,
half* __restrict__ right_p,
half* __restrict__ og_p,
int rows,
int n,
int hidden) {
__shared__ half tile_l[32][33];
__shared__ half tile_r[32][33];
__shared__ half tile_o[32][33];
int x = (blockIdx.x << 5) + threadIdx.x;
int y = (blockIdx.y << 5) + threadIdx.y;
#pragma unroll
for (int i = 0; i < 4; ++i) {
int row = y + (i << 3);
if (x < hidden && row < rows) {
float m = __half2float(mask[row]);
long long base = (long long)row * (5LL * hidden) + x;
float lp = __half2float(y0[base + 0LL * hidden]);
float rp = __half2float(y0[base + 1LL * hidden]);
float lg = __half2float(y0[base + 2LL * hidden]);
float rg = __half2float(y0[base + 3LL * hidden]);
float og = __half2float(y0[base + 4LL * hidden]);
float l = lp * _sigmoid(lg) * m;
float r = rp * _sigmoid(rg) * m;
float o = _sigmoid(og);
tile_l[threadIdx.y + (i << 3)][threadIdx.x] = __float2half_rn(l);
tile_r[threadIdx.y + (i << 3)][threadIdx.x] = __float2half_rn(r);
tile_o[threadIdx.y + (i << 3)][threadIdx.x] = __float2half_rn(o);
}
}
__syncthreads();
int xt = (blockIdx.y << 5) + threadIdx.x;
int yt = (blockIdx.x << 5) + threadIdx.y;
#pragma unroll
for (int i = 0; i < 4; ++i) {
int h = yt + (i << 3);
int row = xt;
if (h < hidden && row < rows) {
int b = row / (n * n);
int local = row - b * (n * n);
int ii = local / n;
int jj = local - ii * n;
long long out_idx = ((long long)(b * hidden + h) * n + ii) * n + jj;
left_p[out_idx] = tile_l[threadIdx.x][threadIdx.y + (i << 3)];
right_p[out_idx] = tile_r[threadIdx.x][threadIdx.y + (i << 3)];
og_p[out_idx] = tile_o[threadIdx.x][threadIdx.y + (i << 3)];
}
}
}
__global__ void ln2_gate_pack_h128(
const half* __restrict__ out_p,
const half* __restrict__ og_p,
const float* __restrict__ w,
const float* __restrict__ b,
half* __restrict__ out2,
int rows_in_b) {
int tx = (int)threadIdx.x;
int ty = (int)threadIdx.y;
int row_base = (int)blockIdx.x << 5;
int bid = (int)blockIdx.y;
int row = row_base + tx;
__shared__ half sv[128][33];
__shared__ half sg[128][33];
float sum = 0.0f;
float sum2 = 0.0f;
#pragma unroll
for (int h_base = 0; h_base < 128; h_base += 32) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int h = h_base + ty + (i << 3);
if (row < rows_in_b) {
long long idx = ((long long)(bid * 128 + h) * rows_in_b) + row;
half hv = out_p[idx];
half hg = og_p[idx];
sv[h][tx] = hv;
sg[h][tx] = hg;
float v = __half2float(hv);
sum += v;
sum2 += v * v;
}
}
}
__shared__ float sh_sum[8][32];
__shared__ float sh_sum2[8][32];
__shared__ float sh_mean[32];
__shared__ float sh_inv[32];
sh_sum[ty][tx] = sum;
sh_sum2[ty][tx] = sum2;
__syncthreads();
if (ty == 0 && row < rows_in_b) {
float s = 0.0f;
float s2 = 0.0f;
#pragma unroll
for (int t = 0; t < 8; ++t) {
s += sh_sum[t][tx];
s2 += sh_sum2[t][tx];
}
float mean = s * (1.0f / 128.0f);
float var = fmaxf(s2 * (1.0f / 128.0f) - mean * mean, 0.0f);
sh_mean[tx] = mean;
sh_inv[tx] = rsqrtf(var + 1e-5f);
}
__syncthreads();
#pragma unroll
for (int h_base = 0; h_base < 128; h_base += 32) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int row_off = ty + (i << 3);
int out_row = row_base + row_off;
int h = h_base + tx;
if (out_row < rows_in_b) {
float v = __half2float(sv[h][row_off]);
float g = __half2float(sg[h][row_off]);
float mean = sh_mean[row_off];
float inv = sh_inv[row_off];
float nv = (v - mean) * inv * w[h] + b[h];
out2[((long long)(bid * rows_in_b + out_row) * 128) + h] = __float2half_rn(nv * g);
}
}
}
}
__global__ void ln2_gate_pack_tiled(
const half* __restrict__ out_p,
const half* __restrict__ og_p,
const float* __restrict__ w,
const float* __restrict__ b,
half* __restrict__ out2,
int rows_in_b,
int hidden) {
int tx = (int)threadIdx.x;
int ty = (int)threadIdx.y;
int row_base = (int)blockIdx.x << 5;
int bid = (int)blockIdx.y;
int row = row_base + tx;
float sum = 0.0f;
float sum2 = 0.0f;
if (row < rows_in_b) {
for (int h = ty; h < hidden; h += 8) {
long long idx = ((long long)(bid * hidden + h) * rows_in_b) + row;
float v = __half2float(out_p[idx]);
sum += v;
sum2 += v * v;
}
}
__shared__ float sh_sum[8][32];
__shared__ float sh_sum2[8][32];
__shared__ float sh_mean[32];
__shared__ float sh_inv[32];
sh_sum[ty][tx] = sum;
sh_sum2[ty][tx] = sum2;
__syncthreads();
if (ty == 0 && row < rows_in_b) {
float s = 0.0f;
float s2 = 0.0f;
#pragma unroll
for (int t = 0; t < 8; ++t) {
s += sh_sum[t][tx];
s2 += sh_sum2[t][tx];
}
float mean = s / (float)hidden;
float var = fmaxf(s2 / (float)hidden - mean * mean, 0.0f);
sh_mean[tx] = mean;
sh_inv[tx] = rsqrtf(var + 1e-5f);
}
__syncthreads();
__shared__ half tile_v[32][33];
__shared__ half tile_g[32][33];
for (int h_base = 0; h_base < hidden; h_base += 32) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int h = h_base + ty + (i << 3);
if (row < rows_in_b && h < hidden) {
long long idx = ((long long)(bid * hidden + h) * rows_in_b) + row;
tile_v[ty + (i << 3)][tx] = out_p[idx];
tile_g[ty + (i << 3)][tx] = og_p[idx];
}
}
__syncthreads();
#pragma unroll
for (int i = 0; i < 4; ++i) {
int row_off = ty + (i << 3);
int out_row = row_base + row_off;
int h = h_base + tx;
if (out_row < rows_in_b && h < hidden) {
float v = __half2float(tile_v[tx][row_off]);
float g = __half2float(tile_g[tx][row_off]);
float mean = sh_mean[row_off];
float inv = sh_inv[row_off];
float nv = (v - mean) * inv * w[h] + b[h];
out2[((long long)(bid * rows_in_b + out_row) * hidden) + h] = __float2half_rn(nv * g);
}
}
__syncthreads();
}
}
struct LtPlanKey {
int m, n, k;
int batch;
int op_b;
int a_type, b_type, c_type, d_type;
};
struct LtPlan {
bool valid;
LtPlanKey key;
cublasLtMatmulAlgo_t algo;
size_t work_bytes;
cublasLtMatmulDesc_t op_desc;
cublasLtMatrixLayout_t a_desc;
cublasLtMatrixLayout_t b_desc;
cublasLtMatrixLayout_t c_desc;
cublasLtMatrixLayout_t d_desc;
};
static cublasLtHandle_t _lt = nullptr;
static LtPlan _plan_g1 = {false};
static LtPlan _plan_g2 = {false};
static LtPlan _plan_ct = {false};
static inline bool _key_eq(const LtPlanKey& a, const LtPlanKey& b) {
return a.m==b.m && a.n==b.n && a.k==b.k && a.batch==b.batch && a.op_b==b.op_b &&
a.a_type==b.a_type && a.b_type==b.b_type && a.c_type==b.c_type && a.d_type==b.d_type;
}
static inline void _lt_init() {
if (_lt) return;
auto st = cublasLtCreate(&_lt);
TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "cublasLtCreate failed");
}
static inline void _lt_plan_destroy(LtPlan* plan) {
if (!plan->valid) return;
if (plan->a_desc) cublasLtMatrixLayoutDestroy(plan->a_desc);
if (plan->b_desc) cublasLtMatrixLayoutDestroy(plan->b_desc);
if (plan->c_desc) cublasLtMatrixLayoutDestroy(plan->c_desc);
if (plan->d_desc) cublasLtMatrixLayoutDestroy(plan->d_desc);
if (plan->op_desc) cublasLtMatmulDescDestroy(plan->op_desc);
plan->a_desc = nullptr;
plan->b_desc = nullptr;
plan->c_desc = nullptr;
plan->d_desc = nullptr;
plan->op_desc = nullptr;
plan->valid = false;
}
static inline cublasLtMatrixLayout_t _lt_make_layout(
cudaDataType type,
int rows,
int cols,
int ld,
int batch,
long long stride,
cublasLtOrder_t order) {
cublasLtMatrixLayout_t out = nullptr;
auto st = cublasLtMatrixLayoutCreate(&out, type, rows, cols, ld);
TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "cublasLtMatrixLayoutCreate failed");
st = cublasLtMatrixLayoutSetAttribute(out, CUBLASLT_MATRIX_LAYOUT_ORDER, &order, sizeof(order));
TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "layout set order failed");
if (batch > 1) {
st = cublasLtMatrixLayoutSetAttribute(out, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, &batch, sizeof(batch));
TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "layout set batch failed");
st = cublasLtMatrixLayoutSetAttribute(out, CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET, &stride, sizeof(stride));
TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "layout set stride failed");
}
return out;
}
static inline void _lt_pick_algo(
LtPlan* plan,
const LtPlanKey& key,
cublasOperation_t op_a,
cublasOperation_t op_b,
cudaDataType a_type,
cudaDataType b_type,
cudaDataType c_type,
cudaDataType d_type,
int lda, int ldb, int ldc, int ldd,
int batch,
long long stride_a,
long long stride_b,
long long stride_c,
long long stride_d,
size_t work_bytes) {
_lt_init();
if (plan->valid) {
_lt_plan_destroy(plan);
}
auto st = cublasLtMatmulDescCreate(&plan->op_desc, CUBLAS_COMPUTE_32F, CUDA_R_32F);
TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "matmul desc create failed");
st = cublasLtMatmulDescSetAttribute(plan->op_desc, CUBLASLT_MATMUL_DESC_TRANSA, &op_a, sizeof(op_a));
TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "set transa failed");
st = cublasLtMatmulDescSetAttribute(plan->op_desc, CUBLASLT_MATMUL_DESC_TRANSB, &op_b, sizeof(op_b));
TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "set transb failed");
cublasLtOrder_t order = CUBLASLT_ORDER_ROW;
int a_rows = (op_a == CUBLAS_OP_N) ? key.m : key.k;
int a_cols = (op_a == CUBLAS_OP_N) ? key.k : key.m;
int b_rows = (op_b == CUBLAS_OP_N) ? key.k : key.n;
int b_cols = (op_b == CUBLAS_OP_N) ? key.n : key.k;
plan->a_desc = _lt_make_layout(a_type, a_rows, a_cols, lda, batch, stride_a, order);
plan->b_desc = _lt_make_layout(b_type, b_rows, b_cols, ldb, batch, stride_b, order);
plan->c_desc = _lt_make_layout(c_type, key.m, key.n, ldc, batch, stride_c, order);
plan->d_desc = _lt_make_layout(d_type, key.m, key.n, ldd, batch, stride_d, order);
cublasLtMatmulPreference_t pref = nullptr;
st = cublasLtMatmulPreferenceCreate(&pref);
TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "pref create failed");
st = cublasLtMatmulPreferenceSetAttribute(pref, CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, &work_bytes, sizeof(work_bytes));
TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "pref set failed");
cublasLtMatmulHeuristicResult_t heur;
int got = 0;
st = cublasLtMatmulAlgoGetHeuristic(_lt, plan->op_desc, plan->a_desc, plan->b_desc, plan->c_desc, plan->d_desc, pref, 1, &heur, &got);
TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS && got > 0, "no cublasLt heuristic algo");
plan->valid = true;
plan->key = key;
plan->algo = heur.algo;
plan->work_bytes = work_bytes;
cublasLtMatmulPreferenceDestroy(pref);
}
static inline void _lt_matmul(
LtPlan* plan,
const LtPlanKey& key,
cublasOperation_t op_a,
cublasOperation_t op_b,
const void* a,
const void* b,
const void* c,
void* d,
cudaDataType a_type,
cudaDataType b_type,
cudaDataType c_type,
cudaDataType d_type,
int lda, int ldb, int ldc, int ldd,
int batch,
long long stride_a,
long long stride_b,
long long stride_c,
long long stride_d,
void* work,
size_t work_bytes) {
_lt_init();
if (!plan->valid || !_key_eq(plan->key, key)) {
_lt_pick_algo(plan, key, op_a, op_b, a_type, b_type, c_type, d_type, lda, ldb, ldc, ldd, batch, stride_a, stride_b, stride_c, stride_d, work_bytes);
}
float alpha = 1.0f;
float beta = 0.0f;
auto st = cublasLtMatmul(_lt, plan->op_desc, &alpha, a, plan->a_desc, b, plan->b_desc, &beta, c, plan->c_desc, d, plan->d_desc, &plan->algo, work, plan->work_bytes, 0);
TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "cublasLtMatmul failed");
}
torch::Tensor trimul_forward(
torch::Tensor x,
torch::Tensor mask_h,
torch::Tensor norm_w,
torch::Tensor norm_b,
torch::Tensor w_stack,
torch::Tensor out_norm_w,
torch::Tensor out_norm_b,
torch::Tensor w_to_out,
torch::Tensor work_u8,
torch::Tensor x_norm_h,
torch::Tensor y0_h,
torch::Tensor left_packed,
torch::Tensor right_packed,
torch::Tensor out_gate_packed,
torch::Tensor out_packed,
torch::Tensor out2_h,
torch::Tensor y_out_f) {
CHECK_CUDA(x);
CHECK_CUDA(mask_h);
CHECK_CUDA(norm_w);
CHECK_CUDA(norm_b);
CHECK_CUDA(w_stack);
CHECK_CUDA(out_norm_w);
CHECK_CUDA(out_norm_b);
CHECK_CUDA(w_to_out);
CHECK_CUDA(work_u8);
CHECK_CUDA(x_norm_h);
CHECK_CUDA(y0_h);
CHECK_CUDA(left_packed);
CHECK_CUDA(right_packed);
CHECK_CUDA(out_gate_packed);
CHECK_CUDA(out_packed);
CHECK_CUDA(out2_h);
CHECK_CUDA(y_out_f);
CHECK_CONTIGUOUS(x);
CHECK_CONTIGUOUS(mask_h);
CHECK_CONTIGUOUS(norm_w);
CHECK_CONTIGUOUS(norm_b);
CHECK_CONTIGUOUS(w_stack);
CHECK_CONTIGUOUS(out_norm_w);
CHECK_CONTIGUOUS(out_norm_b);
CHECK_CONTIGUOUS(w_to_out);
CHECK_CONTIGUOUS(work_u8);
CHECK_CONTIGUOUS(x_norm_h);
CHECK_CONTIGUOUS(y0_h);
CHECK_CONTIGUOUS(left_packed);
CHECK_CONTIGUOUS(right_packed);
CHECK_CONTIGUOUS(out_gate_packed);
CHECK_CONTIGUOUS(out_packed);
CHECK_CONTIGUOUS(out2_h);
CHECK_CONTIGUOUS(y_out_f);
CHECK_DTYPE(x, torch::kFloat32);
CHECK_DTYPE(mask_h, torch::kFloat16);
CHECK_DTYPE(norm_w, torch::kFloat32);
CHECK_DTYPE(norm_b, torch::kFloat32);
CHECK_DTYPE(w_stack, torch::kFloat16);
CHECK_DTYPE(out_norm_w, torch::kFloat32);
CHECK_DTYPE(out_norm_b, torch::kFloat32);
CHECK_DTYPE(w_to_out, torch::kFloat16);
CHECK_DTYPE(work_u8, torch::kUInt8);
CHECK_DTYPE(x_norm_h, torch::kFloat16);
CHECK_DTYPE(y0_h, torch::kFloat16);
CHECK_DTYPE(left_packed, torch::kFloat16);
CHECK_DTYPE(right_packed, torch::kFloat16);
CHECK_DTYPE(out_gate_packed, torch::kFloat16);
CHECK_DTYPE(out_packed, torch::kFloat16);
CHECK_DTYPE(out2_h, torch::kFloat16);
CHECK_DTYPE(y_out_f, torch::kFloat32);
TORCH_CHECK(x.dim() == 4, "x must be [bs,N,N,dim]");
int bs = (int)x.size(0);
int n = (int)x.size(1);
int dim = (int)x.size(3);
TORCH_CHECK((int)x.size(2) == n, "x must be square on N");
TORCH_CHECK((int)mask_h.size(0) == bs && (int)mask_h.size(1) == n && (int)mask_h.size(2) == n, "mask shape");
int hidden5 = (int)w_stack.size(1);
TORCH_CHECK(hidden5 % 5 == 0, "w_stack second dim must be 5*hidden");
int hidden = hidden5 / 5;
int rows = bs * n * n;
int rows_in_b = n * n;
TORCH_CHECK((int)x_norm_h.size(0) == rows && (int)x_norm_h.size(1) == dim, "x_norm_h shape");
TORCH_CHECK((dim & 3) == 0, "dim must be multiple of 4");
int warps = 8;
dim3 block1(32 * warps, 1, 1);
dim3 grid1((rows + warps - 1) / warps, 1, 1);
ln_fwd_fp16<<<grid1, block1>>>(
(const float*)x.data_ptr<float>(),
(half*)x_norm_h.data_ptr<at::Half>(),
(const float*)norm_w.data_ptr<float>(),
(const float*)norm_b.data_ptr<float>(),
rows, dim);
TORCH_CHECK((int)y0_h.size(0) == rows && (int)y0_h.size(1) == 5 * hidden, "y0_h shape");
LtPlanKey k1;
k1.m = rows; k1.n = 5 * hidden; k1.k = dim; k1.batch = 1; k1.op_b = 0;
k1.a_type = (int)CUDA_R_16F; k1.b_type = (int)CUDA_R_16F; k1.c_type = (int)CUDA_R_16F; k1.d_type = (int)CUDA_R_16F;
_lt_matmul(
&_plan_g1, k1,
CUBLAS_OP_N, CUBLAS_OP_N,
x_norm_h.data_ptr<at::Half>(),
w_stack.data_ptr<at::Half>(),
y0_h.data_ptr<at::Half>(),
y0_h.data_ptr<at::Half>(),
CUDA_R_16F, CUDA_R_16F, CUDA_R_16F, CUDA_R_16F,
dim, 5 * hidden, 5 * hidden, 5 * hidden,
1, 0, 0, 0, 0,
work_u8.data_ptr(), (size_t)work_u8.numel());
TORCH_CHECK((int)left_packed.size(0) == bs * hidden && (int)left_packed.size(1) == n && (int)left_packed.size(2) == n, "left_packed shape");
TORCH_CHECK(left_packed.sizes() == right_packed.sizes(), "right_packed shape");
TORCH_CHECK(left_packed.sizes() == out_gate_packed.sizes(), "out_gate_packed shape");
dim3 block2(32, 8, 1);
dim3 grid2((hidden + 31) / 32, (rows + 31) / 32, 1);
pack_lr_og<<<grid2, block2>>>(
(const half*)y0_h.data_ptr<at::Half>(),
(const half*)mask_h.data_ptr<at::Half>(),
(half*)left_packed.data_ptr<at::Half>(),
(half*)right_packed.data_ptr<at::Half>(),
(half*)out_gate_packed.data_ptr<at::Half>(),
rows, n, hidden);
TORCH_CHECK(out_packed.sizes() == left_packed.sizes(), "out_packed shape");
int batch_ct = bs * hidden;
LtPlanKey kc;
kc.m = n; kc.n = n; kc.k = n; kc.batch = batch_ct; kc.op_b = 1;
kc.a_type = (int)CUDA_R_16F; kc.b_type = (int)CUDA_R_16F; kc.c_type = (int)CUDA_R_16F; kc.d_type = (int)CUDA_R_16F;
long long stride_mat = (long long)n * n;
_lt_matmul(
&_plan_ct, kc,
CUBLAS_OP_N, CUBLAS_OP_T,
left_packed.data_ptr<at::Half>(),
right_packed.data_ptr<at::Half>(),
out_packed.data_ptr<at::Half>(),
out_packed.data_ptr<at::Half>(),
CUDA_R_16F, CUDA_R_16F, CUDA_R_16F, CUDA_R_16F,
n, n, n, n,
batch_ct,
stride_mat, stride_mat, stride_mat, stride_mat,
work_u8.data_ptr(), (size_t)work_u8.numel());
TORCH_CHECK((int)out2_h.size(0) == rows && (int)out2_h.size(1) == hidden, "out2_h shape");
TORCH_CHECK((int)out_norm_w.numel() == hidden && (int)out_norm_b.numel() == hidden, "out_norm weight/bias shape");
dim3 block3(32, 8, 1);
dim3 grid3((rows_in_b + 31) / 32, bs, 1);
if (hidden == 128) {
ln2_gate_pack_h128<<<grid3, block3>>>(
(const half*)out_packed.data_ptr<at::Half>(),
(const half*)out_gate_packed.data_ptr<at::Half>(),
(const float*)out_norm_w.data_ptr<float>(),
(const float*)out_norm_b.data_ptr<float>(),
(half*)out2_h.data_ptr<at::Half>(),
rows_in_b);
} else {
ln2_gate_pack_tiled<<<grid3, block3>>>(
(const half*)out_packed.data_ptr<at::Half>(),
(const half*)out_gate_packed.data_ptr<at::Half>(),
(const float*)out_norm_w.data_ptr<float>(),
(const float*)out_norm_b.data_ptr<float>(),
(half*)out2_h.data_ptr<at::Half>(),
rows_in_b, hidden);
}
TORCH_CHECK((int)y_out_f.size(0) == rows && (int)y_out_f.size(1) == dim, "y_out_f shape");
LtPlanKey k2;
k2.m = rows; k2.n = dim; k2.k = hidden; k2.batch = 1; k2.op_b = 0;
k2.a_type = (int)CUDA_R_16F; k2.b_type = (int)CUDA_R_16F; k2.c_type = (int)CUDA_R_32F; k2.d_type = (int)CUDA_R_32F;
_lt_matmul(
&_plan_g2, k2,
CUBLAS_OP_N, CUBLAS_OP_N,
out2_h.data_ptr<at::Half>(),
w_to_out.data_ptr<at::Half>(),
y_out_f.data_ptr<float>(),
y_out_f.data_ptr<float>(),
CUDA_R_16F, CUDA_R_16F, CUDA_R_32F, CUDA_R_32F,
hidden, dim, dim, dim,
1, 0, 0, 0, 0,
work_u8.data_ptr(), (size_t)work_u8.numel());
return y_out_f;
}
"""
name = "trimul_ext_" + hashlib.sha256((cpp_src + cuda_src).encode("utf-8")).hexdigest()[:16]
_EXT = load_inline(
name=name,
cpp_sources=[cpp_src],
cuda_sources=[cuda_src],
functions=None,
extra_cflags=["-O3", "-std=c++17"],
extra_cuda_cflags=["-O3", "--use_fast_math", "-std=c++17"],
extra_ldflags=["-lcublasLt", "-lcublas"],
with_cuda=True,
build_directory=build_dir,
verbose=False,
)
return _EXT
def _prep_weights(weights: Dict[str, "torch.Tensor"], dim: int, hidden_dim: int, device: "torch.device") -> dict[str, Any]:
import torch
key = (
int(weights["left_proj.weight"].data_ptr()),
int(weights["right_proj.weight"].data_ptr()),
int(weights["left_gate.weight"].data_ptr()),
int(weights["right_gate.weight"].data_ptr()),
int(weights["out_gate.weight"].data_ptr()),
int(weights["to_out.weight"].data_ptr()),
int(weights["norm.weight"].data_ptr()),
int(weights["norm.bias"].data_ptr()),
int(weights["to_out_norm.weight"].data_ptr()),
int(weights["to_out_norm.bias"].data_ptr()),
device.index if device.type == "cuda" else -1,
)
cached = _WEIGHT_CACHE.get(key)
if cached is not None:
return cached
def _f32_vec(t: "torch.Tensor") -> "torch.Tensor":
if t.device == device and t.dtype == torch.float32 and t.is_contiguous():
return t
return t.to(device=device, dtype=torch.float32).contiguous()
def _to_half_t(w: "torch.Tensor") -> "torch.Tensor":
return w.to(device=device, dtype=torch.float16).t().contiguous()
w_stack = torch.cat(
[
_to_half_t(weights["left_proj.weight"]),
_to_half_t(weights["right_proj.weight"]),
_to_half_t(weights["left_gate.weight"]),
_to_half_t(weights["right_gate.weight"]),
_to_half_t(weights["out_gate.weight"]),
],
dim=1,
).contiguous()
w_to_out = _to_half_t(weights["to_out.weight"])
out = {
"norm_w": _f32_vec(weights["norm.weight"]),
"norm_b": _f32_vec(weights["norm.bias"]),
"out_norm_w": _f32_vec(weights["to_out_norm.weight"]),
"out_norm_b": _f32_vec(weights["to_out_norm.bias"]),
"w_stack": w_stack,
"w_to_out": w_to_out,
}
_WEIGHT_CACHE[key] = out
return out
def _get_scratch(device: "torch.device", bs: int, n: int, dim: int, hidden: int) -> dict[str, Any]:
import torch
key = (device.type, device.index if device.type == "cuda" else -1, bs, n, dim, hidden)
cached = _KERNEL_CACHE.get(key)
if cached is not None:
return cached
rows = bs * n * n
work_bytes = 32 * 1024 * 1024
scratch = {
"work": torch.empty((work_bytes,), device=device, dtype=torch.uint8),
"x_norm": torch.empty((rows, dim), device=device, dtype=torch.float16),
"y0": torch.empty((rows, 5 * hidden), device=device, dtype=torch.float16),
"left_p": torch.empty((bs * hidden, n, n), device=device, dtype=torch.float16),
"right_p": torch.empty((bs * hidden, n, n), device=device, dtype=torch.float16),
"og_p": torch.empty((bs * hidden, n, n), device=device, dtype=torch.float16),
"out_p": torch.empty((bs * hidden, n, n), device=device, dtype=torch.float16),
"out2": torch.empty((rows, hidden), device=device, dtype=torch.float16),
"y_out": torch.empty((rows, dim), device=device, dtype=torch.float32),
}
_KERNEL_CACHE[key] = scratch
return scratch
def _ensure_mask_half(mask: "torch.Tensor", device: "torch.device") -> "torch.Tensor":
import torch
if mask.dtype == torch.float16 and mask.is_contiguous() and mask.device == device:
return mask
return mask.to(device=device, dtype=torch.float16).contiguous()
def custom_kernel(data: Tuple["torch.Tensor", "torch.Tensor", Dict[str, "torch.Tensor"], Dict[str, Any]]) -> "torch.Tensor":
import torch
x, mask, weights, config = data
dim = int(config["dim"])
hidden_dim = int(config["hidden_dim"])
if not x.is_cuda:
raise RuntimeError("CUDA only")
if x.dtype != torch.float32:
x = x.to(dtype=torch.float32)
if not x.is_contiguous():
x = x.contiguous()
bs, n0, n1, d0 = x.shape
if n0 != n1 or d0 != dim:
raise RuntimeError("shape mismatch")
mask_h = _ensure_mask_half(mask, x.device)
ext = _load_ext()
wpack = _prep_weights(weights, dim=dim, hidden_dim=hidden_dim, device=x.device)
scratch = _get_scratch(x.device, bs=bs, n=n0, dim=dim, hidden=hidden_dim)
y_flat = ext.forward(
x,
mask_h,
wpack["norm_w"],
wpack["norm_b"],
wpack["w_stack"],
wpack["out_norm_w"],
wpack["out_norm_b"],
wpack["w_to_out"],
scratch["work"],
scratch["x_norm"],
scratch["y0"],
scratch["left_p"],
scratch["right_p"],
scratch["og_p"],
scratch["out_p"],
scratch["out2"],
scratch["y_out"],
)
return y_flat.view(bs, n0, n0, dim)
__all__ = ["custom_kernel"]
scrolls · 852 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 417990.
⋯ 1 unchanged linesfrom typing import Any, Dict, Tuple- import torch- import torch.nn.functional as F+ __PRECISION_NOTE__ = "fp16_gemm_fp32_accum"- import cutlass- import cutlass.cute as cute- from cutlass.cute.runtime import make_ptr- from cutlass.cutlass_dsl import for_generate, if_generate, yield_out+ _EXT = None+ _KERNEL_CACHE: dict[tuple, dict[str, Any]] = {}+ _WEIGHT_CACHE: dict[tuple, dict[str, Any]] = {}- _TMN = 32- _TK = 32- _THREADS = 256+ def _load_ext():+ global _EXT+ if _EXT is not None:+ return _EXT+ import hashlib+ import os- class _BatchedGemmF32_32x32x32:- def __init__(self) -> None:- self.threads = _THREADS+ from torch.utils.cpp_extension import load_inline- @cute.jit- def __call__(- self,- a_ptr: "cute.Pointer",- b_ptr: "cute.Pointer",- c_ptr: "cute.Pointer",- problem: tuple,- ):- batch, n = problem+ this_dir = os.path.dirname(os.path.abspath(__file__))+ build_dir = os.path.join(this_dir, ".torch_ext_build")+ os.makedirs(build_dir, exist_ok=True)- stride_batch = n * n- stride_i = n+ os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "9.0")- a = cute.make_tensor(- a_ptr,- cute.make_layout(- (batch, n, n),- stride=(stride_batch, stride_i, 1),- ),- )- b = cute.make_tensor(- b_ptr,- cute.make_layout(- (batch, n, n),- stride=(stride_batch, stride_i, 1),- ),- )- c = cute.make_tensor(- c_ptr,- cute.make_layout(- (batch, n, n),- stride=(stride_batch, stride_i, 1),- ),- )+ cpp_src = r"""+ #include <torch/extension.h>+ torch::Tensor trimul_forward(+ torch::Tensor x,+ torch::Tensor mask_h,+ torch::Tensor norm_w,+ torch::Tensor norm_b,+ torch::Tensor w_stack,+ torch::Tensor out_norm_w,+ torch::Tensor out_norm_b,+ torch::Tensor w_to_out,+ torch::Tensor work_u8,+ torch::Tensor x_norm_h,+ torch::Tensor y0_h,+ torch::Tensor left_packed,+ torch::Tensor right_packed,+ torch::Tensor out_gate_packed,+ torch::Tensor out_packed,+ torch::Tensor out2_h,+ torch::Tensor y_out_f);- if (n & (_TMN - 1)) == 0:- grid_x = n // _TMN- grid_y = n // _TMN- self.kernel_fast(a, b, c, n).launch(- grid=[grid_x, grid_y, batch],- block=[self.threads, 1, 1],- )- else:- grid_x = (n + _TMN - 1) // _TMN- grid_y = (n + _TMN - 1) // _TMN- self.kernel_bnd(a, b, c, n).launch(- grid=[grid_x, grid_y, batch],- block=[self.threads, 1, 1],- )- return+ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {+ m.def("forward", &trimul_forward, "trimul outgoing forward (CUDA)");+ }+ """- @cute.kernel- def kernel_fast(self, a: "cute.Tensor", b: "cute.Tensor", c: "cute.Tensor", n: int):- tx, _, _ = cute.arch.thread_idx()- bx, by, bz = cute.arch.block_idx()+ cuda_src = r"""+ #include <torch/extension.h>+ #include <cuda.h>+ #include <cuda_runtime.h>+ #include <cuda_fp16.h>+ #include <cublasLt.h>- tid = tx- lane_m = tid >> 4- lane_n = tid & 15+ #ifndef CHECK_CUDA+ #define CHECK_CUDA(x) TORCH_CHECK((x).is_cuda(), #x " must be a CUDA tensor")+ #endif+ #ifndef CHECK_CONTIGUOUS+ #define CHECK_CONTIGUOUS(x) TORCH_CHECK((x).is_contiguous(), #x " must be contiguous")+ #endif+ #ifndef CHECK_DTYPE+ #define CHECK_DTYPE(x, dt) TORCH_CHECK((x).dtype() == (dt), #x " dtype mismatch")+ #endif- base_m = by * _TMN- base_n = bx * _TMN+ static __device__ __forceinline__ float _warp_sum(float v) {+ unsigned mask = 0xffffffffu;+ v += __shfl_down_sync(mask, v, 16);+ v += __shfl_down_sync(mask, v, 8);+ v += __shfl_down_sync(mask, v, 4);+ v += __shfl_down_sync(mask, v, 2);+ v += __shfl_down_sync(mask, v, 1);+ return __shfl_sync(mask, v, 0);+ }- i0 = base_m + lane_m- j0 = base_n + lane_n- i1 = i0 + 16- j1 = j0 + 16+ static __device__ __forceinline__ float _sigmoid(float x) {+ return 1.0f / (1.0f + __expf(-x));+ }- acc00 = cutlass.Float32(0.0)- acc01 = cutlass.Float32(0.0)- acc10 = cutlass.Float32(0.0)- acc11 = cutlass.Float32(0.0)+ __global__ void ln_fwd_fp16(+ const float* __restrict__ x,+ half* __restrict__ y,+ const float* __restrict__ w,+ const float* __restrict__ b,+ int rows,+ int dim) {+ int tid = threadIdx.x;+ int warp = tid >> 5;+ int lane = tid & 31;+ int row = (blockIdx.x * (blockDim.x >> 5)) + warp;+ if (row >= rows) return;- smem_stride = _TK + 1- smem_a_elems = _TMN * smem_stride- smem_b_elems = _TMN * smem_stride- smem_ptr = cute.arch.alloc_smem(cutlass.Float32, smem_a_elems + smem_b_elems, alignment=16)+ const float* xrow = x + (long long)row * dim;+ half* yrow = y + (long long)row * dim;- sA = cute.make_tensor(- smem_ptr,- cute.make_layout((_TMN, _TK), stride=(smem_stride, 1)),- )- sB = cute.make_tensor(- smem_ptr + smem_a_elems,- cute.make_layout((_TMN, _TK), stride=(smem_stride, 1)),- )+ float sum = 0.0f;+ float sum2 = 0.0f;- for k0, (acc00, acc01, acc10, acc11), (acc00_out, acc01_out, acc10_out, acc11_out) in for_generate(- 0,- n,- _TK,- iter_args=[acc00, acc01, acc10, acc11],- ):- base = tid << 2- for t in range(4):- idx = base + t- row = idx >> 5- col = idx & 31- sA[row, col] = a[bz, base_m + row, k0 + col]- sB[row, col] = b[bz, base_n + row, k0 + col]+ for (int c4 = lane; c4 < (dim >> 2); c4 += 32) {+ float4 v = reinterpret_cast<const float4*>(xrow)[c4];+ sum += v.x + v.y + v.z + v.w;+ sum2 += v.x * v.x + v.y * v.y + v.z * v.z + v.w * v.w;+ }+ sum = _warp_sum(sum);+ sum2 = _warp_sum(sum2);+ float mean = sum / (float)dim;+ float var = fmaxf(sum2 / (float)dim - mean * mean, 0.0f);+ float inv = rsqrtf(var + 1e-5f);- cute.arch.sync_threads()+ for (int c4 = lane; c4 < (dim >> 2); c4 += 32) {+ float4 v = reinterpret_cast<const float4*>(xrow)[c4];+ int c = c4 << 2;+ float4 ww = make_float4(w[c + 0], w[c + 1], w[c + 2], w[c + 3]);+ float4 bb = make_float4(b[c + 0], b[c + 1], b[c + 2], b[c + 3]);+ float4 o;+ o.x = (v.x - mean) * inv * ww.x + bb.x;+ o.y = (v.y - mean) * inv * ww.y + bb.y;+ o.z = (v.z - mean) * inv * ww.z + bb.z;+ o.w = (v.w - mean) * inv * ww.w + bb.w;+ reinterpret_cast<half2*>(yrow)[c4 * 2 + 0] = __floats2half2_rn(o.x, o.y);+ reinterpret_cast<half2*>(yrow)[c4 * 2 + 1] = __floats2half2_rn(o.z, o.w);+ }+ }- for kk in range(_TK):- a0 = sA[lane_m, kk]- a1 = sA[lane_m + 16, kk]- b0 = sB[lane_n, kk]- b1 = sB[lane_n + 16, kk]+ __global__ void pack_lr_og(+ const half* __restrict__ y0,+ const half* __restrict__ mask,+ half* __restrict__ left_p,+ half* __restrict__ right_p,+ half* __restrict__ og_p,+ int rows,+ int n,+ int hidden) {+ __shared__ half tile_l[32][33];+ __shared__ half tile_r[32][33];+ __shared__ half tile_o[32][33];- acc00 = acc00 + a0 * b0- acc01 = acc01 + a0 * b1- acc10 = acc10 + a1 * b0- acc11 = acc11 + a1 * b1+ int x = (blockIdx.x << 5) + threadIdx.x;+ int y = (blockIdx.y << 5) + threadIdx.y;- cute.arch.sync_threads()- yield_out([acc00, acc01, acc10, acc11])+ #pragma unroll+ for (int i = 0; i < 4; ++i) {+ int row = y + (i << 3);+ if (x < hidden && row < rows) {+ float m = __half2float(mask[row]);+ long long base = (long long)row * (5LL * hidden) + x;+ float lp = __half2float(y0[base + 0LL * hidden]);+ float rp = __half2float(y0[base + 1LL * hidden]);+ float lg = __half2float(y0[base + 2LL * hidden]);+ float rg = __half2float(y0[base + 3LL * hidden]);+ float og = __half2float(y0[base + 4LL * hidden]);- c[bz, i0, j0] = acc00_out- c[bz, i0, j1] = acc01_out- c[bz, i1, j0] = acc10_out- c[bz, i1, j1] = acc11_out+ float l = lp * _sigmoid(lg) * m;+ float r = rp * _sigmoid(rg) * m;+ float o = _sigmoid(og);+ tile_l[threadIdx.y + (i << 3)][threadIdx.x] = __float2half_rn(l);+ tile_r[threadIdx.y + (i << 3)][threadIdx.x] = __float2half_rn(r);+ tile_o[threadIdx.y + (i << 3)][threadIdx.x] = __float2half_rn(o);+ }+ }+ __syncthreads();- @cute.kernel- def kernel_bnd(self, a: "cute.Tensor", b: "cute.Tensor", c: "cute.Tensor", n: int):- tx, _, _ = cute.arch.thread_idx()- bx, by, bz = cute.arch.block_idx()+ int xt = (blockIdx.y << 5) + threadIdx.x;+ int yt = (blockIdx.x << 5) + threadIdx.y;- tid = tx- lane_m = tid >> 4- lane_n = tid & 15+ #pragma unroll+ for (int i = 0; i < 4; ++i) {+ int h = yt + (i << 3);+ int row = xt;+ if (h < hidden && row < rows) {+ int b = row / (n * n);+ int local = row - b * (n * n);+ int ii = local / n;+ int jj = local - ii * n;+ long long out_idx = ((long long)(b * hidden + h) * n + ii) * n + jj;+ left_p[out_idx] = tile_l[threadIdx.x][threadIdx.y + (i << 3)];+ right_p[out_idx] = tile_r[threadIdx.x][threadIdx.y + (i << 3)];+ og_p[out_idx] = tile_o[threadIdx.x][threadIdx.y + (i << 3)];+ }+ }+ }- base_m = by * _TMN- base_n = bx * _TMN+ __global__ void ln2_gate_pack_h128(+ const half* __restrict__ out_p,+ const half* __restrict__ og_p,+ const float* __restrict__ w,+ const float* __restrict__ b,+ half* __restrict__ out2,+ int rows_in_b) {+ int tx = (int)threadIdx.x;+ int ty = (int)threadIdx.y;+ int row_base = (int)blockIdx.x << 5;+ int bid = (int)blockIdx.y;- i0 = base_m + lane_m- j0 = base_n + lane_n- i1 = i0 + 16- j1 = j0 + 16+ int row = row_base + tx;- acc00 = cutlass.Float32(0.0)- acc01 = cutlass.Float32(0.0)- acc10 = cutlass.Float32(0.0)- acc11 = cutlass.Float32(0.0)+ __shared__ half sv[128][33];+ __shared__ half sg[128][33];- smem_stride = _TK + 1- smem_a_elems = _TMN * smem_stride- smem_b_elems = _TMN * smem_stride- smem_ptr = cute.arch.alloc_smem(cutlass.Float32, smem_a_elems + smem_b_elems, alignment=16)+ float sum = 0.0f;+ float sum2 = 0.0f;- sA = cute.make_tensor(- smem_ptr,- cute.make_layout((_TMN, _TK), stride=(smem_stride, 1)),- )- sB = cute.make_tensor(- smem_ptr + smem_a_elems,- cute.make_layout((_TMN, _TK), stride=(smem_stride, 1)),- )+ #pragma unroll+ for (int h_base = 0; h_base < 128; h_base += 32) {+ #pragma unroll+ for (int i = 0; i < 4; ++i) {+ int h = h_base + ty + (i << 3);+ if (row < rows_in_b) {+ long long idx = ((long long)(bid * 128 + h) * rows_in_b) + row;+ half hv = out_p[idx];+ half hg = og_p[idx];+ sv[h][tx] = hv;+ sg[h][tx] = hg;+ float v = __half2float(hv);+ sum += v;+ sum2 += v * v;+ }+ }+ }- for k0, (acc00, acc01, acc10, acc11), (acc00_out, acc01_out, acc10_out, acc11_out) in for_generate(- 0,- n,- _TK,- iter_args=[acc00, acc01, acc10, acc11],- ):- base = tid << 2- for t in range(4):- idx = base + t- row = idx >> 5- col = idx & 31+ __shared__ float sh_sum[8][32];+ __shared__ float sh_sum2[8][32];+ __shared__ float sh_mean[32];+ __shared__ float sh_inv[32];- gi_a = base_m + row- gk_a = k0 + col- gj_b = base_n + row- gk_b = k0 + col+ sh_sum[ty][tx] = sum;+ sh_sum2[ty][tx] = sum2;+ __syncthreads();- def _ld_a():- sA[row, col] = a[bz, gi_a, gk_a]+ if (ty == 0 && row < rows_in_b) {+ float s = 0.0f;+ float s2 = 0.0f;+ #pragma unroll+ for (int t = 0; t < 8; ++t) {+ s += sh_sum[t][tx];+ s2 += sh_sum2[t][tx];+ }+ float mean = s * (1.0f / 128.0f);+ float var = fmaxf(s2 * (1.0f / 128.0f) - mean * mean, 0.0f);+ sh_mean[tx] = mean;+ sh_inv[tx] = rsqrtf(var + 1e-5f);+ }+ __syncthreads();- def _ze_a():- sA[row, col] = cutlass.Float32(0.0)+ #pragma unroll+ for (int h_base = 0; h_base < 128; h_base += 32) {+ #pragma unroll+ for (int i = 0; i < 4; ++i) {+ int row_off = ty + (i << 3);+ int out_row = row_base + row_off;+ int h = h_base + tx;+ if (out_row < rows_in_b) {+ float v = __half2float(sv[h][row_off]);+ float g = __half2float(sg[h][row_off]);+ float mean = sh_mean[row_off];+ float inv = sh_inv[row_off];+ float nv = (v - mean) * inv * w[h] + b[h];+ out2[((long long)(bid * rows_in_b + out_row) * 128) + h] = __float2half_rn(nv * g);+ }+ }+ }+ }- def _ld_b():- sB[row, col] = b[bz, gj_b, gk_b]+ __global__ void ln2_gate_pack_tiled(+ const half* __restrict__ out_p,+ const half* __restrict__ og_p,+ const float* __restrict__ w,+ const float* __restrict__ b,+ half* __restrict__ out2,+ int rows_in_b,+ int hidden) {+ int tx = (int)threadIdx.x;+ int ty = (int)threadIdx.y;+ int row_base = (int)blockIdx.x << 5;+ int bid = (int)blockIdx.y;- def _ze_b():- sB[row, col] = cutlass.Float32(0.0)+ int row = row_base + tx;- if_generate((gi_a < n) & (gk_a < n), _ld_a, _ze_a)- if_generate((gj_b < n) & (gk_b < n), _ld_b, _ze_b)+ float sum = 0.0f;+ float sum2 = 0.0f;- cute.arch.sync_threads()+ if (row < rows_in_b) {+ for (int h = ty; h < hidden; h += 8) {+ long long idx = ((long long)(bid * hidden + h) * rows_in_b) + row;+ float v = __half2float(out_p[idx]);+ sum += v;+ sum2 += v * v;+ }+ }- for kk in range(_TK):- a0 = sA[lane_m, kk]- a1 = sA[lane_m + 16, kk]- b0 = sB[lane_n, kk]- b1 = sB[lane_n + 16, kk]+ __shared__ float sh_sum[8][32];+ __shared__ float sh_sum2[8][32];+ __shared__ float sh_mean[32];+ __shared__ float sh_inv[32];+ sh_sum[ty][tx] = sum;+ sh_sum2[ty][tx] = sum2;+ __syncthreads();- acc00 = acc00 + a0 * b0- acc01 = acc01 + a0 * b1- acc10 = acc10 + a1 * b0- acc11 = acc11 + a1 * b1+ if (ty == 0 && row < rows_in_b) {+ float s = 0.0f;+ float s2 = 0.0f;+ #pragma unroll+ for (int t = 0; t < 8; ++t) {+ s += sh_sum[t][tx];+ s2 += sh_sum2[t][tx];+ }+ float mean = s / (float)hidden;+ float var = fmaxf(s2 / (float)hidden - mean * mean, 0.0f);+ sh_mean[tx] = mean;+ sh_inv[tx] = rsqrtf(var + 1e-5f);+ }+ __syncthreads();- cute.arch.sync_threads()- yield_out([acc00, acc01, acc10, acc11])+ __shared__ half tile_v[32][33];+ __shared__ half tile_g[32][33];- def _st00():- c[bz, i0, j0] = acc00_out+ for (int h_base = 0; h_base < hidden; h_base += 32) {+ #pragma unroll+ for (int i = 0; i < 4; ++i) {+ int h = h_base + ty + (i << 3);+ if (row < rows_in_b && h < hidden) {+ long long idx = ((long long)(bid * hidden + h) * rows_in_b) + row;+ tile_v[ty + (i << 3)][tx] = out_p[idx];+ tile_g[ty + (i << 3)][tx] = og_p[idx];+ }+ }+ __syncthreads();- def _st01():- c[bz, i0, j1] = acc01_out+ #pragma unroll+ for (int i = 0; i < 4; ++i) {+ int row_off = ty + (i << 3);+ int out_row = row_base + row_off;+ int h = h_base + tx;+ if (out_row < rows_in_b && h < hidden) {+ float v = __half2float(tile_v[tx][row_off]);+ float g = __half2float(tile_g[tx][row_off]);+ float mean = sh_mean[row_off];+ float inv = sh_inv[row_off];+ float nv = (v - mean) * inv * w[h] + b[h];+ out2[((long long)(bid * rows_in_b + out_row) * hidden) + h] = __float2half_rn(nv * g);+ }+ }+ __syncthreads();+ }+ }- def _st10():- c[bz, i1, j0] = acc10_out+ struct LtPlanKey {+ int m, n, k;+ int batch;+ int op_b;+ int a_type, b_type, c_type, d_type;+ };- def _st11():- c[bz, i1, j1] = acc11_out+ struct LtPlan {+ bool valid;+ LtPlanKey key;+ cublasLtMatmulAlgo_t algo;+ size_t work_bytes;+ cublasLtMatmulDesc_t op_desc;+ cublasLtMatrixLayout_t a_desc;+ cublasLtMatrixLayout_t b_desc;+ cublasLtMatrixLayout_t c_desc;+ cublasLtMatrixLayout_t d_desc;+ };- if_generate((i0 < n) & (j0 < n), _st00)- if_generate((i0 < n) & (j1 < n), _st01)- if_generate((i1 < n) & (j0 < n), _st10)- if_generate((i1 < n) & (j1 < n), _st11)+ static cublasLtHandle_t _lt = nullptr;+ static LtPlan _plan_g1 = {false};+ static LtPlan _plan_g2 = {false};+ static LtPlan _plan_ct = {false};+ static inline bool _key_eq(const LtPlanKey& a, const LtPlanKey& b) {+ return a.m==b.m && a.n==b.n && a.k==b.k && a.batch==b.batch && a.op_b==b.op_b &&+ a.a_type==b.a_type && a.b_type==b.b_type && a.c_type==b.c_type && a.d_type==b.d_type;+ }- _CONTRACT = _BatchedGemmF32_32x32x32()- _CONTRACT_C = None+ static inline void _lt_init() {+ if (_lt) return;+ auto st = cublasLtCreate(&_lt);+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "cublasLtCreate failed");+ }+ static inline void _lt_plan_destroy(LtPlan* plan) {+ if (!plan->valid) return;+ if (plan->a_desc) cublasLtMatrixLayoutDestroy(plan->a_desc);+ if (plan->b_desc) cublasLtMatrixLayoutDestroy(plan->b_desc);+ if (plan->c_desc) cublasLtMatrixLayoutDestroy(plan->c_desc);+ if (plan->d_desc) cublasLtMatrixLayoutDestroy(plan->d_desc);+ if (plan->op_desc) cublasLtMatmulDescDestroy(plan->op_desc);+ plan->a_desc = nullptr;+ plan->b_desc = nullptr;+ plan->c_desc = nullptr;+ plan->d_desc = nullptr;+ plan->op_desc = nullptr;+ plan->valid = false;+ }- def _get_contract_compiled():- global _CONTRACT_C- if _CONTRACT_C is not None:- return _CONTRACT_C+ static inline cublasLtMatrixLayout_t _lt_make_layout(+ cudaDataType type,+ int rows,+ int cols,+ int ld,+ int batch,+ long long stride,+ cublasLtOrder_t order) {+ cublasLtMatrixLayout_t out = nullptr;+ auto st = cublasLtMatrixLayoutCreate(&out, type, rows, cols, ld);+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "cublasLtMatrixLayoutCreate failed");+ st = cublasLtMatrixLayoutSetAttribute(out, CUBLASLT_MATRIX_LAYOUT_ORDER, &order, sizeof(order));+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "layout set order failed");+ if (batch > 1) {+ st = cublasLtMatrixLayoutSetAttribute(out, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, &batch, sizeof(batch));+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "layout set batch failed");+ st = cublasLtMatrixLayoutSetAttribute(out, CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET, &stride, sizeof(stride));+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "layout set stride failed");+ }+ return out;+ }- a_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)- b_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)- c_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)- _CONTRACT_C = cute.compile(- _CONTRACT,- a_ptr,- b_ptr,- c_ptr,- (0, 0),- options="--opt-level 3",+ static inline void _lt_pick_algo(+ LtPlan* plan,+ const LtPlanKey& key,+ cublasOperation_t op_a,+ cublasOperation_t op_b,+ cudaDataType a_type,+ cudaDataType b_type,+ cudaDataType c_type,+ cudaDataType d_type,+ int lda, int ldb, int ldc, int ldd,+ int batch,+ long long stride_a,+ long long stride_b,+ long long stride_c,+ long long stride_d,+ size_t work_bytes) {+ _lt_init();++ if (plan->valid) {+ _lt_plan_destroy(plan);+ }++ auto st = cublasLtMatmulDescCreate(&plan->op_desc, CUBLAS_COMPUTE_32F, CUDA_R_32F);+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "matmul desc create failed");+ st = cublasLtMatmulDescSetAttribute(plan->op_desc, CUBLASLT_MATMUL_DESC_TRANSA, &op_a, sizeof(op_a));+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "set transa failed");+ st = cublasLtMatmulDescSetAttribute(plan->op_desc, CUBLASLT_MATMUL_DESC_TRANSB, &op_b, sizeof(op_b));+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "set transb failed");++ cublasLtOrder_t order = CUBLASLT_ORDER_ROW;+ int a_rows = (op_a == CUBLAS_OP_N) ? key.m : key.k;+ int a_cols = (op_a == CUBLAS_OP_N) ? key.k : key.m;+ int b_rows = (op_b == CUBLAS_OP_N) ? key.k : key.n;+ int b_cols = (op_b == CUBLAS_OP_N) ? key.n : key.k;+ plan->a_desc = _lt_make_layout(a_type, a_rows, a_cols, lda, batch, stride_a, order);+ plan->b_desc = _lt_make_layout(b_type, b_rows, b_cols, ldb, batch, stride_b, order);+ plan->c_desc = _lt_make_layout(c_type, key.m, key.n, ldc, batch, stride_c, order);+ plan->d_desc = _lt_make_layout(d_type, key.m, key.n, ldd, batch, stride_d, order);++ cublasLtMatmulPreference_t pref = nullptr;+ st = cublasLtMatmulPreferenceCreate(&pref);+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "pref create failed");+ st = cublasLtMatmulPreferenceSetAttribute(pref, CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, &work_bytes, sizeof(work_bytes));+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "pref set failed");++ cublasLtMatmulHeuristicResult_t heur;+ int got = 0;+ st = cublasLtMatmulAlgoGetHeuristic(_lt, plan->op_desc, plan->a_desc, plan->b_desc, plan->c_desc, plan->d_desc, pref, 1, &heur, &got);+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS && got > 0, "no cublasLt heuristic algo");++ plan->valid = true;+ plan->key = key;+ plan->algo = heur.algo;+ plan->work_bytes = work_bytes;++ cublasLtMatmulPreferenceDestroy(pref);+ }++ static inline void _lt_matmul(+ LtPlan* plan,+ const LtPlanKey& key,+ cublasOperation_t op_a,+ cublasOperation_t op_b,+ const void* a,+ const void* b,+ const void* c,+ void* d,+ cudaDataType a_type,+ cudaDataType b_type,+ cudaDataType c_type,+ cudaDataType d_type,+ int lda, int ldb, int ldc, int ldd,+ int batch,+ long long stride_a,+ long long stride_b,+ long long stride_c,+ long long stride_d,+ void* work,+ size_t work_bytes) {+ _lt_init();+ if (!plan->valid || !_key_eq(plan->key, key)) {+ _lt_pick_algo(plan, key, op_a, op_b, a_type, b_type, c_type, d_type, lda, ldb, ldc, ldd, batch, stride_a, stride_b, stride_c, stride_d, work_bytes);+ }++ float alpha = 1.0f;+ float beta = 0.0f;+ auto st = cublasLtMatmul(_lt, plan->op_desc, &alpha, a, plan->a_desc, b, plan->b_desc, &beta, c, plan->c_desc, d, plan->d_desc, &plan->algo, work, plan->work_bytes, 0);+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "cublasLtMatmul failed");+ }++ torch::Tensor trimul_forward(+ torch::Tensor x,+ torch::Tensor mask_h,+ torch::Tensor norm_w,+ torch::Tensor norm_b,+ torch::Tensor w_stack,+ torch::Tensor out_norm_w,+ torch::Tensor out_norm_b,+ torch::Tensor w_to_out,+ torch::Tensor work_u8,+ torch::Tensor x_norm_h,+ torch::Tensor y0_h,+ torch::Tensor left_packed,+ torch::Tensor right_packed,+ torch::Tensor out_gate_packed,+ torch::Tensor out_packed,+ torch::Tensor out2_h,+ torch::Tensor y_out_f) {+ CHECK_CUDA(x);+ CHECK_CUDA(mask_h);+ CHECK_CUDA(norm_w);+ CHECK_CUDA(norm_b);+ CHECK_CUDA(w_stack);+ CHECK_CUDA(out_norm_w);+ CHECK_CUDA(out_norm_b);+ CHECK_CUDA(w_to_out);+ CHECK_CUDA(work_u8);+ CHECK_CUDA(x_norm_h);+ CHECK_CUDA(y0_h);+ CHECK_CUDA(left_packed);+ CHECK_CUDA(right_packed);+ CHECK_CUDA(out_gate_packed);+ CHECK_CUDA(out_packed);+ CHECK_CUDA(out2_h);+ CHECK_CUDA(y_out_f);++ CHECK_CONTIGUOUS(x);+ CHECK_CONTIGUOUS(mask_h);+ CHECK_CONTIGUOUS(norm_w);+ CHECK_CONTIGUOUS(norm_b);+ CHECK_CONTIGUOUS(w_stack);+ CHECK_CONTIGUOUS(out_norm_w);+ CHECK_CONTIGUOUS(out_norm_b);+ CHECK_CONTIGUOUS(w_to_out);+ CHECK_CONTIGUOUS(work_u8);+ CHECK_CONTIGUOUS(x_norm_h);+ CHECK_CONTIGUOUS(y0_h);+ CHECK_CONTIGUOUS(left_packed);+ CHECK_CONTIGUOUS(right_packed);+ CHECK_CONTIGUOUS(out_gate_packed);+ CHECK_CONTIGUOUS(out_packed);+ CHECK_CONTIGUOUS(out2_h);+ CHECK_CONTIGUOUS(y_out_f);++ CHECK_DTYPE(x, torch::kFloat32);+ CHECK_DTYPE(mask_h, torch::kFloat16);+ CHECK_DTYPE(norm_w, torch::kFloat32);+ CHECK_DTYPE(norm_b, torch::kFloat32);+ CHECK_DTYPE(w_stack, torch::kFloat16);+ CHECK_DTYPE(out_norm_w, torch::kFloat32);+ CHECK_DTYPE(out_norm_b, torch::kFloat32);+ CHECK_DTYPE(w_to_out, torch::kFloat16);+ CHECK_DTYPE(work_u8, torch::kUInt8);+ CHECK_DTYPE(x_norm_h, torch::kFloat16);+ CHECK_DTYPE(y0_h, torch::kFloat16);+ CHECK_DTYPE(left_packed, torch::kFloat16);+ CHECK_DTYPE(right_packed, torch::kFloat16);+ CHECK_DTYPE(out_gate_packed, torch::kFloat16);+ CHECK_DTYPE(out_packed, torch::kFloat16);+ CHECK_DTYPE(out2_h, torch::kFloat16);+ CHECK_DTYPE(y_out_f, torch::kFloat32);++ TORCH_CHECK(x.dim() == 4, "x must be [bs,N,N,dim]");+ int bs = (int)x.size(0);+ int n = (int)x.size(1);+ int dim = (int)x.size(3);+ TORCH_CHECK((int)x.size(2) == n, "x must be square on N");+ TORCH_CHECK((int)mask_h.size(0) == bs && (int)mask_h.size(1) == n && (int)mask_h.size(2) == n, "mask shape");++ int hidden5 = (int)w_stack.size(1);+ TORCH_CHECK(hidden5 % 5 == 0, "w_stack second dim must be 5*hidden");+ int hidden = hidden5 / 5;++ int rows = bs * n * n;+ int rows_in_b = n * n;++ TORCH_CHECK((int)x_norm_h.size(0) == rows && (int)x_norm_h.size(1) == dim, "x_norm_h shape");+ TORCH_CHECK((dim & 3) == 0, "dim must be multiple of 4");+ int warps = 8;+ dim3 block1(32 * warps, 1, 1);+ dim3 grid1((rows + warps - 1) / warps, 1, 1);+ ln_fwd_fp16<<<grid1, block1>>>(+ (const float*)x.data_ptr<float>(),+ (half*)x_norm_h.data_ptr<at::Half>(),+ (const float*)norm_w.data_ptr<float>(),+ (const float*)norm_b.data_ptr<float>(),+ rows, dim);++ TORCH_CHECK((int)y0_h.size(0) == rows && (int)y0_h.size(1) == 5 * hidden, "y0_h shape");+ LtPlanKey k1;+ k1.m = rows; k1.n = 5 * hidden; k1.k = dim; k1.batch = 1; k1.op_b = 0;+ k1.a_type = (int)CUDA_R_16F; k1.b_type = (int)CUDA_R_16F; k1.c_type = (int)CUDA_R_16F; k1.d_type = (int)CUDA_R_16F;+ _lt_matmul(+ &_plan_g1, k1,+ CUBLAS_OP_N, CUBLAS_OP_N,+ x_norm_h.data_ptr<at::Half>(),+ w_stack.data_ptr<at::Half>(),+ y0_h.data_ptr<at::Half>(),+ y0_h.data_ptr<at::Half>(),+ CUDA_R_16F, CUDA_R_16F, CUDA_R_16F, CUDA_R_16F,+ dim, 5 * hidden, 5 * hidden, 5 * hidden,+ 1, 0, 0, 0, 0,+ work_u8.data_ptr(), (size_t)work_u8.numel());++ TORCH_CHECK((int)left_packed.size(0) == bs * hidden && (int)left_packed.size(1) == n && (int)left_packed.size(2) == n, "left_packed shape");+ TORCH_CHECK(left_packed.sizes() == right_packed.sizes(), "right_packed shape");+ TORCH_CHECK(left_packed.sizes() == out_gate_packed.sizes(), "out_gate_packed shape");+ dim3 block2(32, 8, 1);+ dim3 grid2((hidden + 31) / 32, (rows + 31) / 32, 1);+ pack_lr_og<<<grid2, block2>>>(+ (const half*)y0_h.data_ptr<at::Half>(),+ (const half*)mask_h.data_ptr<at::Half>(),+ (half*)left_packed.data_ptr<at::Half>(),+ (half*)right_packed.data_ptr<at::Half>(),+ (half*)out_gate_packed.data_ptr<at::Half>(),+ rows, n, hidden);++ TORCH_CHECK(out_packed.sizes() == left_packed.sizes(), "out_packed shape");+ int batch_ct = bs * hidden;+ LtPlanKey kc;+ kc.m = n; kc.n = n; kc.k = n; kc.batch = batch_ct; kc.op_b = 1;+ kc.a_type = (int)CUDA_R_16F; kc.b_type = (int)CUDA_R_16F; kc.c_type = (int)CUDA_R_16F; kc.d_type = (int)CUDA_R_16F;+ long long stride_mat = (long long)n * n;+ _lt_matmul(+ &_plan_ct, kc,+ CUBLAS_OP_N, CUBLAS_OP_T,+ left_packed.data_ptr<at::Half>(),+ right_packed.data_ptr<at::Half>(),+ out_packed.data_ptr<at::Half>(),+ out_packed.data_ptr<at::Half>(),+ CUDA_R_16F, CUDA_R_16F, CUDA_R_16F, CUDA_R_16F,+ n, n, n, n,+ batch_ct,+ stride_mat, stride_mat, stride_mat, stride_mat,+ work_u8.data_ptr(), (size_t)work_u8.numel());++ TORCH_CHECK((int)out2_h.size(0) == rows && (int)out2_h.size(1) == hidden, "out2_h shape");+ TORCH_CHECK((int)out_norm_w.numel() == hidden && (int)out_norm_b.numel() == hidden, "out_norm weight/bias shape");+ dim3 block3(32, 8, 1);+ dim3 grid3((rows_in_b + 31) / 32, bs, 1);+ if (hidden == 128) {+ ln2_gate_pack_h128<<<grid3, block3>>>(+ (const half*)out_packed.data_ptr<at::Half>(),+ (const half*)out_gate_packed.data_ptr<at::Half>(),+ (const float*)out_norm_w.data_ptr<float>(),+ (const float*)out_norm_b.data_ptr<float>(),+ (half*)out2_h.data_ptr<at::Half>(),+ rows_in_b);+ } else {+ ln2_gate_pack_tiled<<<grid3, block3>>>(+ (const half*)out_packed.data_ptr<at::Half>(),+ (const half*)out_gate_packed.data_ptr<at::Half>(),+ (const float*)out_norm_w.data_ptr<float>(),+ (const float*)out_norm_b.data_ptr<float>(),+ (half*)out2_h.data_ptr<at::Half>(),+ rows_in_b, hidden);+ }++ TORCH_CHECK((int)y_out_f.size(0) == rows && (int)y_out_f.size(1) == dim, "y_out_f shape");+ LtPlanKey k2;+ k2.m = rows; k2.n = dim; k2.k = hidden; k2.batch = 1; k2.op_b = 0;+ k2.a_type = (int)CUDA_R_16F; k2.b_type = (int)CUDA_R_16F; k2.c_type = (int)CUDA_R_32F; k2.d_type = (int)CUDA_R_32F;+ _lt_matmul(+ &_plan_g2, k2,+ CUBLAS_OP_N, CUBLAS_OP_N,+ out2_h.data_ptr<at::Half>(),+ w_to_out.data_ptr<at::Half>(),+ y_out_f.data_ptr<float>(),+ y_out_f.data_ptr<float>(),+ CUDA_R_16F, CUDA_R_16F, CUDA_R_32F, CUDA_R_32F,+ hidden, dim, dim, dim,+ 1, 0, 0, 0, 0,+ work_u8.data_ptr(), (size_t)work_u8.numel());++ return y_out_f;+ }+ """++ name = "trimul_ext_" + hashlib.sha256((cpp_src + cuda_src).encode("utf-8")).hexdigest()[:16]+ _EXT = load_inline(+ name=name,+ cpp_sources=[cpp_src],+ cuda_sources=[cuda_src],+ functions=None,+ extra_cflags=["-O3", "-std=c++17"],+ extra_cuda_cflags=["-O3", "--use_fast_math", "-std=c++17"],+ extra_ldflags=["-lcublasLt", "-lcublas"],+ with_cuda=True,+ build_directory=build_dir,+ verbose=False,)- return _CONTRACT_C+ return _EXT- def _contract_outgoing(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:- bs, n, _, hidden = left.shape+ def _prep_weights(weights: Dict[str, "torch.Tensor"], dim: int, hidden_dim: int, device: "torch.device") -> dict[str, Any]:+ import torch- if not (left.is_cuda and right.is_cuda):- raise RuntimeError("该实现仅支持 CUDA 张量。")- if left.dtype != torch.float32 or right.dtype != torch.float32:- raise RuntimeError("该实现期望 left/right 为 float32。")+ key = (+ int(weights["left_proj.weight"].data_ptr()),+ int(weights["right_proj.weight"].data_ptr()),+ int(weights["left_gate.weight"].data_ptr()),+ int(weights["right_gate.weight"].data_ptr()),+ int(weights["out_gate.weight"].data_ptr()),+ int(weights["to_out.weight"].data_ptr()),+ int(weights["norm.weight"].data_ptr()),+ int(weights["norm.bias"].data_ptr()),+ int(weights["to_out_norm.weight"].data_ptr()),+ int(weights["to_out_norm.bias"].data_ptr()),+ device.index if device.type == "cuda" else -1,+ )+ cached = _WEIGHT_CACHE.get(key)+ if cached is not None:+ return cached-- left_t = left.permute(0, 3, 1, 2).contiguous()- right_t = right.permute(0, 3, 1, 2).contiguous()+ def _f32_vec(t: "torch.Tensor") -> "torch.Tensor":+ if t.device == device and t.dtype == torch.float32 and t.is_contiguous():+ return t+ return t.to(device=device, dtype=torch.float32).contiguous()- bh = bs * hidden- a = left_t.view(bh, n, n)- b = right_t.view(bh, n, n)- c = torch.empty((bh, n, n), device=left.device, dtype=torch.float32)+ def _to_half_t(w: "torch.Tensor") -> "torch.Tensor":+ return w.to(device=device, dtype=torch.float16).t().contiguous()- compiled = _get_contract_compiled()- a_ptr = make_ptr(cutlass.Float32, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)- b_ptr = make_ptr(cutlass.Float32, b.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)- c_ptr = make_ptr(cutlass.Float32, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)- compiled(a_ptr, b_ptr, c_ptr, (bh, n))+ w_stack = torch.cat(+ [+ _to_half_t(weights["left_proj.weight"]),+ _to_half_t(weights["right_proj.weight"]),+ _to_half_t(weights["left_gate.weight"]),+ _to_half_t(weights["right_gate.weight"]),+ _to_half_t(weights["out_gate.weight"]),+ ],+ dim=1,+ ).contiguous()- out = c.view(bs, hidden, n, n).permute(0, 2, 3, 1).contiguous()+ w_to_out = _to_half_t(weights["to_out.weight"])++ out = {+ "norm_w": _f32_vec(weights["norm.weight"]),+ "norm_b": _f32_vec(weights["norm.bias"]),+ "out_norm_w": _f32_vec(weights["to_out_norm.weight"]),+ "out_norm_b": _f32_vec(weights["to_out_norm.bias"]),+ "w_stack": w_stack,+ "w_to_out": w_to_out,+ }+ _WEIGHT_CACHE[key] = outreturn out- @torch.inference_mode()- def custom_kernel(data: Tuple[torch.Tensor, torch.Tensor, Dict[str, torch.Tensor], Dict[str, Any]]) -> torch.Tensor:- x, mask, weights, config = data+ def _get_scratch(device: "torch.device", bs: int, n: int, dim: int, hidden: int) -> dict[str, Any]:+ import torch+ key = (device.type, device.index if device.type == "cuda" else -1, bs, n, dim, hidden)+ cached = _KERNEL_CACHE.get(key)+ if cached is not None:+ return cached++ rows = bs * n * n+ work_bytes = 32 * 1024 * 1024+ scratch = {+ "work": torch.empty((work_bytes,), device=device, dtype=torch.uint8),+ "x_norm": torch.empty((rows, dim), device=device, dtype=torch.float16),+ "y0": torch.empty((rows, 5 * hidden), device=device, dtype=torch.float16),+ "left_p": torch.empty((bs * hidden, n, n), device=device, dtype=torch.float16),+ "right_p": torch.empty((bs * hidden, n, n), device=device, dtype=torch.float16),+ "og_p": torch.empty((bs * hidden, n, n), device=device, dtype=torch.float16),+ "out_p": torch.empty((bs * hidden, n, n), device=device, dtype=torch.float16),+ "out2": torch.empty((rows, hidden), device=device, dtype=torch.float16),+ "y_out": torch.empty((rows, dim), device=device, dtype=torch.float32),+ }+ _KERNEL_CACHE[key] = scratch+ return scratch+++ def _ensure_mask_half(mask: "torch.Tensor", device: "torch.device") -> "torch.Tensor":+ import torch++ if mask.dtype == torch.float16 and mask.is_contiguous() and mask.device == device:+ return mask+ return mask.to(device=device, dtype=torch.float16).contiguous()+++ def custom_kernel(data: Tuple["torch.Tensor", "torch.Tensor", Dict[str, "torch.Tensor"], Dict[str, Any]]) -> "torch.Tensor":+ import torch++ x, mask, weights, config = datadim = int(config["dim"])hidden_dim = int(config["hidden_dim"])if not x.is_cuda:- raise RuntimeError("该实现要求 x 在 CUDA 上。")-+ raise RuntimeError("CUDA only")if x.dtype != torch.float32:x = x.to(dtype=torch.float32)+ if not x.is_contiguous():+ x = x.contiguous()- try:- torch.backends.cuda.matmul.allow_tf32 = True- torch.backends.cudnn.allow_tf32 = True- except Exception:- pass+ bs, n0, n1, d0 = x.shape+ if n0 != n1 or d0 != dim:+ raise RuntimeError("shape mismatch")- x = F.layer_norm(x, (dim,), weights["norm.weight"], weights["norm.bias"], 1e-5)+ mask_h = _ensure_mask_half(mask, x.device)- left = F.linear(x, weights["left_proj.weight"], None)- right = F.linear(x, weights["right_proj.weight"], None)+ ext = _load_ext()+ wpack = _prep_weights(weights, dim=dim, hidden_dim=hidden_dim, device=x.device)+ scratch = _get_scratch(x.device, bs=bs, n=n0, dim=dim, hidden=hidden_dim)- mask_f = mask.unsqueeze(-1)- if mask_f.dtype != left.dtype:- mask_f = mask_f.to(dtype=left.dtype)- left.mul_(mask_f)- right.mul_(mask_f)+ y_flat = ext.forward(+ x,+ mask_h,+ wpack["norm_w"],+ wpack["norm_b"],+ wpack["w_stack"],+ wpack["out_norm_w"],+ wpack["out_norm_b"],+ wpack["w_to_out"],+ scratch["work"],+ scratch["x_norm"],+ scratch["y0"],+ scratch["left_p"],+ scratch["right_p"],+ scratch["og_p"],+ scratch["out_p"],+ scratch["out2"],+ scratch["y_out"],+ )- left_gate = F.linear(x, weights["left_gate.weight"], None)- right_gate = F.linear(x, weights["right_gate.weight"], None)- out_gate = F.linear(x, weights["out_gate.weight"], None)- left_gate.sigmoid_()- right_gate.sigmoid_()- out_gate.sigmoid_()- left.mul_(left_gate)- right.mul_(right_gate)+ return y_flat.view(bs, n0, n0, dim)- out = _contract_outgoing(left, right)- out = F.layer_norm(out, (hidden_dim,), weights["to_out_norm.weight"], weights["to_out_norm.bias"], 1e-5)- out.mul_(out_gate)- out = F.linear(out, weights["to_out.weight"], None)- return out--__all__ = ["custom_kernel"]
scrolls · 1115 diff lines total
Best evidence level for this revision: reported
JSON