submission 418614
shiyegao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 869 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-418614?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:bcea436d35510a9ade8ddd6fcaadb762c107415122930b1bda8a14220963ad85
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 warp_s[3];vector-width = float4
const float4* x4 = reinterpret_cast<const float4*>(x + row * 128);Kernel source
submission.py869 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_h,
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 <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");
// 强制启用 tensor op math(half 输入的 GEMM 明确走张量核)
checkCublas(cublasSetMathMode(handle, CUBLAS_TENSOR_OP_MATH), "cublasSetMathMode");
}
~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) {
float z = __expf(-x);
return 1.0f / (1.0f + z);
}
// ---------------- LN1 ----------------
__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);
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_384_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 tid = (int)threadIdx.x; // 0..95
int lane = tid & 31; // 0..31
int warp_id = tid >> 5; // 0..2
const float4* x4 = reinterpret_cast<const float4*>(x + row * 384);
float4 v = x4[tid];
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);
__shared__ float warp_s[3];
__shared__ float warp_ss[3];
__shared__ float tot_s;
__shared__ float tot_ss;
if (lane == 0) {
warp_s[warp_id] = s;
warp_ss[warp_id] = ss;
}
__syncthreads();
float sum = 0.0f;
float sq = 0.0f;
if (warp_id == 0) {
if (lane < 3) {
sum = warp_s[lane];
sq = warp_ss[lane];
}
sum = warp_sum(sum);
sq = warp_sum(sq);
}
if (warp_id == 0 && lane == 0) {
tot_s = sum;
tot_ss = sq;
}
__syncthreads();
sum = tot_s;
sq = tot_ss;
float mean = sum * (1.0f / 384.0f);
float var = sq * (1.0f / 384.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[tid];
float4 gb = b4[tid];
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 * 384 + tid * 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);
}
}
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 if (dim == 384) {
dim3 block(96, 1, 1);
dim3 grid((unsigned)rows, 1, 1);
ln1_384_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_384_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");
}
}
// --------------- pack (projT -> left/right) ---------------
// projT 形状为 [5H, M](row-major,最后一维 M 连续),避免额外转置。
// 关键点:不再物化 out_gate(ogate)到单独张量,减少一次全量读写;
// LN2 阶段直接从 projT 的第 5 组读取并做 sigmoid。
__global__ void pack4_dmaj_f16(
const half* __restrict__ projT, // [5H, M]
const half* __restrict__ mask_h, // [bs, nn]
half* __restrict__ left, // [bs*H, nn]
half* __restrict__ right, // [bs*H, nn]
int nn,
int M,
int hidden) {
int bd = (int)blockIdx.y; // 0..bs*hidden-1
int d = bd - (bd / hidden) * hidden;
int b = bd / hidden;
int p0 = (int)blockIdx.x * ((int)blockDim.x * 2) + (int)threadIdx.x * 2;
if (p0 >= nn) return;
const half* mask_row = mask_h + b * nn;
int idx0 = b * nn + p0;
const half* row_l = projT + (d * M);
const half* row_r = projT + ((hidden + d) * M);
const half* row_gl = projT + ((2 * hidden + d) * M);
const half* row_gr = projT + ((3 * hidden + d) * M);
half* out_l = left + bd * nn;
half* out_r = right + bd * nn;
if (p0 + 1 < nn) {
half2 m2 = *reinterpret_cast<const half2*>(mask_row + p0);
half2 l2 = *reinterpret_cast<const half2*>(row_l + idx0);
half2 r2 = *reinterpret_cast<const half2*>(row_r + idx0);
half2 gl2 = *reinterpret_cast<const half2*>(row_gl + idx0);
half2 gr2 = *reinterpret_cast<const half2*>(row_gr + idx0);
float2 mf = __half22float2(m2);
float2 lf = __half22float2(l2);
float2 rf = __half22float2(r2);
float2 glf = __half22float2(gl2);
float2 grf = __half22float2(gr2);
glf.x = fast_sigmoid(glf.x);
glf.y = fast_sigmoid(glf.y);
grf.x = fast_sigmoid(grf.x);
grf.y = fast_sigmoid(grf.y);
float lo0 = lf.x * glf.x * mf.x;
float lo1 = lf.y * glf.y * mf.y;
float ro0 = rf.x * grf.x * mf.x;
float ro1 = rf.y * grf.y * mf.y;
*reinterpret_cast<half2*>(out_l + p0) = __floats2half2_rn(lo0, lo1);
*reinterpret_cast<half2*>(out_r + p0) = __floats2half2_rn(ro0, ro1);
return;
}
float m0 = __half2float(mask_row[p0]);
float l0 = __half2float(row_l[idx0]);
float r0 = __half2float(row_r[idx0]);
float gl0 = fast_sigmoid(__half2float(row_gl[idx0]));
float gr0 = fast_sigmoid(__half2float(row_gr[idx0]));
out_l[p0] = __float2half_rn(l0 * gl0 * m0);
out_r[p0] = __float2half_rn(r0 * gr0 * m0);
}
static void launch_pack4(
torch::Tensor projT,
torch::Tensor mask_h,
torch::Tensor left,
torch::Tensor right,
int bs,
int n,
int hidden) {
int nn = n * n;
int M = bs * nn;
dim3 block(256, 1, 1);
dim3 grid((unsigned)((nn + (int)block.x * 2 - 1) / ((int)block.x * 2)), (unsigned)(bs * hidden), 1);
pack4_dmaj_f16<<<grid, block>>>(
(const half*)projT.data_ptr(),
(const half*)mask_h.data_ptr(),
(half*)left.data_ptr(),
(half*)right.data_ptr(),
nn,
M,
hidden);
checkCuda(cudaGetLastError(), "pack4_dmaj_f16");
}
// --------------- LN2 + gate ---------------
// out_acc: [bs*H, nn];go_T: [H, M](d-major,最后一维 M 连续,存的是 pre-sigmoid)
// 输出 out_norm_T: [H, M](row-major,最后一维 M 连续)。
// 本版本不读取/存储 ogate 中间张量,直接从 go_T 做 sigmoid。
template<int WARPS>
__global__ void ln2_gate_tile32_f16(
const half* __restrict__ out_acc, // [bs*H, nn] (half)
const half* __restrict__ go_T, // [H, M] (half, pre-sigmoid)
const float* __restrict__ w, // [H]
const float* __restrict__ b, // [H]
half* __restrict__ out_norm_T, // [H, M]
int nn,
int hidden) {
int bb = (int)blockIdx.y;
int p0 = (int)blockIdx.x * 32;
int tid = (int)threadIdx.x;
int lane = tid & 31;
int wid = tid >> 5;
int p = p0 + lane;
int M = nn * (int)gridDim.y;
int out_p = bb * nn + p;
extern __shared__ unsigned char smem_u8[];
half* sh_x = (half*)smem_u8; // hidden*32
float* sh_sum = (float*)(smem_u8 + (size_t)((int64_t)hidden * 32) * sizeof(half));
float* sh_sq = sh_sum + WARPS * 32;
float* sh_mean = sh_sq + WARPS * 32;
float* sh_inv = sh_mean + 32;
float sum = 0.0f;
float sq = 0.0f;
for (int d = wid; d < hidden; d += WARPS) {
half hx = __float2half_rn(0.0f);
if (p < nn) {
int idx = (bb * hidden + d) * nn + p;
hx = out_acc[idx];
}
sh_x[(int64_t)d * 32 + lane] = hx;
float x = __half2float(hx);
sum += x;
sq += x * x;
}
sh_sum[wid * 32 + lane] = sum;
sh_sq[wid * 32 + lane] = sq;
__syncthreads();
if (wid == 0) {
float tot = 0.0f;
float tot_sq = 0.0f;
#pragma unroll
for (int w_id = 0; w_id < WARPS; ++w_id) {
tot += sh_sum[w_id * 32 + lane];
tot_sq += sh_sq[w_id * 32 + lane];
}
float inv_n = 1.0f / (float)hidden;
float mean = tot * inv_n;
float var = tot_sq * inv_n - mean * mean;
sh_mean[lane] = mean;
sh_inv[lane] = rsqrtf(var + 1e-5f);
}
__syncthreads();
if (p >= nn) return;
float mean = sh_mean[lane];
float inv = sh_inv[lane];
for (int d = wid; d < hidden; d += WARPS) {
float x = __half2float(sh_x[(int64_t)d * 32 + lane]);
float y = (x - mean) * inv * w[d] + b[d];
float g = fast_sigmoid(__half2float(go_T[d * M + out_p]));
out_norm_T[d * M + out_p] = __float2half_rn(y * g);
}
}
static void launch_ln2(
torch::Tensor out_acc,
torch::Tensor go_T,
torch::Tensor w,
torch::Tensor b,
torch::Tensor out_norm_T,
int bs,
int n,
int hidden) {
int nn = n * n;
dim3 grid((unsigned)((nn + 31) / 32), (unsigned)bs, 1);
if (hidden <= 32) {
constexpr int WARPS = 1;
dim3 block(WARPS * 32, 1, 1);
size_t shmem = (size_t)((int64_t)hidden * 32) * sizeof(half)
+ (size_t)((int64_t)WARPS * 32 * 2 + 64) * sizeof(float);
ln2_gate_tile32_f16<WARPS><<<grid, block, shmem>>>(
(const half*)out_acc.data_ptr(),
(const half*)go_T.data_ptr(),
(const float*)w.data_ptr(),
(const float*)b.data_ptr(),
(half*)out_norm_T.data_ptr(),
nn,
hidden);
checkCuda(cudaGetLastError(), "ln2_gate_tile32_f16_1w");
} else if (hidden <= 64) {
constexpr int WARPS = 2;
dim3 block(WARPS * 32, 1, 1);
size_t shmem = (size_t)((int64_t)hidden * 32) * sizeof(half)
+ (size_t)((int64_t)WARPS * 32 * 2 + 64) * sizeof(float);
ln2_gate_tile32_f16<WARPS><<<grid, block, shmem>>>(
(const half*)out_acc.data_ptr(),
(const half*)go_T.data_ptr(),
(const float*)w.data_ptr(),
(const float*)b.data_ptr(),
(half*)out_norm_T.data_ptr(),
nn,
hidden);
checkCuda(cudaGetLastError(), "ln2_gate_tile32_f16_2w");
} else if (hidden <= 128) {
constexpr int WARPS = 4;
dim3 block(WARPS * 32, 1, 1);
size_t shmem = (size_t)((int64_t)hidden * 32) * sizeof(half)
+ (size_t)((int64_t)WARPS * 32 * 2 + 64) * sizeof(float);
ln2_gate_tile32_f16<WARPS><<<grid, block, shmem>>>(
(const half*)out_acc.data_ptr(),
(const half*)go_T.data_ptr(),
(const float*)w.data_ptr(),
(const float*)b.data_ptr(),
(half*)out_norm_T.data_ptr(),
nn,
hidden);
checkCuda(cudaGetLastError(), "ln2_gate_tile32_f16_4w");
} else if (hidden <= 256) {
constexpr int WARPS = 8;
dim3 block(WARPS * 32, 1, 1);
size_t shmem = (size_t)((int64_t)hidden * 32) * sizeof(half)
+ (size_t)((int64_t)WARPS * 32 * 2 + 64) * sizeof(float);
ln2_gate_tile32_f16<WARPS><<<grid, block, shmem>>>(
(const half*)out_acc.data_ptr(),
(const half*)go_T.data_ptr(),
(const float*)w.data_ptr(),
(const float*)b.data_ptr(),
(half*)out_norm_T.data_ptr(),
nn,
hidden);
checkCuda(cudaGetLastError(), "ln2_gate_tile32_f16_8w");
} else {
throw std::runtime_error("hidden_dim too large");
}
}
// ---------------- GEMM helpers ----------------
// 约定:所有矩阵都来自 row-major Tensor,但用 cuBLAS 的 column-major 语义解释,
// 通过精确设置 m/n/k 与 lda/ldb/ldc 得到想要的布局,避免额外转置核。
static void gemm1_x_wt_to_dmaj_f16(
cublasHandle_t h,
const half* x_rm, // [M, K] row-major
const half* w_rm, // [N, K] row-major
half* c_dmaj_rm, // [N, M] row-major (等价于 column-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)M, (int)N, (int)K,
&alpha,
x_rm, CUDA_R_16F, (int)K,
w_rm, CUDA_R_16F, (int)K,
&beta,
c_dmaj_rm, CUDA_R_16F, (int)M,
CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
"cublasGemmEx_gemm1");
}
static void gemm2_dmaj_to_y_t_f16_f16(
cublasHandle_t h,
const half* a_dmaj_rm, // [K, M] row-major (等价于 column-major [M, K])
const half* w_rm, // [N, K] row-major (等价于 column-major [K, N])
half* y_t_rm, // [N, M] row-major (等价于 column-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_N, CUBLAS_OP_N,
(int)M, (int)N, (int)K,
&alpha,
a_dmaj_rm, CUDA_R_16F, (int)M,
w_rm, CUDA_R_16F, (int)K,
&beta,
y_t_rm, CUDA_R_16F, (int)M,
CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
"cublasGemmEx_gemm2");
}
static void gemm_contract_batched_f16_f16(
cublasHandle_t h,
const half* left_row, // [B,M,K] row-major [batch,n,n]
const half* right_row, // [B,N,K] row-major [batch,n,n]
half* out_row, // [B,M,N] row-major [batch,n,n] (half)
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;
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_16F, n, strideC,
batch,
CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
"cublasGemmStridedBatchedEx_contract");
}
} // namespace
torch::Tensor trimul_fwd(
torch::Tensor x,
torch::Tensor mask_h,
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_h.is_cuda()) {
throw std::runtime_error("cuda only");
}
if (x.scalar_type() != torch::kFloat32) {
throw std::runtime_error("x must be float32");
}
if (mask_h.scalar_type() != torch::kFloat16) {
throw std::runtime_error("mask must be float16");
}
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 nn = (int64_t)n * (int64_t)n;
int64_t M = (int64_t)bs * nn;
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;
// projT: [5H, M](d-major,最后一维连续)
auto projT = torch::empty({out_ch, M}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
auto* holder = get_cublas();
cublasHandle_t h = holder->handle;
gemm1_x_wt_to_dmaj_f16(
h,
(const half*)xhat.data_ptr(),
(const half*)w_cat.data_ptr(),
(half*)projT.data_ptr(),
M,
out_ch,
dim);
auto left = torch::empty({bs * (int)hidden, nn}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
auto right = torch::empty({bs * (int)hidden, nn}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
// pack:projT[5H,M] + mask[bs,nn] -> left/right[bs*H,nn]
launch_pack4(projT, mask_h.view({bs, nn}), left, right, bs, n, (int)hidden);
auto left3 = left.view({bs * (int)hidden, n, n});
auto right3 = right.view({bs * (int)hidden, n, n});
// out_acc: half(仍由 GEMM 做 FP32 累加,只降低写回精度)
auto out_acc = torch::empty({bs * (int)hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
gemm_contract_batched_f16_f16(
h,
(const half*)left3.data_ptr(),
(const half*)right3.data_ptr(),
(half*)out_acc.data_ptr(),
bs * (int)hidden, n);
// go_T: [H, M],从 projT 的第 5 组切片得到(pre-sigmoid)
auto go_T = projT.narrow(0, (int64_t)4 * hidden, hidden);
// LN2 + gate:输出 out_norm_T[H,M] half
auto out_norm_T = torch::empty({hidden, M}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
launch_ln2(out_acc, go_T, ln2_w, ln2_b, out_norm_T, bs, n, (int)hidden);
// gemm2:y_T[dim,M] half(降低最终写回带宽;仍 FP32 累加)
auto y_T = torch::empty({dim, M}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
gemm2_dmaj_to_y_t_f16_f16(
h,
(const half*)out_norm_T.data_ptr(),
(const half*)w_out.data_ptr(),
(half*)y_T.data_ptr(),
M,
dim,
hidden);
return y_T.view({dim, bs, n, n}).permute({1, 2, 3, 0});
}
"""
name = "trimul_ext_mod9"
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.float16:
mask = mask.to(dtype=torch.float16)
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 · 869 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 418609.
⋯ diff truncated: revisions differ almost entirely
Best evidence level for this revision: reported
JSON