submission 417619
shiyegao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 695 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-417619?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:713057414abea34eb71c3d647b29904478e3eee5b54afa51c57345c261021266
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__ float shm_sum[256];vector-width = float4
const float4* x4 = reinterpret_cast<const float4*>(x + row * 128);Kernel source
submission.py695 lines
from __future__ import annotations
from typing import Any, Dict, Tuple
import os
import torch
_EXT = None
_EXT_LOCK = None
def _lazy_import_extension_utils():
from torch.utils.cpp_extension import load_inline
return load_inline
def _get_ext():
global _EXT, _EXT_LOCK
if _EXT is not None:
return _EXT
if _EXT_LOCK is None:
import threading
_EXT_LOCK = threading.Lock()
with _EXT_LOCK:
if _EXT is not None:
return _EXT
load_inline = _lazy_import_extension_utils()
if "TORCH_CUDA_ARCH_LIST" not in os.environ:
os.environ["TORCH_CUDA_ARCH_LIST"] = "9.0a"
cpp_src = r"""
#include <torch/extension.h>
torch::Tensor trimul_fwd(
torch::Tensor x,
torch::Tensor mask,
torch::Tensor ln1_w,
torch::Tensor ln1_b,
torch::Tensor w_cat,
torch::Tensor ln2_w,
torch::Tensor ln2_b,
torch::Tensor w_out,
int64_t dim,
int64_t hidden);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("fwd", &trimul_fwd, "trimul forward (cuda)");
}
"""
cuda_src = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda.h>
#include <cuda_fp16.h>
#include <cublas_v2.h>
#include <mutex>
namespace {
static inline void checkCuda(cudaError_t e, const char* msg) {
if (e != cudaSuccess) {
throw std::runtime_error(std::string(msg) + ": " + cudaGetErrorString(e));
}
}
static inline void checkCublas(cublasStatus_t s, const char* msg) {
if (s != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(std::string(msg) + ": cublas status=" + std::to_string((int)s));
}
}
struct CublasHandleHolder {
cublasHandle_t handle = nullptr;
CublasHandleHolder() {
checkCublas(cublasCreate(&handle), "cublasCreate");
// 默认数学模式即可;这里不做额外设置,减少环境依赖
}
~CublasHandleHolder() {
if (handle) {
cublasDestroy(handle);
handle = nullptr;
}
}
};
static CublasHandleHolder* get_cublas() {
static std::once_flag once;
static CublasHandleHolder* holder = nullptr;
std::call_once(once, []() { holder = new CublasHandleHolder(); });
return holder;
}
__device__ __forceinline__ float warp_sum(float v) {
for (int d = 16; d > 0; d >>= 1) {
v += __shfl_down_sync(0xffffffff, v, d);
}
return v;
}
__device__ __forceinline__ float fast_sigmoid(float x) {
// 使用 __expf,精度在题面容忍范围内通常足够
float z = __expf(-x);
return 1.0f / (1.0f + z);
}
__global__ void ln1_128_f16(
const float* __restrict__ x,
const float* __restrict__ w,
const float* __restrict__ b,
half* __restrict__ y,
int64_t rows) {
int64_t row = (int64_t)blockIdx.x;
if (row >= rows) return;
int lane = (int)threadIdx.x; // 0..31
const float4* x4 = reinterpret_cast<const float4*>(x + row * 128);
float4 v = x4[lane];
float s = v.x + v.y + v.z + v.w;
float ss = v.x * v.x + v.y * v.y + v.z * v.z + v.w * v.w;
s = warp_sum(s);
ss = warp_sum(ss);
// warp_sum 仅保证 lane0 得到全和,需要广播到全 warp
s = __shfl_sync(0xffffffff, s, 0);
ss = __shfl_sync(0xffffffff, ss, 0);
float mean = s * (1.0f / 128.0f);
float var = ss * (1.0f / 128.0f) - mean * mean;
float inv = rsqrtf(var + 1e-5f);
const float4* w4 = reinterpret_cast<const float4*>(w);
const float4* b4 = reinterpret_cast<const float4*>(b);
float4 gw = w4[lane];
float4 gb = b4[lane];
float y0 = (v.x - mean) * inv * gw.x + gb.x;
float y1 = (v.y - mean) * inv * gw.y + gb.y;
float y2 = (v.z - mean) * inv * gw.z + gb.z;
float y3 = (v.w - mean) * inv * gw.w + gb.w;
half2 h0 = __floats2half2_rn(y0, y1);
half2 h1 = __floats2half2_rn(y2, y3);
half2* y2p = reinterpret_cast<half2*>(y + row * 128 + lane * 4);
y2p[0] = h0;
y2p[1] = h1;
}
__global__ void ln1_generic_f16(
const float* __restrict__ x,
const float* __restrict__ w,
const float* __restrict__ b,
half* __restrict__ y,
int dim,
int64_t rows) {
int64_t row = (int64_t)blockIdx.x;
if (row >= rows) return;
// 计算均值与二阶矩
float sum = 0.0f;
float sq = 0.0f;
int64_t base = row * (int64_t)dim;
for (int c = (int)threadIdx.x; c < dim; c += (int)blockDim.x) {
float v = x[base + c];
sum += v;
sq += v * v;
}
__shared__ float shm_sum[256];
__shared__ float shm_sq[256];
int t = (int)threadIdx.x;
shm_sum[t] = sum;
shm_sq[t] = sq;
__syncthreads();
for (int stride = ((int)blockDim.x) / 2; stride > 0; stride >>= 1) {
if (t < stride) {
shm_sum[t] += shm_sum[t + stride];
shm_sq[t] += shm_sq[t + stride];
}
__syncthreads();
}
float mean = shm_sum[0] / (float)dim;
float var = shm_sq[0] / (float)dim - mean * mean;
float inv = rsqrtf(var + 1e-5f);
for (int c = (int)threadIdx.x; c < dim; c += (int)blockDim.x) {
float v = x[base + c];
float yv = (v - mean) * inv * w[c] + b[c];
y[base + c] = __float2half_rn(yv);
}
}
__global__ void pack_proj_f16(
const half* __restrict__ proj,
const float* __restrict__ mask,
half* __restrict__ left,
half* __restrict__ right,
half* __restrict__ ogate,
int bs,
int n,
int hidden) {
int64_t row = (int64_t)blockIdx.x;
int d = (int)threadIdx.x;
if (d >= hidden) return;
int64_t nn = (int64_t)n * (int64_t)n;
int b0 = (int)(row / nn);
int64_t rem = row - (int64_t)b0 * nn;
int i = (int)(rem / n);
int j = (int)(rem - (int64_t)i * n);
float m = mask[row];
int out = hidden * 5;
const half* p = proj + row * out;
float l = __half2float(p[d]);
float r = __half2float(p[hidden + d]);
float gl = fast_sigmoid(__half2float(p[2 * hidden + d]));
float gr = fast_sigmoid(__half2float(p[3 * hidden + d]));
float go = fast_sigmoid(__half2float(p[4 * hidden + d]));
float l2 = l * gl * m;
float r2 = r * gr * m;
int64_t base = (((int64_t)b0 * hidden + d) * n + i) * n + j;
left[base] = __float2half_rn(l2);
right[base] = __float2half_rn(r2);
ogate[base] = __float2half_rn(go);
}
template<int MAX_H>
__global__ void ln2_gate_store_f16(
const float* __restrict__ out_acc,
const half* __restrict__ ogate,
const float* __restrict__ w,
const float* __restrict__ b,
half* __restrict__ out_norm,
int bs,
int n,
int hidden) {
// blockIdx.x 对应 (b,i,j)
int64_t row = (int64_t)blockIdx.x;
int tid = (int)threadIdx.x;
if (tid >= MAX_H) return;
int64_t nn = (int64_t)n * (int64_t)n;
int b0 = (int)(row / nn);
int64_t rem = row - (int64_t)b0 * nn;
int i = (int)(rem / n);
int j = (int)(rem - (int64_t)i * n);
float v = 0.0f;
float vv = 0.0f;
if (tid < hidden) {
int64_t idx = (((int64_t)b0 * hidden + tid) * n + i) * n + j;
float x = out_acc[idx];
v = x;
vv = x * x;
}
// 归约:MAX_H 固定,使用共享内存
__shared__ float shm_sum[MAX_H];
__shared__ float shm_sq[MAX_H];
shm_sum[tid] = v;
shm_sq[tid] = vv;
__syncthreads();
for (int stride = MAX_H / 2; stride > 0; stride >>= 1) {
if (tid < stride) {
shm_sum[tid] += shm_sum[tid + stride];
shm_sq[tid] += shm_sq[tid + stride];
}
__syncthreads();
}
float mean = shm_sum[0] / (float)hidden;
float var = shm_sq[0] / (float)hidden - mean * mean;
float inv = rsqrtf(var + 1e-5f);
if (tid < hidden) {
int64_t idx_in = (((int64_t)b0 * hidden + tid) * n + i) * n + j;
float x = out_acc[idx_in];
float y = (x - mean) * inv * w[tid] + b[tid];
float g = __half2float(ogate[idx_in]);
float z = y * g;
int64_t idx_out = ((row * hidden) + tid);
out_norm[idx_out] = __float2half_rn(z);
}
}
static void launch_ln1(torch::Tensor x, torch::Tensor w, torch::Tensor b, torch::Tensor y) {
int dim = (int)x.size(1);
auto rows = x.size(0);
if (dim == 128) {
dim3 block(32, 1, 1);
dim3 grid((unsigned)rows, 1, 1);
ln1_128_f16<<<grid, block>>>(
(const float*)x.data_ptr(),
(const float*)w.data_ptr(),
(const float*)b.data_ptr(),
(half*)y.data_ptr(),
(int64_t)rows);
checkCuda(cudaGetLastError(), "ln1_128_f16");
} else {
dim3 block(256, 1, 1);
dim3 grid((unsigned)rows, 1, 1);
ln1_generic_f16<<<grid, block>>>(
(const float*)x.data_ptr(),
(const float*)w.data_ptr(),
(const float*)b.data_ptr(),
(half*)y.data_ptr(),
dim,
(int64_t)rows);
checkCuda(cudaGetLastError(), "ln1_generic_f16");
}
}
static void launch_pack(torch::Tensor proj, torch::Tensor mask, torch::Tensor left, torch::Tensor right, torch::Tensor og, int bs, int n, int hidden) {
int64_t rows = (int64_t)bs * (int64_t)n * (int64_t)n;
dim3 block((unsigned)hidden, 1, 1);
dim3 grid((unsigned)rows, 1, 1);
pack_proj_f16<<<grid, block>>>(
(const half*)proj.data_ptr(),
(const float*)mask.data_ptr(),
(half*)left.data_ptr(),
(half*)right.data_ptr(),
(half*)og.data_ptr(),
bs, n, hidden);
checkCuda(cudaGetLastError(), "pack_proj_f16");
}
static void launch_ln2(torch::Tensor out_acc, torch::Tensor og, torch::Tensor w, torch::Tensor b, torch::Tensor out_norm, int bs, int n, int hidden) {
int64_t rows = (int64_t)bs * (int64_t)n * (int64_t)n;
if (hidden <= 32) {
dim3 block(32, 1, 1);
dim3 grid((unsigned)rows, 1, 1);
ln2_gate_store_f16<32><<<grid, block>>>(
(const float*)out_acc.data_ptr(),
(const half*)og.data_ptr(),
(const float*)w.data_ptr(),
(const float*)b.data_ptr(),
(half*)out_norm.data_ptr(),
bs, n, hidden);
checkCuda(cudaGetLastError(), "ln2_gate_store_f16_32");
} else if (hidden <= 64) {
dim3 block(64, 1, 1);
dim3 grid((unsigned)rows, 1, 1);
ln2_gate_store_f16<64><<<grid, block>>>(
(const float*)out_acc.data_ptr(),
(const half*)og.data_ptr(),
(const float*)w.data_ptr(),
(const float*)b.data_ptr(),
(half*)out_norm.data_ptr(),
bs, n, hidden);
checkCuda(cudaGetLastError(), "ln2_gate_store_f16_64");
} else if (hidden <= 128) {
dim3 block(128, 1, 1);
dim3 grid((unsigned)rows, 1, 1);
ln2_gate_store_f16<128><<<grid, block>>>(
(const float*)out_acc.data_ptr(),
(const half*)og.data_ptr(),
(const float*)w.data_ptr(),
(const float*)b.data_ptr(),
(half*)out_norm.data_ptr(),
bs, n, hidden);
checkCuda(cudaGetLastError(), "ln2_gate_store_f16_128");
} else {
throw std::runtime_error("hidden_dim too large");
}
}
static void gemm_x_wt_f16_f16(
cublasHandle_t h,
const half* x_row, // row-major [M,K]
const half* w_row, // row-major [N,K]
half* y_row, // row-major [M,N]
int64_t M,
int64_t N,
int64_t K) {
// 使用列主序 trick:输出按列主序 (N x M) 写入,即等价于 row-major (M x N)
float alpha = 1.0f;
float beta = 0.0f;
checkCublas(
cublasGemmEx(
h,
CUBLAS_OP_T, CUBLAS_OP_N,
(int)N, (int)M, (int)K,
&alpha,
w_row, CUDA_R_16F, (int)K,
x_row, CUDA_R_16F, (int)K,
&beta,
y_row, CUDA_R_16F, (int)N,
CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
"cublasGemmEx");
}
static void gemm_x_wt_f16_f32(
cublasHandle_t h,
const half* x_row, // row-major [M,K]
const half* w_row, // row-major [N,K]
float* y_row, // row-major [M,N]
int64_t M,
int64_t N,
int64_t K) {
float alpha = 1.0f;
float beta = 0.0f;
checkCublas(
cublasGemmEx(
h,
CUBLAS_OP_T, CUBLAS_OP_N,
(int)N, (int)M, (int)K,
&alpha,
w_row, CUDA_R_16F, (int)K,
x_row, CUDA_R_16F, (int)K,
&beta,
y_row, CUDA_R_32F, (int)N,
CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
"cublasGemmEx");
}
static void gemm_contract_batched_f16_f32(
cublasHandle_t h,
const half* left_row, // row-major [B,M,K]
const half* right_row, // row-major [B,N,K] (这里 N=M=K=n)
float* out_row, // row-major [B,M,N]
int batch,
int n) {
float alpha = 1.0f;
float beta = 0.0f;
long long strideA = (long long)n * (long long)n;
long long strideB = (long long)n * (long long)n;
long long strideC = (long long)n * (long long)n;
// 计算 C = L * R^T
// 采用列主序 trick:C^T = R * L^T
checkCublas(
cublasGemmStridedBatchedEx(
h,
CUBLAS_OP_T, CUBLAS_OP_N,
n, n, n,
&alpha,
right_row, CUDA_R_16F, n, strideB,
left_row, CUDA_R_16F, n, strideA,
&beta,
out_row, CUDA_R_32F, n, strideC,
batch,
CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
"cublasGemmStridedBatchedEx");
}
} // namespace
torch::Tensor trimul_fwd(
torch::Tensor x,
torch::Tensor mask,
torch::Tensor ln1_w,
torch::Tensor ln1_b,
torch::Tensor w_cat,
torch::Tensor ln2_w,
torch::Tensor ln2_b,
torch::Tensor w_out,
int64_t dim,
int64_t hidden) {
if (!x.is_cuda() || !mask.is_cuda()) {
throw std::runtime_error("cuda only");
}
if (x.scalar_type() != torch::kFloat32) {
throw std::runtime_error("x must be float32");
}
if (mask.scalar_type() != torch::kFloat32) {
throw std::runtime_error("mask must be float32");
}
if (dim != x.size(3)) {
throw std::runtime_error("dim mismatch");
}
if (w_cat.scalar_type() != torch::kFloat16 || w_out.scalar_type() != torch::kFloat16) {
throw std::runtime_error("weights must be float16");
}
int bs = (int)x.size(0);
int n = (int)x.size(1);
int64_t M = (int64_t)bs * (int64_t)n * (int64_t)n;
auto x2d = x.view({M, dim});
auto xhat = torch::empty({M, dim}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
launch_ln1(x2d, ln1_w, ln1_b, xhat);
int64_t out_ch = hidden * 5;
auto proj = torch::empty({M, out_ch}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
auto* holder = get_cublas();
cublasHandle_t h = holder->handle;
// gemm1: proj = xhat @ w_cat^T
gemm_x_wt_f16_f16(h, (const half*)xhat.data_ptr(), (const half*)w_cat.data_ptr(), (half*)proj.data_ptr(), M, out_ch, dim);
auto left = torch::empty({bs, hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
auto right = torch::empty({bs, hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
auto og = torch::empty({bs, hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
launch_pack(proj, mask.view({M}), left, right, og, bs, n, (int)hidden);
auto left3 = left.view({bs * (int)hidden, n, n});
auto right3 = right.view({bs * (int)hidden, n, n});
auto out_acc = torch::empty({bs * (int)hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat32));
gemm_contract_batched_f16_f32(
h,
(const half*)left3.data_ptr(),
(const half*)right3.data_ptr(),
(float*)out_acc.data_ptr(),
bs * (int)hidden, n);
auto out_norm = torch::empty({M, hidden}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
launch_ln2(out_acc, og.view({bs * (int)hidden, n, n}), ln2_w, ln2_b, out_norm, bs, n, (int)hidden);
// gemm2: y = out_norm @ w_out^T
auto y = torch::empty({M, dim}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat32));
gemm_x_wt_f16_f32(h, (const half*)out_norm.data_ptr(), (const half*)w_out.data_ptr(), (float*)y.data_ptr(), M, dim, hidden);
return y.view({bs, n, n, dim});
}
"""
name = "trimul_ext_mod"
extra_cuda_cflags = [
"-O3",
"--use_fast_math",
]
extra_cflags = [
"-O3",
]
extra_ldflags = [
"-lcublas",
]
_EXT = load_inline(
name=name,
cpp_sources=cpp_src,
cuda_sources=cuda_src,
functions=None,
extra_cflags=extra_cflags,
extra_cuda_cflags=extra_cuda_cflags,
extra_ldflags=extra_ldflags,
with_cuda=True,
verbose=False,
)
return _EXT
class _WeightCache:
__slots__ = ("key", "w_cat", "w_out")
def __init__(self) -> None:
self.key = None
self.w_cat = None
self.w_out = None
_W_CACHE = _WeightCache()
def _prepare_weights(weights: Dict[str, torch.Tensor], dim: int, hidden: int):
k = (
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()),
)
if _W_CACHE.key == k and _W_CACHE.w_cat is not None and _W_CACHE.w_out is not None:
return _W_CACHE.w_cat, _W_CACHE.w_out
w_left = weights["left_proj.weight"]
w_right = weights["right_proj.weight"]
w_lg = weights["left_gate.weight"]
w_rg = weights["right_gate.weight"]
w_og = weights["out_gate.weight"]
w_out = weights["to_out.weight"]
if w_left.shape != (hidden, dim):
raise RuntimeError("left_proj.weight shape mismatch")
if w_right.shape != (hidden, dim):
raise RuntimeError("right_proj.weight shape mismatch")
if w_lg.shape != (hidden, dim):
raise RuntimeError("left_gate.weight shape mismatch")
if w_rg.shape != (hidden, dim):
raise RuntimeError("right_gate.weight shape mismatch")
if w_og.shape != (hidden, dim):
raise RuntimeError("out_gate.weight shape mismatch")
if w_out.shape != (dim, hidden):
raise RuntimeError("to_out.weight shape mismatch")
w_cat = torch.cat([w_left, w_right, w_lg, w_rg, w_og], dim=0).contiguous().to(dtype=torch.float16)
w_out_h = w_out.contiguous().to(dtype=torch.float16)
_W_CACHE.key = k
_W_CACHE.w_cat = w_cat
_W_CACHE.w_out = w_out_h
return w_cat, w_out_h
@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
dim = int(config["dim"])
hidden = int(config["hidden_dim"])
if not x.is_cuda:
raise RuntimeError("x must be CUDA tensor")
if not mask.is_cuda:
raise RuntimeError("mask must be CUDA tensor")
if x.dtype != torch.float32:
x = x.to(dtype=torch.float32)
if mask.dtype != torch.float32:
mask = mask.to(dtype=torch.float32)
x = x.contiguous()
mask = mask.contiguous()
if x.ndim != 4:
raise RuntimeError("x must be 4D")
if mask.ndim != 3:
raise RuntimeError("mask must be 3D")
if x.shape[:3] != mask.shape:
raise RuntimeError("x/mask shape mismatch")
if x.shape[3] != dim:
raise RuntimeError("dim mismatch")
for k in (
"norm.weight",
"norm.bias",
"left_proj.weight",
"right_proj.weight",
"left_gate.weight",
"right_gate.weight",
"out_gate.weight",
"to_out_norm.weight",
"to_out_norm.bias",
"to_out.weight",
):
if not weights[k].is_cuda:
raise RuntimeError(f"weight {k} must be CUDA tensor")
if weights[k].dtype != torch.float32:
raise RuntimeError(f"weight {k} must be float32")
ln1_w = weights["norm.weight"].contiguous()
ln1_b = weights["norm.bias"].contiguous()
if ln1_w.shape != (dim,) or ln1_b.shape != (dim,):
raise RuntimeError("norm params shape mismatch")
ln2_w = weights["to_out_norm.weight"].contiguous()
ln2_b = weights["to_out_norm.bias"].contiguous()
if ln2_w.shape != (hidden,) or ln2_b.shape != (hidden,):
raise RuntimeError("to_out_norm params shape mismatch")
w_cat, w_out = _prepare_weights(weights, dim, hidden)
ext = _get_ext()
return ext.fwd(x, mask, ln1_w, ln1_b, w_cat, ln2_w, ln2_b, w_out, dim, hidden)
__all__ = ["custom_kernel"]
scrolls · 695 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 417301.
⋯ 1 unchanged linesfrom typing import Any, Dict, Tuple+ import os+import torch- import torch.nn.functional as F- import cutlass- import cutlass.cute as cute- from cutlass.cute.runtime import make_ptr- from cutlass.cutlass_dsl import for_generate, yield_out+ _EXT = None+ _EXT_LOCK = None+ def _lazy_import_extension_utils():++ from torch.utils.cpp_extension import load_inline+ return load_inline- _TILE_MN = 32- _TILE_K = 32- _THREADS = 256+ def _get_ext():+ global _EXT, _EXT_LOCK+ if _EXT is not None:+ return _EXT+ if _EXT_LOCK is None:+ import threading- class _OutgoingContractBatchedKernel:- def __init__(self) -> None:- self.threads = _THREADS+ _EXT_LOCK = threading.Lock()+ with _EXT_LOCK:+ if _EXT is not None:+ return _EXT- @cute.jit- def __call__(- self,- a_ptr: "cute.Pointer",- b_ptr: "cute.Pointer",- c_ptr: "cute.Pointer",- problem: tuple,- ):- bh, n = problem+ load_inline = _lazy_import_extension_utils()- stride_bh = n * n- stride_m = n- stride_n = 1++ if "TORCH_CUDA_ARCH_LIST" not in os.environ:+ os.environ["TORCH_CUDA_ARCH_LIST"] = "9.0a"- a = cute.make_tensor(- a_ptr,- cute.make_layout((bh, n, n), stride=(stride_bh, stride_m, stride_n)),- )- b = cute.make_tensor(- b_ptr,- cute.make_layout((bh, n, n), stride=(stride_bh, stride_m, stride_n)),- )- c = cute.make_tensor(- c_ptr,- cute.make_layout((bh, n, n), stride=(stride_bh, stride_m, stride_n)),- )++- grid_x = n // _TILE_MN- grid_y = n // _TILE_MN- grid_z = bh- self.kernel(a, b, c, n).launch(- grid=[grid_x, grid_y, grid_z],- block=[self.threads, 1, 1],- )- return+ cpp_src = r"""+ #include <torch/extension.h>- @cute.kernel- def kernel(- self,- a: "cute.Tensor",- b: "cute.Tensor",- c: "cute.Tensor",- n: int,- ):- tx, _, _ = cute.arch.thread_idx()- bx, by, bz = cute.arch.block_idx()+ torch::Tensor trimul_fwd(+ torch::Tensor x,+ torch::Tensor mask,+ torch::Tensor ln1_w,+ torch::Tensor ln1_b,+ torch::Tensor w_cat,+ torch::Tensor ln2_w,+ torch::Tensor ln2_b,+ torch::Tensor w_out,+ int64_t dim,+ int64_t hidden);- tid = tx- lane_m = tid >> 4- lane_n = tid & 15+ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {+ m.def("fwd", &trimul_fwd, "trimul forward (cuda)");+ }+ """- base_m = by * _TILE_MN- base_n = bx * _TILE_MN+ cuda_src = r"""+ #include <torch/extension.h>+ #include <ATen/cuda/CUDAContext.h>+ #include <cuda.h>+ #include <cuda_fp16.h>+ #include <cublas_v2.h>- i0 = base_m + lane_m- j0 = base_n + lane_n- i1 = i0 + 16- j1 = j0 + 16+ #include <mutex>- acc00 = cutlass.Float32(0.0)- acc01 = cutlass.Float32(0.0)- acc10 = cutlass.Float32(0.0)- acc11 = cutlass.Float32(0.0)+ namespace {-- smem_stride = _TILE_K + 1- smem_a_elems = _TILE_MN * smem_stride- smem_b_elems = _TILE_MN * smem_stride- smem_ptr = cute.arch.alloc_smem(cutlass.Float32, smem_a_elems + smem_b_elems)+ static inline void checkCuda(cudaError_t e, const char* msg) {+ if (e != cudaSuccess) {+ throw std::runtime_error(std::string(msg) + ": " + cudaGetErrorString(e));+ }+ }- sA = cute.make_tensor(- smem_ptr,- cute.make_layout((_TILE_MN, _TILE_K), stride=(smem_stride, 1)),- )- sB = cute.make_tensor(- smem_ptr + smem_a_elems,- cute.make_layout((_TILE_MN, _TILE_K), stride=(smem_stride, 1)),- )+ static inline void checkCublas(cublasStatus_t s, const char* msg) {+ if (s != CUBLAS_STATUS_SUCCESS) {+ throw std::runtime_error(std::string(msg) + ": cublas status=" + std::to_string((int)s));+ }+ }- for k0, (acc00, acc01, acc10, acc11), (acc00_out, acc01_out, acc10_out, acc11_out) in for_generate(- 0,- n,- _TILE_K,- 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]+ struct CublasHandleHolder {+ cublasHandle_t handle = nullptr;+ CublasHandleHolder() {+ checkCublas(cublasCreate(&handle), "cublasCreate");+ // 默认数学模式即可;这里不做额外设置,减少环境依赖+ }+ ~CublasHandleHolder() {+ if (handle) {+ cublasDestroy(handle);+ handle = nullptr;+ }+ }+ };- cute.arch.sync_threads()+ static CublasHandleHolder* get_cublas() {+ static std::once_flag once;+ static CublasHandleHolder* holder = nullptr;+ std::call_once(once, []() { holder = new CublasHandleHolder(); });+ return holder;+ }-- for kk in range(_TILE_K):- a0 = sA[lane_m, kk]- a1 = sA[lane_m + 16, kk]- b0 = sB[lane_n, kk]- b1 = sB[lane_n + 16, kk]+ __device__ __forceinline__ float warp_sum(float v) {+ for (int d = 16; d > 0; d >>= 1) {+ v += __shfl_down_sync(0xffffffff, v, d);+ }+ return v;+ }- acc00 = acc00 + a0 * b0- acc01 = acc01 + a0 * b1- acc10 = acc10 + a1 * b0- acc11 = acc11 + a1 * b1+ __device__ __forceinline__ float fast_sigmoid(float x) {+ // 使用 __expf,精度在题面容忍范围内通常足够+ float z = __expf(-x);+ return 1.0f / (1.0f + z);+ }- cute.arch.sync_threads()- yield_out([acc00, acc01, acc10, acc11])+ __global__ void ln1_128_f16(+ const float* __restrict__ x,+ const float* __restrict__ w,+ const float* __restrict__ b,+ half* __restrict__ y,+ int64_t rows) {+ int64_t row = (int64_t)blockIdx.x;+ if (row >= rows) return;+ int lane = (int)threadIdx.x; // 0..31- c[bz, i0, j0] = acc00_out- c[bz, i0, j1] = acc01_out- c[bz, i1, j0] = acc10_out- c[bz, i1, j1] = acc11_out+ const float4* x4 = reinterpret_cast<const float4*>(x + row * 128);+ float4 v = x4[lane];+ float s = v.x + v.y + v.z + v.w;+ float ss = v.x * v.x + v.y * v.y + v.z * v.z + v.w * v.w;+ s = warp_sum(s);+ ss = warp_sum(ss);+ // warp_sum 仅保证 lane0 得到全和,需要广播到全 warp+ s = __shfl_sync(0xffffffff, s, 0);+ ss = __shfl_sync(0xffffffff, ss, 0);+ float mean = s * (1.0f / 128.0f);+ float var = ss * (1.0f / 128.0f) - mean * mean;+ float inv = rsqrtf(var + 1e-5f);+ const float4* w4 = reinterpret_cast<const float4*>(w);+ const float4* b4 = reinterpret_cast<const float4*>(b);+ float4 gw = w4[lane];+ float4 gb = b4[lane];- _CONTRACT = _OutgoingContractBatchedKernel()- _CONTRACT_COMPILED = None+ float y0 = (v.x - mean) * inv * gw.x + gb.x;+ float y1 = (v.y - mean) * inv * gw.y + gb.y;+ float y2 = (v.z - mean) * inv * gw.z + gb.z;+ float y3 = (v.w - mean) * inv * gw.w + gb.w;+ half2 h0 = __floats2half2_rn(y0, y1);+ half2 h1 = __floats2half2_rn(y2, y3);- def _get_contract_compiled():- global _CONTRACT_COMPILED- if _CONTRACT_COMPILED is not None:- return _CONTRACT_COMPILED+ half2* y2p = reinterpret_cast<half2*>(y + row * 128 + lane * 4);+ y2p[0] = h0;+ y2p[1] = h1;+ }- 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_COMPILED = cute.compile(- _CONTRACT,- a_ptr,- b_ptr,- c_ptr,- (0, 0),- options="--opt-level 3",- )- return _CONTRACT_COMPILED+ __global__ void ln1_generic_f16(+ const float* __restrict__ x,+ const float* __restrict__ w,+ const float* __restrict__ b,+ half* __restrict__ y,+ int dim,+ int64_t rows) {+ int64_t row = (int64_t)blockIdx.x;+ if (row >= rows) return;+ // 计算均值与二阶矩+ float sum = 0.0f;+ float sq = 0.0f;+ int64_t base = row * (int64_t)dim;+ for (int c = (int)threadIdx.x; c < dim; c += (int)blockDim.x) {+ float v = x[base + c];+ sum += v;+ sq += v * v;+ }- def _contract_outgoing(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:- bs, n, _, hidden = left.shape+ __shared__ float shm_sum[256];+ __shared__ float shm_sq[256];+ int t = (int)threadIdx.x;+ shm_sum[t] = sum;+ shm_sq[t] = sq;+ __syncthreads();- 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。")+ for (int stride = ((int)blockDim.x) / 2; stride > 0; stride >>= 1) {+ if (t < stride) {+ shm_sum[t] += shm_sum[t + stride];+ shm_sq[t] += shm_sq[t + stride];+ }+ __syncthreads();+ }-- if hidden != 128:- raise RuntimeError("仅支持 hidden_dim=128 的特化路径。")+ float mean = shm_sum[0] / (float)dim;+ float var = shm_sq[0] / (float)dim - mean * mean;+ float inv = rsqrtf(var + 1e-5f);- if (n & (_TILE_MN - 1)) != 0:- raise RuntimeError("仅支持 N 为 32 的倍数的特化路径。")+ for (int c = (int)threadIdx.x; c < dim; c += (int)blockDim.x) {+ float v = x[base + c];+ float yv = (v - mean) * inv * w[c] + b[c];+ y[base + c] = __float2half_rn(yv);+ }+ }+ __global__ void pack_proj_f16(+ const half* __restrict__ proj,+ const float* __restrict__ mask,+ half* __restrict__ left,+ half* __restrict__ right,+ half* __restrict__ ogate,+ int bs,+ int n,+ int hidden) {+ int64_t row = (int64_t)blockIdx.x;+ int d = (int)threadIdx.x;+ if (d >= hidden) return;++ int64_t nn = (int64_t)n * (int64_t)n;+ int b0 = (int)(row / nn);+ int64_t rem = row - (int64_t)b0 * nn;+ int i = (int)(rem / n);+ int j = (int)(rem - (int64_t)i * n);++ float m = mask[row];++ int out = hidden * 5;+ const half* p = proj + row * out;++ float l = __half2float(p[d]);+ float r = __half2float(p[hidden + d]);+ float gl = fast_sigmoid(__half2float(p[2 * hidden + d]));+ float gr = fast_sigmoid(__half2float(p[3 * hidden + d]));+ float go = fast_sigmoid(__half2float(p[4 * hidden + d]));++ float l2 = l * gl * m;+ float r2 = r * gr * m;++ int64_t base = (((int64_t)b0 * hidden + d) * n + i) * n + j;+ left[base] = __float2half_rn(l2);+ right[base] = __float2half_rn(r2);+ ogate[base] = __float2half_rn(go);+ }++ template<int MAX_H>+ __global__ void ln2_gate_store_f16(+ const float* __restrict__ out_acc,+ const half* __restrict__ ogate,+ const float* __restrict__ w,+ const float* __restrict__ b,+ half* __restrict__ out_norm,+ int bs,+ int n,+ int hidden) {+ // blockIdx.x 对应 (b,i,j)+ int64_t row = (int64_t)blockIdx.x;+ int tid = (int)threadIdx.x;+ if (tid >= MAX_H) return;++ int64_t nn = (int64_t)n * (int64_t)n;+ int b0 = (int)(row / nn);+ int64_t rem = row - (int64_t)b0 * nn;+ int i = (int)(rem / n);+ int j = (int)(rem - (int64_t)i * n);++ float v = 0.0f;+ float vv = 0.0f;+ if (tid < hidden) {+ int64_t idx = (((int64_t)b0 * hidden + tid) * n + i) * n + j;+ float x = out_acc[idx];+ v = x;+ vv = x * x;+ }++ // 归约:MAX_H 固定,使用共享内存+ __shared__ float shm_sum[MAX_H];+ __shared__ float shm_sq[MAX_H];+ shm_sum[tid] = v;+ shm_sq[tid] = vv;+ __syncthreads();++ for (int stride = MAX_H / 2; stride > 0; stride >>= 1) {+ if (tid < stride) {+ shm_sum[tid] += shm_sum[tid + stride];+ shm_sq[tid] += shm_sq[tid + stride];+ }+ __syncthreads();+ }++ float mean = shm_sum[0] / (float)hidden;+ float var = shm_sq[0] / (float)hidden - mean * mean;+ float inv = rsqrtf(var + 1e-5f);++ if (tid < hidden) {+ int64_t idx_in = (((int64_t)b0 * hidden + tid) * n + i) * n + j;+ float x = out_acc[idx_in];+ float y = (x - mean) * inv * w[tid] + b[tid];+ float g = __half2float(ogate[idx_in]);+ float z = y * g;++ int64_t idx_out = ((row * hidden) + tid);+ out_norm[idx_out] = __float2half_rn(z);+ }+ }++ static void launch_ln1(torch::Tensor x, torch::Tensor w, torch::Tensor b, torch::Tensor y) {+ int dim = (int)x.size(1);+ auto rows = x.size(0);+ if (dim == 128) {+ dim3 block(32, 1, 1);+ dim3 grid((unsigned)rows, 1, 1);+ ln1_128_f16<<<grid, block>>>(+ (const float*)x.data_ptr(),+ (const float*)w.data_ptr(),+ (const float*)b.data_ptr(),+ (half*)y.data_ptr(),+ (int64_t)rows);+ checkCuda(cudaGetLastError(), "ln1_128_f16");+ } else {+ dim3 block(256, 1, 1);+ dim3 grid((unsigned)rows, 1, 1);+ ln1_generic_f16<<<grid, block>>>(+ (const float*)x.data_ptr(),+ (const float*)w.data_ptr(),+ (const float*)b.data_ptr(),+ (half*)y.data_ptr(),+ dim,+ (int64_t)rows);+ checkCuda(cudaGetLastError(), "ln1_generic_f16");+ }+ }++ static void launch_pack(torch::Tensor proj, torch::Tensor mask, torch::Tensor left, torch::Tensor right, torch::Tensor og, int bs, int n, int hidden) {+ int64_t rows = (int64_t)bs * (int64_t)n * (int64_t)n;+ dim3 block((unsigned)hidden, 1, 1);+ dim3 grid((unsigned)rows, 1, 1);+ pack_proj_f16<<<grid, block>>>(+ (const half*)proj.data_ptr(),+ (const float*)mask.data_ptr(),+ (half*)left.data_ptr(),+ (half*)right.data_ptr(),+ (half*)og.data_ptr(),+ bs, n, hidden);+ checkCuda(cudaGetLastError(), "pack_proj_f16");+ }++ static void launch_ln2(torch::Tensor out_acc, torch::Tensor og, torch::Tensor w, torch::Tensor b, torch::Tensor out_norm, int bs, int n, int hidden) {+ int64_t rows = (int64_t)bs * (int64_t)n * (int64_t)n;+ if (hidden <= 32) {+ dim3 block(32, 1, 1);+ dim3 grid((unsigned)rows, 1, 1);+ ln2_gate_store_f16<32><<<grid, block>>>(+ (const float*)out_acc.data_ptr(),+ (const half*)og.data_ptr(),+ (const float*)w.data_ptr(),+ (const float*)b.data_ptr(),+ (half*)out_norm.data_ptr(),+ bs, n, hidden);+ checkCuda(cudaGetLastError(), "ln2_gate_store_f16_32");+ } else if (hidden <= 64) {+ dim3 block(64, 1, 1);+ dim3 grid((unsigned)rows, 1, 1);+ ln2_gate_store_f16<64><<<grid, block>>>(+ (const float*)out_acc.data_ptr(),+ (const half*)og.data_ptr(),+ (const float*)w.data_ptr(),+ (const float*)b.data_ptr(),+ (half*)out_norm.data_ptr(),+ bs, n, hidden);+ checkCuda(cudaGetLastError(), "ln2_gate_store_f16_64");+ } else if (hidden <= 128) {+ dim3 block(128, 1, 1);+ dim3 grid((unsigned)rows, 1, 1);+ ln2_gate_store_f16<128><<<grid, block>>>(+ (const float*)out_acc.data_ptr(),+ (const half*)og.data_ptr(),+ (const float*)w.data_ptr(),+ (const float*)b.data_ptr(),+ (half*)out_norm.data_ptr(),+ bs, n, hidden);+ checkCuda(cudaGetLastError(), "ln2_gate_store_f16_128");+ } else {+ throw std::runtime_error("hidden_dim too large");+ }+ }++ static void gemm_x_wt_f16_f16(+ cublasHandle_t h,+ const half* x_row, // row-major [M,K]+ const half* w_row, // row-major [N,K]+ half* y_row, // row-major [M,N]+ int64_t M,+ int64_t N,+ int64_t K) {+ // 使用列主序 trick:输出按列主序 (N x M) 写入,即等价于 row-major (M x N)+ float alpha = 1.0f;+ float beta = 0.0f;+ checkCublas(+ cublasGemmEx(+ h,+ CUBLAS_OP_T, CUBLAS_OP_N,+ (int)N, (int)M, (int)K,+ &alpha,+ w_row, CUDA_R_16F, (int)K,+ x_row, CUDA_R_16F, (int)K,+ &beta,+ y_row, CUDA_R_16F, (int)N,+ CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),+ "cublasGemmEx");+ }++ static void gemm_x_wt_f16_f32(+ cublasHandle_t h,+ const half* x_row, // row-major [M,K]+ const half* w_row, // row-major [N,K]+ float* y_row, // row-major [M,N]+ int64_t M,+ int64_t N,+ int64_t K) {+ float alpha = 1.0f;+ float beta = 0.0f;+ checkCublas(+ cublasGemmEx(+ h,+ CUBLAS_OP_T, CUBLAS_OP_N,+ (int)N, (int)M, (int)K,+ &alpha,+ w_row, CUDA_R_16F, (int)K,+ x_row, CUDA_R_16F, (int)K,+ &beta,+ y_row, CUDA_R_32F, (int)N,+ CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),+ "cublasGemmEx");+ }++ static void gemm_contract_batched_f16_f32(+ cublasHandle_t h,+ const half* left_row, // row-major [B,M,K]+ const half* right_row, // row-major [B,N,K] (这里 N=M=K=n)+ float* out_row, // row-major [B,M,N]+ int batch,+ int n) {+ float alpha = 1.0f;+ float beta = 0.0f;+ long long strideA = (long long)n * (long long)n;+ long long strideB = (long long)n * (long long)n;+ long long strideC = (long long)n * (long long)n;++ // 计算 C = L * R^T+ // 采用列主序 trick:C^T = R * L^T+ checkCublas(+ cublasGemmStridedBatchedEx(+ h,+ CUBLAS_OP_T, CUBLAS_OP_N,+ n, n, n,+ &alpha,+ right_row, CUDA_R_16F, n, strideB,+ left_row, CUDA_R_16F, n, strideA,+ &beta,+ out_row, CUDA_R_32F, n, strideC,+ batch,+ CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),+ "cublasGemmStridedBatchedEx");+ }++ } // namespace++ torch::Tensor trimul_fwd(+ torch::Tensor x,+ torch::Tensor mask,+ torch::Tensor ln1_w,+ torch::Tensor ln1_b,+ torch::Tensor w_cat,+ torch::Tensor ln2_w,+ torch::Tensor ln2_b,+ torch::Tensor w_out,+ int64_t dim,+ int64_t hidden) {+ if (!x.is_cuda() || !mask.is_cuda()) {+ throw std::runtime_error("cuda only");+ }+ if (x.scalar_type() != torch::kFloat32) {+ throw std::runtime_error("x must be float32");+ }+ if (mask.scalar_type() != torch::kFloat32) {+ throw std::runtime_error("mask must be float32");+ }+ if (dim != x.size(3)) {+ throw std::runtime_error("dim mismatch");+ }+ if (w_cat.scalar_type() != torch::kFloat16 || w_out.scalar_type() != torch::kFloat16) {+ throw std::runtime_error("weights must be float16");+ }++ int bs = (int)x.size(0);+ int n = (int)x.size(1);+ int64_t M = (int64_t)bs * (int64_t)n * (int64_t)n;++ auto x2d = x.view({M, dim});+ auto xhat = torch::empty({M, dim}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));+ launch_ln1(x2d, ln1_w, ln1_b, xhat);++ int64_t out_ch = hidden * 5;+ auto proj = torch::empty({M, out_ch}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));++ auto* holder = get_cublas();+ cublasHandle_t h = holder->handle;++ // gemm1: proj = xhat @ w_cat^T+ gemm_x_wt_f16_f16(h, (const half*)xhat.data_ptr(), (const half*)w_cat.data_ptr(), (half*)proj.data_ptr(), M, out_ch, dim);++ auto left = torch::empty({bs, hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));+ auto right = torch::empty({bs, hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));+ auto og = torch::empty({bs, hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));+ launch_pack(proj, mask.view({M}), left, right, og, bs, n, (int)hidden);++ auto left3 = left.view({bs * (int)hidden, n, n});+ auto right3 = right.view({bs * (int)hidden, n, n});+ auto out_acc = torch::empty({bs * (int)hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat32));++ gemm_contract_batched_f16_f32(+ h,+ (const half*)left3.data_ptr(),+ (const half*)right3.data_ptr(),+ (float*)out_acc.data_ptr(),+ bs * (int)hidden, n);++ auto out_norm = torch::empty({M, hidden}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));+ launch_ln2(out_acc, og.view({bs * (int)hidden, n, n}), ln2_w, ln2_b, out_norm, bs, n, (int)hidden);++ // gemm2: y = out_norm @ w_out^T+ auto y = torch::empty({M, dim}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat32));+ gemm_x_wt_f16_f32(h, (const half*)out_norm.data_ptr(), (const half*)w_out.data_ptr(), (float*)y.data_ptr(), M, dim, hidden);++ return y.view({bs, n, n, dim});+ }+ """++ name = "trimul_ext_mod"++ extra_cuda_cflags = [+ "-O3",+ "--use_fast_math",+ ]+ extra_cflags = [+ "-O3",+ ]++ extra_ldflags = [+ "-lcublas",+ ]++ _EXT = load_inline(+ name=name,+ cpp_sources=cpp_src,+ cuda_sources=cuda_src,+ functions=None,+ extra_cflags=extra_cflags,+ extra_cuda_cflags=extra_cuda_cflags,+ extra_ldflags=extra_ldflags,+ with_cuda=True,+ verbose=False,+ )+ return _EXT+++ class _WeightCache:+ __slots__ = ("key", "w_cat", "w_out")++ def __init__(self) -> None:+ self.key = None+ self.w_cat = None+ self.w_out = None+++ _W_CACHE = _WeightCache()+++ def _prepare_weights(weights: Dict[str, torch.Tensor], dim: int, hidden: int):- left_t = left.permute(0, 3, 1, 2).contiguous()- right_t = right.permute(0, 3, 1, 2).contiguous()- bh = bs * hidden- a = left_t.view(bh, n, n)- b = right_t.view(bh, n, n)+ k = (+ 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()),+ )+ if _W_CACHE.key == k and _W_CACHE.w_cat is not None and _W_CACHE.w_out is not None:+ return _W_CACHE.w_cat, _W_CACHE.w_out- c = torch.empty((bh, n, n), device=left.device, dtype=torch.float32)- 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_left = weights["left_proj.weight"]+ w_right = weights["right_proj.weight"]+ w_lg = weights["left_gate.weight"]+ w_rg = weights["right_gate.weight"]+ w_og = weights["out_gate.weight"]+ w_out = weights["to_out.weight"]- out = c.view(bs, hidden, n, n).permute(0, 2, 3, 1).contiguous()- return out+ if w_left.shape != (hidden, dim):+ raise RuntimeError("left_proj.weight shape mismatch")+ if w_right.shape != (hidden, dim):+ raise RuntimeError("right_proj.weight shape mismatch")+ if w_lg.shape != (hidden, dim):+ raise RuntimeError("left_gate.weight shape mismatch")+ if w_rg.shape != (hidden, dim):+ raise RuntimeError("right_gate.weight shape mismatch")+ if w_og.shape != (hidden, dim):+ raise RuntimeError("out_gate.weight shape mismatch")+ if w_out.shape != (dim, hidden):+ raise RuntimeError("to_out.weight shape mismatch")++ w_cat = torch.cat([w_left, w_right, w_lg, w_rg, w_og], dim=0).contiguous().to(dtype=torch.float16)+ w_out_h = w_out.contiguous().to(dtype=torch.float16)+ _W_CACHE.key = k+ _W_CACHE.w_cat = w_cat+ _W_CACHE.w_out = w_out_h+ return w_cat, w_out_h++@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-dim = int(config["dim"])- hidden_dim = int(config["hidden_dim"])+ hidden = int(config["hidden_dim"])+ if not x.is_cuda:+ raise RuntimeError("x must be CUDA tensor")+ if not mask.is_cuda:+ raise RuntimeError("mask must be CUDA tensor")+if x.dtype != torch.float32:x = x.to(dtype=torch.float32)- torch.backends.cuda.matmul.allow_tf32 = True- torch.backends.cudnn.allow_tf32 = True+ if mask.dtype != torch.float32:+ mask = mask.to(dtype=torch.float32)- x = F.layer_norm(x, (dim,), weights["norm.weight"], weights["norm.bias"], 1e-5)+ x = x.contiguous()+ mask = mask.contiguous()- left = F.linear(x, weights["left_proj.weight"], None)- right = F.linear(x, weights["right_proj.weight"], None)+ if x.ndim != 4:+ raise RuntimeError("x must be 4D")+ if mask.ndim != 3:+ raise RuntimeError("mask must be 3D")+ if x.shape[:3] != mask.shape:+ raise RuntimeError("x/mask shape mismatch")+ if x.shape[3] != dim:+ raise RuntimeError("dim mismatch")- mask_f = mask.unsqueeze(-1)- if mask_f.dtype != left.dtype:- mask_f = mask_f.to(dtype=left.dtype)- left = left * mask_f- right = right * mask_f++ for k in (+ "norm.weight",+ "norm.bias",+ "left_proj.weight",+ "right_proj.weight",+ "left_gate.weight",+ "right_gate.weight",+ "out_gate.weight",+ "to_out_norm.weight",+ "to_out_norm.bias",+ "to_out.weight",+ ):+ if not weights[k].is_cuda:+ raise RuntimeError(f"weight {k} must be CUDA tensor")+ if weights[k].dtype != torch.float32:+ raise RuntimeError(f"weight {k} must be float32")- left_gate = torch.sigmoid(F.linear(x, weights["left_gate.weight"], None))- right_gate = torch.sigmoid(F.linear(x, weights["right_gate.weight"], None))- out_gate = torch.sigmoid(F.linear(x, weights["out_gate.weight"], None))++ ln1_w = weights["norm.weight"].contiguous()+ ln1_b = weights["norm.bias"].contiguous()+ if ln1_w.shape != (dim,) or ln1_b.shape != (dim,):+ raise RuntimeError("norm params shape mismatch")- left = left * left_gate- right = right * right_gate+ ln2_w = weights["to_out_norm.weight"].contiguous()+ ln2_b = weights["to_out_norm.bias"].contiguous()+ if ln2_w.shape != (hidden,) or ln2_b.shape != (hidden,):+ raise RuntimeError("to_out_norm params shape mismatch")- 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 = out * out_gate- out = F.linear(out, weights["to_out.weight"], None)- return out++ w_cat, w_out = _prepare_weights(weights, dim, hidden)++ ext = _get_ext()+ return ext.fwd(x, mask, ln1_w, ln1_b, w_cat, ln2_w, ln2_b, w_out, dim, hidden)+__all__ = ["custom_kernel"]
scrolls · 874 diff lines total
Best evidence level for this revision: reported
JSON