submission 417882
shiyegao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 950 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-417882?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:e16a58c1fabe47092486feb0b35670be0448a98818e57ef8760b0af182f0aaaf
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 sh_sum[3];vector-width = float4
const float4* x4 = reinterpret_cast<const float4*>(x + row * 128);Kernel source
submission.py950 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 <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");
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) {
v += __shfl_down_sync(0xffffffff, v, 16);
v += __shfl_down_sync(0xffffffff, v, 8);
v += __shfl_down_sync(0xffffffff, v, 4);
v += __shfl_down_sync(0xffffffff, v, 2);
v += __shfl_down_sync(0xffffffff, v, 1);
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;
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;
int warp = 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 sh_sum[3];
__shared__ float sh_sq[3];
if (lane == 0) {
sh_sum[warp] = s;
sh_sq[warp] = ss;
}
__syncthreads();
__shared__ float mean_sh;
__shared__ float inv_sh;
if (tid == 0) {
float sum = sh_sum[0] + sh_sum[1] + sh_sum[2];
float sq = sh_sq[0] + sh_sq[1] + sh_sq[2];
float mean = sum * (1.0f / 384.0f);
float var = sq * (1.0f / 384.0f) - mean * mean;
mean_sh = mean;
inv_sh = rsqrtf(var + 1e-5f);
}
__syncthreads();
float mean = mean_sh;
float inv = inv_sh;
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;
int64_t off = row * 384 + (int64_t)tid * 4;
half2 h0 = __floats2half2_rn(y0, y1);
half2 h1 = __floats2half2_rn(y2, y3);
half2* y2p = reinterpret_cast<half2*>(y + off);
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 (proj -> left/right/ogate) ---------------
template<int D_TILE, int TILE_P>
__global__ void pack_proj_tiled2_f16(
const half* __restrict__ proj,
const float* __restrict__ mask,
half* __restrict__ left,
half* __restrict__ right,
half* __restrict__ ogate,
int64_t nn,
int hidden) {
int b = (int)blockIdx.y;
int64_t p0 = (int64_t)blockIdx.x * (int64_t)TILE_P;
int tid = (int)threadIdx.x;
int p_l = tid / D_TILE;
int d0 = tid - p_l * D_TILE;
int64_t p = p0 + (int64_t)p_l;
__shared__ float m_sh[TILE_P];
if (d0 == 0) {
float mv = 0.0f;
if (p < nn) {
mv = mask[(int64_t)b * nn + p];
}
m_sh[p_l] = mv;
}
__syncthreads();
__shared__ half sh_l[128 * TILE_P];
__shared__ half sh_r[128 * TILE_P];
__shared__ half sh_g[128 * TILE_P];
half out_l0 = __float2half_rn(0.0f);
half out_r0 = __float2half_rn(0.0f);
half out_g0 = __float2half_rn(0.0f);
half out_l1 = __float2half_rn(0.0f);
half out_r1 = __float2half_rn(0.0f);
half out_g1 = __float2half_rn(0.0f);
if (p < nn) {
int64_t row = (int64_t)b * nn + p;
int out_ch = hidden * 5;
const half* base = proj + row * (int64_t)out_ch;
float m = m_sh[p_l];
int d = d0;
if (d < hidden) {
float l = __half2float(base[d]);
float r = __half2float(base[hidden + d]);
float gl = fast_sigmoid(__half2float(base[2 * hidden + d]));
float gr = fast_sigmoid(__half2float(base[3 * hidden + d]));
float go = fast_sigmoid(__half2float(base[4 * hidden + d]));
out_l0 = __float2half_rn(l * gl * m);
out_r0 = __float2half_rn(r * gr * m);
out_g0 = __float2half_rn(go);
}
int d1 = d0 + 64;
if (d1 < hidden) {
float l = __half2float(base[d1]);
float r = __half2float(base[hidden + d1]);
float gl = fast_sigmoid(__half2float(base[2 * hidden + d1]));
float gr = fast_sigmoid(__half2float(base[3 * hidden + d1]));
float go = fast_sigmoid(__half2float(base[4 * hidden + d1]));
out_l1 = __float2half_rn(l * gl * m);
out_r1 = __float2half_rn(r * gr * m);
out_g1 = __float2half_rn(go);
}
}
sh_l[d0 * TILE_P + p_l] = out_l0;
sh_r[d0 * TILE_P + p_l] = out_r0;
sh_g[d0 * TILE_P + p_l] = out_g0;
sh_l[(d0 + 64) * TILE_P + p_l] = out_l1;
sh_r[(d0 + 64) * TILE_P + p_l] = out_r1;
sh_g[(d0 + 64) * TILE_P + p_l] = out_g1;
__syncthreads();
int d2 = tid / TILE_P; // 0..63
int p2 = tid - d2 * TILE_P;
int64_t p_out = p0 + (int64_t)p2;
if (p_out < nn) {
int d = d2;
if (d < hidden) {
int64_t out_idx = ((int64_t)b * (int64_t)hidden + (int64_t)d) * nn + p_out;
left[out_idx] = sh_l[d * TILE_P + p2];
right[out_idx] = sh_r[d * TILE_P + p2];
ogate[out_idx] = sh_g[d * TILE_P + p2];
}
int d3 = d2 + 64;
if (d3 < hidden) {
int64_t out_idx = ((int64_t)b * (int64_t)hidden + (int64_t)d3) * nn + p_out;
left[out_idx] = sh_l[d3 * TILE_P + p2];
right[out_idx] = sh_r[d3 * TILE_P + p2];
ogate[out_idx] = sh_g[d3 * TILE_P + p2];
}
}
}
template<int MAX_H, int TILE_P>
__global__ void pack_proj_tiled_f16(
const half* __restrict__ proj,
const float* __restrict__ mask,
half* __restrict__ left,
half* __restrict__ right,
half* __restrict__ ogate,
int64_t nn,
int hidden) {
int b = (int)blockIdx.y;
int64_t p0 = (int64_t)blockIdx.x * (int64_t)TILE_P;
int tid = (int)threadIdx.x;
int p_l = tid / MAX_H;
int d = tid - p_l * MAX_H;
int64_t p = p0 + (int64_t)p_l;
__shared__ float m_sh[TILE_P];
if (d == 0) {
float mv = 0.0f;
if (p < nn) {
mv = mask[(int64_t)b * nn + p];
}
m_sh[p_l] = mv;
}
__syncthreads();
__shared__ half sh_l[MAX_H * TILE_P];
__shared__ half sh_r[MAX_H * TILE_P];
__shared__ half sh_g[MAX_H * TILE_P];
half out_l = __float2half_rn(0.0f);
half out_r = __float2half_rn(0.0f);
half out_g = __float2half_rn(0.0f);
if (d < hidden && p < nn) {
int64_t row = (int64_t)b * nn + p;
int out_ch = hidden * 5;
const half* base = proj + row * (int64_t)out_ch;
float l = __half2float(base[d]);
float r = __half2float(base[hidden + d]);
float gl = fast_sigmoid(__half2float(base[2 * hidden + d]));
float gr = fast_sigmoid(__half2float(base[3 * hidden + d]));
float go = fast_sigmoid(__half2float(base[4 * hidden + d]));
float m = m_sh[p_l];
out_l = __float2half_rn(l * gl * m);
out_r = __float2half_rn(r * gr * m);
out_g = __float2half_rn(go);
}
sh_l[d * TILE_P + p_l] = out_l;
sh_r[d * TILE_P + p_l] = out_r;
sh_g[d * TILE_P + p_l] = out_g;
__syncthreads();
int d2 = tid / TILE_P;
int p2 = tid - d2 * TILE_P;
int64_t p_out = p0 + (int64_t)p2;
if (d2 < hidden && p_out < nn) {
int64_t out_idx = ((int64_t)b * (int64_t)hidden + (int64_t)d2) * nn + p_out;
left[out_idx] = sh_l[d2 * TILE_P + p2];
right[out_idx] = sh_r[d2 * TILE_P + p2];
ogate[out_idx] = sh_g[d2 * TILE_P + p2];
}
}
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 nn = (int64_t)n * (int64_t)n;
if (hidden <= 32) {
constexpr int MAX_H = 32;
constexpr int TILE_P = 16;
dim3 block(MAX_H * TILE_P, 1, 1);
dim3 grid((unsigned)((nn + TILE_P - 1) / TILE_P), (unsigned)bs, 1);
pack_proj_tiled_f16<MAX_H, TILE_P><<<grid, block>>>(
(const half*)proj.data_ptr(),
(const float*)mask.data_ptr(),
(half*)left.data_ptr(),
(half*)right.data_ptr(),
(half*)og.data_ptr(),
nn,
hidden);
checkCuda(cudaGetLastError(), "pack_proj_tiled_f16_32");
} else if (hidden <= 64) {
constexpr int MAX_H = 64;
constexpr int TILE_P = 8;
dim3 block(MAX_H * TILE_P, 1, 1);
dim3 grid((unsigned)((nn + TILE_P - 1) / TILE_P), (unsigned)bs, 1);
pack_proj_tiled_f16<MAX_H, TILE_P><<<grid, block>>>(
(const half*)proj.data_ptr(),
(const float*)mask.data_ptr(),
(half*)left.data_ptr(),
(half*)right.data_ptr(),
(half*)og.data_ptr(),
nn,
hidden);
checkCuda(cudaGetLastError(), "pack_proj_tiled_f16_64");
} else if (hidden <= 128) {
constexpr int D_TILE = 64;
constexpr int TILE_P = 8;
dim3 block(D_TILE * TILE_P, 1, 1);
dim3 grid((unsigned)((nn + TILE_P - 1) / TILE_P), (unsigned)bs, 1);
pack_proj_tiled2_f16<D_TILE, TILE_P><<<grid, block>>>(
(const half*)proj.data_ptr(),
(const float*)mask.data_ptr(),
(half*)left.data_ptr(),
(half*)right.data_ptr(),
(half*)og.data_ptr(),
nn,
hidden);
checkCuda(cudaGetLastError(), "pack_proj_tiled2_f16_128");
} else {
throw std::runtime_error("hidden_dim too large");
}
}
// --------------- LN2 + gate + store ---------------
template<int HMAX, int PTILE>
__global__ void ln2_gate_store_warp_f16(
const half* __restrict__ out_acc,
const half* __restrict__ ogate,
const float* __restrict__ w,
const float* __restrict__ b,
half* __restrict__ out_norm,
int64_t nn,
int hidden) {
int bb = (int)blockIdx.y;
int64_t p0 = (int64_t)blockIdx.x * (int64_t)PTILE;
int tid = (int)threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
constexpr int WP = HMAX / 32;
int p_l = warp / WP;
int w_in = warp - p_l * WP;
int64_t p = p0 + (int64_t)p_l;
int d = w_in * 32 + lane;
float x = 0.0f;
float g = 0.0f;
if (p < nn && d < hidden) {
int64_t idx = ((int64_t)bb * (int64_t)hidden + (int64_t)d) * nn + p;
x = __half2float(out_acc[idx]);
g = __half2float(ogate[idx]);
}
float s = (d < hidden && p < nn) ? x : 0.0f;
float ss = (d < hidden && p < nn) ? x * x : 0.0f;
s = warp_sum(s);
ss = warp_sum(ss);
__shared__ float sh_sum[PTILE * WP];
__shared__ float sh_sq[PTILE * WP];
if (lane == 0) {
sh_sum[p_l * WP + w_in] = s;
sh_sq[p_l * WP + w_in] = ss;
}
__syncthreads();
__shared__ float mean_sh[PTILE];
__shared__ float inv_sh[PTILE];
if (lane == 0 && w_in == 0) {
float sum = 0.0f;
float sq = 0.0f;
#pragma unroll
for (int t = 0; t < WP; ++t) {
sum += sh_sum[p_l * WP + t];
sq += sh_sq[p_l * WP + t];
}
float inv_n = 1.0f / (float)hidden;
float mean = sum * inv_n;
float var = sq * inv_n - mean * mean;
mean_sh[p_l] = mean;
inv_sh[p_l] = rsqrtf(var + 1e-5f);
}
__syncthreads();
if (p < nn && d < hidden) {
float mean = mean_sh[p_l];
float inv = inv_sh[p_l];
float y = (x - mean) * inv * w[d] + b[d];
float yg = y * g;
int64_t row = (int64_t)bb * nn + p;
out_norm[row * (int64_t)hidden + (int64_t)d] = __float2half_rn(yg);
}
}
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 nn = (int64_t)n * (int64_t)n;
if (hidden <= 32) {
constexpr int HMAX = 32;
constexpr int PTILE = 16;
dim3 block(HMAX * PTILE, 1, 1);
dim3 grid((unsigned)((nn + PTILE - 1) / PTILE), (unsigned)bs, 1);
ln2_gate_store_warp_f16<HMAX, PTILE><<<grid, block>>>(
(const half*)out_acc.data_ptr(),
(const half*)og.data_ptr(),
(const float*)w.data_ptr(),
(const float*)b.data_ptr(),
(half*)out_norm.data_ptr(),
nn,
hidden);
checkCuda(cudaGetLastError(), "ln2_gate_store_warp_f16_32");
} else if (hidden <= 64) {
constexpr int HMAX = 64;
constexpr int PTILE = 8;
dim3 block(HMAX * PTILE, 1, 1);
dim3 grid((unsigned)((nn + PTILE - 1) / PTILE), (unsigned)bs, 1);
ln2_gate_store_warp_f16<HMAX, PTILE><<<grid, block>>>(
(const half*)out_acc.data_ptr(),
(const half*)og.data_ptr(),
(const float*)w.data_ptr(),
(const float*)b.data_ptr(),
(half*)out_norm.data_ptr(),
nn,
hidden);
checkCuda(cudaGetLastError(), "ln2_gate_store_warp_f16_64");
} else if (hidden <= 128) {
constexpr int HMAX = 128;
constexpr int PTILE = 4;
dim3 block(HMAX * PTILE, 1, 1);
dim3 grid((unsigned)((nn + PTILE - 1) / PTILE), (unsigned)bs, 1);
ln2_gate_store_warp_f16<HMAX, PTILE><<<grid, block>>>(
(const half*)out_acc.data_ptr(),
(const half*)og.data_ptr(),
(const float*)w.data_ptr(),
(const float*)b.data_ptr(),
(half*)out_norm.data_ptr(),
nn,
hidden);
checkCuda(cudaGetLastError(), "ln2_gate_store_warp_f16_128");
} else {
throw std::runtime_error("hidden_dim too large");
}
}
// ---------------- GEMM helpers ----------------
static void gemm_x_wt_f16_f16(
cublasHandle_t h,
const half* x_row,
const half* w_row,
half* y_row,
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_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,
const half* w_row,
float* y_row,
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_f16(
cublasHandle_t h,
const half* left_row,
const half* right_row,
half* out_row,
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");
}
} // 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;
int64_t nn = (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;
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, (int)hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
auto right = torch::empty({bs, (int)hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
auto og = torch::empty({bs, (int)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::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);
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);
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_mod3"
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 · 950 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 417619.
⋯ 6 unchanged linesimport torch---_EXT = None_EXT_LOCK = Nonedef _lazy_import_extension_utils():-from torch.utils.cpp_extension import load_inlinereturn load_inline⋯ 13 unchanged linesload_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>⋯ 16 unchanged linescuda_src = r"""#include <torch/extension.h>- #include <ATen/cuda/CUDAContext.h>#include <cuda.h>#include <cuda_fp16.h>#include <cublas_v2.h>⋯ 18 unchanged linescublasHandle_t handle = nullptr;CublasHandleHolder() {checkCublas(cublasCreate(&handle), "cublasCreate");- // 默认数学模式即可;这里不做额外设置,减少环境依赖+ checkCublas(cublasSetMathMode(handle, CUBLAS_TENSOR_OP_MATH), "cublasSetMathMode");}~CublasHandleHolder() {if (handle) {⋯ 11 unchanged lines}__device__ __forceinline__ float warp_sum(float v) {- for (int d = 16; d > 0; d >>= 1) {- v += __shfl_down_sync(0xffffffff, v, d);- }+ v += __shfl_down_sync(0xffffffff, v, 16);+ v += __shfl_down_sync(0xffffffff, v, 8);+ v += __shfl_down_sync(0xffffffff, v, 4);+ v += __shfl_down_sync(0xffffffff, v, 2);+ v += __shfl_down_sync(0xffffffff, v, 1);return v;}__device__ __forceinline__ float fast_sigmoid(float x) {- // 使用 __expf,精度在题面容忍范围内通常足够float z = __expf(-x);return 1.0f / (1.0f + z);}+ // ---------------- LN1 ----------------+__global__ void ln1_128_f16(const float* __restrict__ x,const float* __restrict__ w,⋯ 2 unchanged linesint64_t rows) {int64_t row = (int64_t)blockIdx.x;if (row >= rows) return;- int lane = (int)threadIdx.x; // 0..31+ int lane = (int)threadIdx.x;const float4* x4 = reinterpret_cast<const float4*>(x + row * 128);float4 v = x4[lane];⋯ 1 unchanged linesfloat 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 得到全和,需要广播到全 warps = __shfl_sync(0xffffffff, s, 0);ss = __shfl_sync(0xffffffff, ss, 0);float mean = s * (1.0f / 128.0f);⋯ 12 unchanged lineshalf2 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;+ int warp = 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 sh_sum[3];+ __shared__ float sh_sq[3];+ if (lane == 0) {+ sh_sum[warp] = s;+ sh_sq[warp] = ss;+ }+ __syncthreads();++ __shared__ float mean_sh;+ __shared__ float inv_sh;+ if (tid == 0) {+ float sum = sh_sum[0] + sh_sum[1] + sh_sum[2];+ float sq = sh_sq[0] + sh_sq[1] + sh_sq[2];+ float mean = sum * (1.0f / 384.0f);+ float var = sq * (1.0f / 384.0f) - mean * mean;+ mean_sh = mean;+ inv_sh = rsqrtf(var + 1e-5f);+ }+ __syncthreads();++ float mean = mean_sh;+ float inv = inv_sh;++ 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;++ int64_t off = row * 384 + (int64_t)tid * 4;+ half2 h0 = __floats2half2_rn(y0, y1);+ half2 h1 = __floats2half2_rn(y2, y3);+ half2* y2p = reinterpret_cast<half2*>(y + off);+ y2p[0] = h0;+ y2p[1] = h1;+ }+__global__ void ln1_generic_f16(const float* __restrict__ x,const float* __restrict__ w,⋯ 4 unchanged linesint64_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;⋯ 29 unchanged lines}}- __global__ void pack_proj_f16(+ 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 (proj -> left/right/ogate) ---------------++ template<int D_TILE, int TILE_P>+ __global__ void pack_proj_tiled2_f16(const half* __restrict__ proj,const float* __restrict__ mask,half* __restrict__ left,half* __restrict__ right,half* __restrict__ ogate,- int bs,- int n,+ int64_t nn,int hidden) {- int64_t row = (int64_t)blockIdx.x;- int d = (int)threadIdx.x;- if (d >= hidden) return;+ int b = (int)blockIdx.y;+ int64_t p0 = (int64_t)blockIdx.x * (int64_t)TILE_P;- 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);+ int tid = (int)threadIdx.x;+ int p_l = tid / D_TILE;+ int d0 = tid - p_l * D_TILE;+ int64_t p = p0 + (int64_t)p_l;- float m = mask[row];+ __shared__ float m_sh[TILE_P];+ if (d0 == 0) {+ float mv = 0.0f;+ if (p < nn) {+ mv = mask[(int64_t)b * nn + p];+ }+ m_sh[p_l] = mv;+ }+ __syncthreads();- int out = hidden * 5;- const half* p = proj + row * out;+ __shared__ half sh_l[128 * TILE_P];+ __shared__ half sh_r[128 * TILE_P];+ __shared__ half sh_g[128 * TILE_P];- 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]));+ half out_l0 = __float2half_rn(0.0f);+ half out_r0 = __float2half_rn(0.0f);+ half out_g0 = __float2half_rn(0.0f);+ half out_l1 = __float2half_rn(0.0f);+ half out_r1 = __float2half_rn(0.0f);+ half out_g1 = __float2half_rn(0.0f);- float l2 = l * gl * m;- float r2 = r * gr * m;+ if (p < nn) {+ int64_t row = (int64_t)b * nn + p;+ int out_ch = hidden * 5;+ const half* base = proj + row * (int64_t)out_ch;+ float m = m_sh[p_l];- 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);+ int d = d0;+ if (d < hidden) {+ float l = __half2float(base[d]);+ float r = __half2float(base[hidden + d]);+ float gl = fast_sigmoid(__half2float(base[2 * hidden + d]));+ float gr = fast_sigmoid(__half2float(base[3 * hidden + d]));+ float go = fast_sigmoid(__half2float(base[4 * hidden + d]));+ out_l0 = __float2half_rn(l * gl * m);+ out_r0 = __float2half_rn(r * gr * m);+ out_g0 = __float2half_rn(go);+ }++ int d1 = d0 + 64;+ if (d1 < hidden) {+ float l = __half2float(base[d1]);+ float r = __half2float(base[hidden + d1]);+ float gl = fast_sigmoid(__half2float(base[2 * hidden + d1]));+ float gr = fast_sigmoid(__half2float(base[3 * hidden + d1]));+ float go = fast_sigmoid(__half2float(base[4 * hidden + d1]));+ out_l1 = __float2half_rn(l * gl * m);+ out_r1 = __float2half_rn(r * gr * m);+ out_g1 = __float2half_rn(go);+ }+ }++ sh_l[d0 * TILE_P + p_l] = out_l0;+ sh_r[d0 * TILE_P + p_l] = out_r0;+ sh_g[d0 * TILE_P + p_l] = out_g0;+ sh_l[(d0 + 64) * TILE_P + p_l] = out_l1;+ sh_r[(d0 + 64) * TILE_P + p_l] = out_r1;+ sh_g[(d0 + 64) * TILE_P + p_l] = out_g1;+ __syncthreads();++ int d2 = tid / TILE_P; // 0..63+ int p2 = tid - d2 * TILE_P;+ int64_t p_out = p0 + (int64_t)p2;+ if (p_out < nn) {+ int d = d2;+ if (d < hidden) {+ int64_t out_idx = ((int64_t)b * (int64_t)hidden + (int64_t)d) * nn + p_out;+ left[out_idx] = sh_l[d * TILE_P + p2];+ right[out_idx] = sh_r[d * TILE_P + p2];+ ogate[out_idx] = sh_g[d * TILE_P + p2];+ }+ int d3 = d2 + 64;+ if (d3 < hidden) {+ int64_t out_idx = ((int64_t)b * (int64_t)hidden + (int64_t)d3) * nn + p_out;+ left[out_idx] = sh_l[d3 * TILE_P + p2];+ right[out_idx] = sh_r[d3 * TILE_P + p2];+ ogate[out_idx] = sh_g[d3 * TILE_P + p2];+ }+ }}- template<int MAX_H>- __global__ void ln2_gate_store_f16(- const float* __restrict__ out_acc,+ template<int MAX_H, int TILE_P>+ __global__ void pack_proj_tiled_f16(+ const half* __restrict__ proj,+ const float* __restrict__ mask,+ half* __restrict__ left,+ half* __restrict__ right,+ half* __restrict__ ogate,+ int64_t nn,+ int hidden) {+ int b = (int)blockIdx.y;+ int64_t p0 = (int64_t)blockIdx.x * (int64_t)TILE_P;++ int tid = (int)threadIdx.x;+ int p_l = tid / MAX_H;+ int d = tid - p_l * MAX_H;+ int64_t p = p0 + (int64_t)p_l;++ __shared__ float m_sh[TILE_P];+ if (d == 0) {+ float mv = 0.0f;+ if (p < nn) {+ mv = mask[(int64_t)b * nn + p];+ }+ m_sh[p_l] = mv;+ }+ __syncthreads();++ __shared__ half sh_l[MAX_H * TILE_P];+ __shared__ half sh_r[MAX_H * TILE_P];+ __shared__ half sh_g[MAX_H * TILE_P];++ half out_l = __float2half_rn(0.0f);+ half out_r = __float2half_rn(0.0f);+ half out_g = __float2half_rn(0.0f);++ if (d < hidden && p < nn) {+ int64_t row = (int64_t)b * nn + p;+ int out_ch = hidden * 5;+ const half* base = proj + row * (int64_t)out_ch;++ float l = __half2float(base[d]);+ float r = __half2float(base[hidden + d]);+ float gl = fast_sigmoid(__half2float(base[2 * hidden + d]));+ float gr = fast_sigmoid(__half2float(base[3 * hidden + d]));+ float go = fast_sigmoid(__half2float(base[4 * hidden + d]));++ float m = m_sh[p_l];+ out_l = __float2half_rn(l * gl * m);+ out_r = __float2half_rn(r * gr * m);+ out_g = __float2half_rn(go);+ }++ sh_l[d * TILE_P + p_l] = out_l;+ sh_r[d * TILE_P + p_l] = out_r;+ sh_g[d * TILE_P + p_l] = out_g;+ __syncthreads();++ int d2 = tid / TILE_P;+ int p2 = tid - d2 * TILE_P;+ int64_t p_out = p0 + (int64_t)p2;+ if (d2 < hidden && p_out < nn) {+ int64_t out_idx = ((int64_t)b * (int64_t)hidden + (int64_t)d2) * nn + p_out;+ left[out_idx] = sh_l[d2 * TILE_P + p2];+ right[out_idx] = sh_r[d2 * TILE_P + p2];+ ogate[out_idx] = sh_g[d2 * TILE_P + p2];+ }+ }++ 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 nn = (int64_t)n * (int64_t)n;+ if (hidden <= 32) {+ constexpr int MAX_H = 32;+ constexpr int TILE_P = 16;+ dim3 block(MAX_H * TILE_P, 1, 1);+ dim3 grid((unsigned)((nn + TILE_P - 1) / TILE_P), (unsigned)bs, 1);+ pack_proj_tiled_f16<MAX_H, TILE_P><<<grid, block>>>(+ (const half*)proj.data_ptr(),+ (const float*)mask.data_ptr(),+ (half*)left.data_ptr(),+ (half*)right.data_ptr(),+ (half*)og.data_ptr(),+ nn,+ hidden);+ checkCuda(cudaGetLastError(), "pack_proj_tiled_f16_32");+ } else if (hidden <= 64) {+ constexpr int MAX_H = 64;+ constexpr int TILE_P = 8;+ dim3 block(MAX_H * TILE_P, 1, 1);+ dim3 grid((unsigned)((nn + TILE_P - 1) / TILE_P), (unsigned)bs, 1);+ pack_proj_tiled_f16<MAX_H, TILE_P><<<grid, block>>>(+ (const half*)proj.data_ptr(),+ (const float*)mask.data_ptr(),+ (half*)left.data_ptr(),+ (half*)right.data_ptr(),+ (half*)og.data_ptr(),+ nn,+ hidden);+ checkCuda(cudaGetLastError(), "pack_proj_tiled_f16_64");+ } else if (hidden <= 128) {+ constexpr int D_TILE = 64;+ constexpr int TILE_P = 8;+ dim3 block(D_TILE * TILE_P, 1, 1);+ dim3 grid((unsigned)((nn + TILE_P - 1) / TILE_P), (unsigned)bs, 1);+ pack_proj_tiled2_f16<D_TILE, TILE_P><<<grid, block>>>(+ (const half*)proj.data_ptr(),+ (const float*)mask.data_ptr(),+ (half*)left.data_ptr(),+ (half*)right.data_ptr(),+ (half*)og.data_ptr(),+ nn,+ hidden);+ checkCuda(cudaGetLastError(), "pack_proj_tiled2_f16_128");+ } else {+ throw std::runtime_error("hidden_dim too large");+ }+ }++ // --------------- LN2 + gate + store ---------------++ template<int HMAX, int PTILE>+ __global__ void ln2_gate_store_warp_f16(+ const half* __restrict__ out_acc,const half* __restrict__ ogate,const float* __restrict__ w,const float* __restrict__ b,half* __restrict__ out_norm,- int bs,- int n,+ int64_t nn,int hidden) {- // blockIdx.x 对应 (b,i,j)- int64_t row = (int64_t)blockIdx.x;+ int bb = (int)blockIdx.y;+ int64_t p0 = (int64_t)blockIdx.x * (int64_t)PTILE;+int tid = (int)threadIdx.x;- if (tid >= MAX_H) return;+ int warp = tid >> 5;+ int lane = tid & 31;- 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);+ constexpr int WP = HMAX / 32;+ int p_l = warp / WP;+ int w_in = warp - p_l * WP;+ int64_t p = p0 + (int64_t)p_l;+ int d = w_in * 32 + lane;- 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;+ float x = 0.0f;+ float g = 0.0f;+ if (p < nn && d < hidden) {+ int64_t idx = ((int64_t)bb * (int64_t)hidden + (int64_t)d) * nn + p;+ x = __half2float(out_acc[idx]);+ g = __half2float(ogate[idx]);}- // 归约:MAX_H 固定,使用共享内存- __shared__ float shm_sum[MAX_H];- __shared__ float shm_sq[MAX_H];- shm_sum[tid] = v;- shm_sq[tid] = vv;+ float s = (d < hidden && p < nn) ? x : 0.0f;+ float ss = (d < hidden && p < nn) ? x * x : 0.0f;+ s = warp_sum(s);+ ss = warp_sum(ss);++ __shared__ float sh_sum[PTILE * WP];+ __shared__ float sh_sq[PTILE * WP];+ if (lane == 0) {+ sh_sum[p_l * WP + w_in] = s;+ sh_sq[p_l * WP + w_in] = ss;+ }__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];+ __shared__ float mean_sh[PTILE];+ __shared__ float inv_sh[PTILE];+ if (lane == 0 && w_in == 0) {+ float sum = 0.0f;+ float sq = 0.0f;+ #pragma unroll+ for (int t = 0; t < WP; ++t) {+ sum += sh_sum[p_l * WP + t];+ sq += sh_sq[p_l * WP + t];}- __syncthreads();+ float inv_n = 1.0f / (float)hidden;+ float mean = sum * inv_n;+ float var = sq * inv_n - mean * mean;+ mean_sh[p_l] = mean;+ inv_sh[p_l] = rsqrtf(var + 1e-5f);}+ __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);+ if (p < nn && d < hidden) {+ float mean = mean_sh[p_l];+ float inv = inv_sh[p_l];+ float y = (x - mean) * inv * w[d] + b[d];+ float yg = y * g;+ int64_t row = (int64_t)bb * nn + p;+ out_norm[row * (int64_t)hidden + (int64_t)d] = __float2half_rn(yg);}}- 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;+ 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 nn = (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(),+ constexpr int HMAX = 32;+ constexpr int PTILE = 16;+ dim3 block(HMAX * PTILE, 1, 1);+ dim3 grid((unsigned)((nn + PTILE - 1) / PTILE), (unsigned)bs, 1);+ ln2_gate_store_warp_f16<HMAX, PTILE><<<grid, block>>>(+ (const half*)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");+ nn,+ hidden);+ checkCuda(cudaGetLastError(), "ln2_gate_store_warp_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(),+ constexpr int HMAX = 64;+ constexpr int PTILE = 8;+ dim3 block(HMAX * PTILE, 1, 1);+ dim3 grid((unsigned)((nn + PTILE - 1) / PTILE), (unsigned)bs, 1);+ ln2_gate_store_warp_f16<HMAX, PTILE><<<grid, block>>>(+ (const half*)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");+ nn,+ hidden);+ checkCuda(cudaGetLastError(), "ln2_gate_store_warp_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(),+ constexpr int HMAX = 128;+ constexpr int PTILE = 4;+ dim3 block(HMAX * PTILE, 1, 1);+ dim3 grid((unsigned)((nn + PTILE - 1) / PTILE), (unsigned)bs, 1);+ ln2_gate_store_warp_f16<HMAX, PTILE><<<grid, block>>>(+ (const half*)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");+ nn,+ hidden);+ checkCuda(cudaGetLastError(), "ln2_gate_store_warp_f16_128");} else {throw std::runtime_error("hidden_dim too large");}}+ // ---------------- GEMM helpers ----------------+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]+ const half* x_row,+ const half* w_row,+ half* y_row,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(⋯ 12 unchanged linesstatic 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]+ const half* x_row,+ const half* w_row,+ float* y_row,int64_t M,int64_t N,int64_t K) {⋯ 13 unchanged lines"cublasGemmEx");}- static void gemm_contract_batched_f16_f32(+ static void gemm_contract_batched_f16_f16(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]+ const half* left_row,+ const half* right_row,+ half* out_row,int batch,int n) {float alpha = 1.0f;⋯ 2 unchanged lineslong 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^TcheckCublas(cublasGemmStridedBatchedEx(h,⋯ 3 unchanged linesright_row, CUDA_R_16F, n, strideB,left_row, CUDA_R_16F, n, strideA,&beta,- out_row, CUDA_R_32F, n, strideC,+ out_row, CUDA_R_16F, n, strideC,batch,CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),"cublasGemmStridedBatchedEx");⋯ 31 unchanged linesint bs = (int)x.size(0);int n = (int)x.size(1);int64_t M = (int64_t)bs * (int64_t)n * (int64_t)n;+ int64_t nn = (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));⋯ 5 unchanged linesauto* holder = get_cublas();cublasHandle_t h = holder->handle;- // gemm1: proj = xhat @ w_cat^Tgemm_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));+ auto left = torch::empty({bs, (int)hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));+ auto right = torch::empty({bs, (int)hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));+ auto og = torch::empty({bs, (int)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));+ auto out_acc = torch::empty({bs * (int)hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));- gemm_contract_batched_f16_f32(+ gemm_contract_batched_f16_f16(h,(const half*)left3.data_ptr(),(const half*)right3.data_ptr(),- (float*)out_acc.data_ptr(),+ (half*)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^Tauto 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);⋯ 1 unchanged lines}"""- name = "trimul_ext_mod"+ name = "trimul_ext_mod3"extra_cuda_cflags = ["-O3",⋯ 34 unchanged linesdef _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()),⋯ 12 unchanged linesw_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):⋯ 7 unchanged linesif 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)⋯ 17 unchanged linesif x.dtype != torch.float32:x = x.to(dtype=torch.float32)--if mask.dtype != torch.float32:mask = mask.to(dtype=torch.float32)⋯ 9 unchanged linesif x.shape[3] != dim:raise RuntimeError("dim mismatch")-for k in ("norm.weight","norm.bias",⋯ 11 unchanged linesif 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,):⋯ 4 unchanged linesif 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 · 879 diff lines total
Best evidence level for this revision: reported
JSON