submission 419233
shiyegao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 827 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-419233?include=source"interfacepython
Compatibility
measured onNVIDIA A100
declared hardwareNVIDIA A100
architecturessm_80
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:ed693ac809ebc05286e2c564d43c135cd25a3ba511224449c667b1dc6e13fbb4
license declaredunknown
license concludedunknown
authorsshiyegao
imported2026-08-15
Kernel source
submission.py827 lines
from __future__ import annotations
from typing import Any, Dict, Optional, Tuple
import ctypes
import torch
import cutlass
import cutlass.cute as cute
from cutlass.cute.runtime import make_ptr
from cutlass.cutlass_dsl import for_generate, if_generate, yield_out
class _Cublas:
_STATUS_SUCCESS = 0
_OP_N = 0
_CUDA_R_32F = 0
_CUDA_R_16F = 2
_COMPUTE_32F = 68
_GEMM_DEFAULT_TENSOR_OP = 99
_MATH_TENSOR_OP = 1
def __init__(self) -> None:
lib = None
last_err: Optional[Exception] = None
for name in ("libcublas.so.13", "libcublas.so"):
try:
lib = ctypes.CDLL(name)
break
except OSError as exc:
last_err = exc
continue
if lib is None:
raise RuntimeError(f"加载 cuBLAS 失败:{last_err}")
self.lib = lib
self.handle = ctypes.c_void_p()
self.lib.cublasCreate_v2.argtypes = [ctypes.POINTER(ctypes.c_void_p)]
self.lib.cublasCreate_v2.restype = ctypes.c_int
self.lib.cublasDestroy_v2.argtypes = [ctypes.c_void_p]
self.lib.cublasDestroy_v2.restype = ctypes.c_int
self.lib.cublasSetMathMode.argtypes = [ctypes.c_void_p, ctypes.c_int]
self.lib.cublasSetMathMode.restype = ctypes.c_int
self.lib.cublasGemmEx.argtypes = [
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
ctypes.c_int,
ctypes.c_int,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
ctypes.c_int,
ctypes.c_int,
]
self.lib.cublasGemmEx.restype = ctypes.c_int
self.lib.cublasGemmStridedBatchedEx.argtypes = [
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
ctypes.c_int,
ctypes.c_int,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
ctypes.c_longlong,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
ctypes.c_longlong,
ctypes.c_void_p,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
ctypes.c_longlong,
ctypes.c_int,
ctypes.c_int,
ctypes.c_int,
]
self.lib.cublasGemmStridedBatchedEx.restype = ctypes.c_int
st = self.lib.cublasCreate_v2(ctypes.byref(self.handle))
if st != self._STATUS_SUCCESS:
raise RuntimeError(f"cublasCreate_v2 失败:status={st}")
st = self.lib.cublasSetMathMode(self.handle, self._MATH_TENSOR_OP)
if st != self._STATUS_SUCCESS:
raise RuntimeError(f"cublasSetMathMode 失败:status={st}")
def __del__(self) -> None:
try:
if getattr(self, "handle", None):
self.lib.cublasDestroy_v2(self.handle)
except Exception:
pass
@staticmethod
def _ptr(t: torch.Tensor) -> ctypes.c_void_p:
return ctypes.c_void_p(int(t.data_ptr()))
def gemm_rm_f16_f16_to_f16(self, a: torch.Tensor, b: torch.Tensor, c: torch.Tensor) -> None:
m, k = a.shape
kb, n = b.shape
if kb != k:
raise RuntimeError("GEMM 维度不匹配。")
if c.shape != (m, n):
raise RuntimeError("GEMM 输出维度不匹配。")
if a.dtype != torch.float16 or b.dtype != torch.float16 or c.dtype != torch.float16:
raise RuntimeError("gemm_rm_f16_f16_to_f16 仅支持 fp16 输入输出。")
if not (a.is_cuda and b.is_cuda and c.is_cuda):
raise RuntimeError("gemm_rm_f16_f16_to_f16 仅支持 CUDA 张量。")
alpha = ctypes.c_float(1.0)
beta = ctypes.c_float(0.0)
st = self.lib.cublasGemmEx(
self.handle,
self._OP_N,
self._OP_N,
ctypes.c_int(n),
ctypes.c_int(m),
ctypes.c_int(k),
ctypes.c_void_p(ctypes.addressof(alpha)),
self._ptr(b),
self._CUDA_R_16F,
ctypes.c_int(n),
self._ptr(a),
self._CUDA_R_16F,
ctypes.c_int(k),
ctypes.c_void_p(ctypes.addressof(beta)),
self._ptr(c),
self._CUDA_R_16F,
ctypes.c_int(n),
ctypes.c_int(self._COMPUTE_32F),
ctypes.c_int(self._GEMM_DEFAULT_TENSOR_OP),
)
if st != self._STATUS_SUCCESS:
raise RuntimeError(f"cublasGemmEx(fp16->fp16) 失败:status={st}")
def gemm_rm_f16_f16_to_f32(self, a: torch.Tensor, b: torch.Tensor, c: torch.Tensor) -> None:
m, k = a.shape
kb, n = b.shape
if kb != k:
raise RuntimeError("GEMM 维度不匹配。")
if c.shape != (m, n):
raise RuntimeError("GEMM 输出维度不匹配。")
if a.dtype != torch.float16 or b.dtype != torch.float16 or c.dtype != torch.float32:
raise RuntimeError("gemm_rm_f16_f16_to_f32 仅支持 A/B fp16,C fp32。")
if not (a.is_cuda and b.is_cuda and c.is_cuda):
raise RuntimeError("gemm_rm_f16_f16_to_f32 仅支持 CUDA 张量。")
alpha = ctypes.c_float(1.0)
beta = ctypes.c_float(0.0)
st = self.lib.cublasGemmEx(
self.handle,
self._OP_N,
self._OP_N,
ctypes.c_int(n),
ctypes.c_int(m),
ctypes.c_int(k),
ctypes.c_void_p(ctypes.addressof(alpha)),
self._ptr(b),
self._CUDA_R_16F,
ctypes.c_int(n),
self._ptr(a),
self._CUDA_R_16F,
ctypes.c_int(k),
ctypes.c_void_p(ctypes.addressof(beta)),
self._ptr(c),
self._CUDA_R_32F,
ctypes.c_int(n),
ctypes.c_int(self._COMPUTE_32F),
ctypes.c_int(self._GEMM_DEFAULT_TENSOR_OP),
)
if st != self._STATUS_SUCCESS:
raise RuntimeError(f"cublasGemmEx(fp16->fp32) 失败:status={st}")
def gemm_rm_f16_f16_to_f32_strided_batched(
self,
a: torch.Tensor,
b: torch.Tensor,
c: torch.Tensor,
batch: int,
m: int,
n: int,
k: int,
stride_a: int,
stride_b: int,
stride_c: int,
) -> None:
if a.dtype != torch.float16 or b.dtype != torch.float16 or c.dtype != torch.float32:
raise RuntimeError("batched GEMM 仅支持 A/B fp16,C fp32。")
if not (a.is_cuda and b.is_cuda and c.is_cuda):
raise RuntimeError("batched GEMM 仅支持 CUDA 张量。")
alpha = ctypes.c_float(1.0)
beta = ctypes.c_float(0.0)
st = self.lib.cublasGemmStridedBatchedEx(
self.handle,
self._OP_N,
self._OP_N,
ctypes.c_int(n),
ctypes.c_int(m),
ctypes.c_int(k),
ctypes.c_void_p(ctypes.addressof(alpha)),
self._ptr(b),
self._CUDA_R_16F,
ctypes.c_int(n),
ctypes.c_longlong(stride_b),
self._ptr(a),
self._CUDA_R_16F,
ctypes.c_int(k),
ctypes.c_longlong(stride_a),
ctypes.c_void_p(ctypes.addressof(beta)),
self._ptr(c),
self._CUDA_R_32F,
ctypes.c_int(n),
ctypes.c_longlong(stride_c),
ctypes.c_int(batch),
ctypes.c_int(self._COMPUTE_32F),
ctypes.c_int(self._GEMM_DEFAULT_TENSOR_OP),
)
if st != self._STATUS_SUCCESS:
raise RuntimeError(f"cublasGemmStridedBatchedEx 失败:status={st}")
_CUBLAS: Optional[_Cublas] = None
def _get_cublas() -> _Cublas:
global _CUBLAS
if _CUBLAS is None:
_CUBLAS = _Cublas()
return _CUBLAS
class _LayerNormLastDimF32ToF16:
def __init__(self, threads: int = 256) -> None:
self.threads = int(threads)
self.warps = self.threads // 32
@cute.jit
def __call__(
self,
x_ptr: "cute.Pointer",
w_ptr: "cute.Pointer",
b_ptr: "cute.Pointer",
y_ptr: "cute.Pointer",
problem: tuple,
):
bs, n, d = problem
stride_bs = n * n * d
stride_i = n * d
stride_j = d
x = cute.make_tensor(
x_ptr,
cute.make_layout((bs, n, n, d), stride=(stride_bs, stride_i, stride_j, 1)),
)
w = cute.make_tensor(w_ptr, cute.make_layout((d,), stride=(1,)))
b = cute.make_tensor(b_ptr, cute.make_layout((d,), stride=(1,)))
y = cute.make_tensor(
y_ptr,
cute.make_layout((bs, n, n, d), stride=(stride_bs, stride_i, stride_j, 1)),
)
total = bs * n * n
grid = (total + self.warps - 1) // self.warps
self.kernel(x, w, b, y, bs, n, d).launch(
grid=[grid, 1, 1],
block=[self.threads, 1, 1],
)
return
@cute.kernel
def kernel(
self,
x: "cute.Tensor",
w: "cute.Tensor",
b: "cute.Tensor",
y: "cute.Tensor",
bs: int,
n: int,
d: int,
):
tx, _, _ = cute.arch.thread_idx()
bx, _, _ = cute.arch.block_idx()
warp_id = tx >> 5
lane = tx & 31
idx = bx * self.warps + warp_id
def _do_one():
j = idx % n
t0 = idx // n
i = t0 % n
bb = t0 // n
sum0_init = cutlass.Float32(0.0)
sumsq0_init = cutlass.Float32(0.0)
for dd, acc, acc_out in for_generate(lane, d, 32, iter_args=[sum0_init, sumsq0_init]):
sum0_it = acc[0]
sumsq0_it = acc[1]
v = x[bb, i, j, dd]
sum0_it = sum0_it + v
sumsq0_it = sumsq0_it + v * v
yield_out([sum0_it, sumsq0_it])
sum0 = cute.arch.warp_reduction_sum(acc_out[0])
sumsq0 = cute.arch.warp_reduction_sum(acc_out[1])
inv_d = cutlass.Float32(1.0) / cutlass.Float32(d)
mean = sum0 * inv_d
var = sumsq0 * inv_d - mean * mean
inv_std = cute.rsqrt(var + cutlass.Float32(1e-5), fastmath=True)
for dd in for_generate(lane, d, 32):
v = x[bb, i, j, dd]
nrm = (v - mean) * inv_std
out = nrm * w[dd] + b[dd]
y[bb, i, j, dd] = out.to(cutlass.Float16)
yield_out()
if_generate(idx < bs * n * n, _do_one)
def _sigmoid_f16(x: cutlass.Float16) -> cutlass.Float16:
xx = x.to(cutlass.Float32)
ee = cute.exp(cutlass.Float32(0.0) - xx, fastmath=True)
yy = cutlass.Float32(1.0) / (cutlass.Float32(1.0) + ee)
return yy.to(cutlass.Float16)
class _ProcessProj:
def __init__(self, threads: int = 256) -> None:
self.threads = int(threads)
@cute.jit
def __call__(
self,
proj_ptr: "cute.Pointer",
mask_ptr: "cute.Pointer",
left_ptr: "cute.Pointer",
right_ptr: "cute.Pointer",
gate_ptr: "cute.Pointer",
problem: tuple,
):
bs, n, h = problem
stride_proj_bs = n * n * (5 * h)
stride_proj_i = n * (5 * h)
stride_proj_j = 5 * h
proj = cute.make_tensor(
proj_ptr,
cute.make_layout((bs, n, n, 5 * h), stride=(stride_proj_bs, stride_proj_i, stride_proj_j, 1)),
)
mask = cute.make_tensor(mask_ptr, cute.make_layout((bs, n, n), stride=(n * n, n, 1)))
left = cute.make_tensor(left_ptr, cute.make_layout((bs, h, n, n), stride=(h * n * n, n * n, n, 1)))
right = cute.make_tensor(right_ptr, cute.make_layout((bs, h, n, n), stride=(h * n * n, n * n, n, 1)))
gate = cute.make_tensor(gate_ptr, cute.make_layout((bs, n, n, h), stride=(n * n * h, n * h, h, 1)))
total = bs * n * n * h
grid = (total + self.threads - 1) // self.threads
self.kernel(proj, mask, left, right, gate, bs, n, h).launch(
grid=[grid, 1, 1],
block=[self.threads, 1, 1],
)
return
@cute.kernel
def kernel(
self,
proj: "cute.Tensor",
mask: "cute.Tensor",
left: "cute.Tensor",
right: "cute.Tensor",
gate: "cute.Tensor",
bs: int,
n: int,
h: int,
):
tx, _, _ = cute.arch.thread_idx()
bx, _, _ = cute.arch.block_idx()
bdx, _, _ = cute.arch.block_dim()
idx = bx * bdx + tx
total = bs * n * n * h
def _do_one():
hh = idx % h
t0 = idx // h
j = t0 % n
t1 = t0 // n
i = t1 % n
bb = t1 // n
m = mask[bb, i, j].to(cutlass.Float16)
lp = proj[bb, i, j, hh]
rp = proj[bb, i, j, hh + h]
lg = proj[bb, i, j, hh + 2 * h]
rg = proj[bb, i, j, hh + 3 * h]
og = proj[bb, i, j, hh + 4 * h]
gl = _sigmoid_f16(lg)
gr = _sigmoid_f16(rg)
go = _sigmoid_f16(og)
left[bb, hh, i, j] = (lp * gl * m).to(cutlass.Float16)
right[bb, hh, j, i] = (rp * gr * m).to(cutlass.Float16)
gate[bb, i, j, hh] = go
if_generate(idx < total, _do_one)
class _LayerNormHiddenF32ToF16:
def __init__(self, threads: int = 256) -> None:
self.threads = int(threads)
self.warps = self.threads // 32
@cute.jit
def __call__(
self,
x_ptr: "cute.Pointer",
w_ptr: "cute.Pointer",
b_ptr: "cute.Pointer",
g_ptr: "cute.Pointer",
y_ptr: "cute.Pointer",
problem: tuple,
):
bs, n, h = problem
stride_bs = h * n * n
stride_h = n * n
stride_i = n
x = cute.make_tensor(x_ptr, cute.make_layout((bs, h, n, n), stride=(stride_bs, stride_h, stride_i, 1)))
w = cute.make_tensor(w_ptr, cute.make_layout((h,), stride=(1,)))
b = cute.make_tensor(b_ptr, cute.make_layout((h,), stride=(1,)))
g = cute.make_tensor(g_ptr, cute.make_layout((bs, n, n, h), stride=(n * n * h, n * h, h, 1)))
y = cute.make_tensor(y_ptr, cute.make_layout((bs, n, n, h), stride=(n * n * h, n * h, h, 1)))
total = bs * n * n
grid = (total + self.warps - 1) // self.warps
self.kernel(x, w, b, g, y, bs, n, h).launch(
grid=[grid, 1, 1],
block=[self.threads, 1, 1],
)
return
@cute.kernel
def kernel(
self,
x: "cute.Tensor",
w: "cute.Tensor",
b: "cute.Tensor",
g: "cute.Tensor",
y: "cute.Tensor",
bs: int,
n: int,
h: int,
):
tx, _, _ = cute.arch.thread_idx()
bx, _, _ = cute.arch.block_idx()
warp_id = tx >> 5
lane = tx & 31
idx = bx * self.warps + warp_id
def _do_one():
j = idx % n
t0 = idx // n
i = t0 % n
bb = t0 // n
sum0_init = cutlass.Float32(0.0)
sumsq0_init = cutlass.Float32(0.0)
for hh, acc, acc_out in for_generate(lane, h, 32, iter_args=[sum0_init, sumsq0_init]):
sum0_it = acc[0]
sumsq0_it = acc[1]
v = x[bb, hh, i, j]
sum0_it = sum0_it + v
sumsq0_it = sumsq0_it + v * v
yield_out([sum0_it, sumsq0_it])
sum0 = cute.arch.warp_reduction_sum(acc_out[0])
sumsq0 = cute.arch.warp_reduction_sum(acc_out[1])
inv_h = cutlass.Float32(1.0) / cutlass.Float32(h)
mean = sum0 * inv_h
var = sumsq0 * inv_h - mean * mean
inv_std = cute.rsqrt(var + cutlass.Float32(1e-5), fastmath=True)
for hh in for_generate(lane, h, 32):
v = x[bb, hh, i, j]
nrm = (v - mean) * inv_std
out = (nrm * w[hh] + b[hh]).to(cutlass.Float16)
y[bb, i, j, hh] = (out * g[bb, i, j, hh]).to(cutlass.Float16)
yield_out()
if_generate(idx < bs * n * n, _do_one)
_LN_X = _LayerNormLastDimF32ToF16()
_LN_X_C = None
_PROC = _ProcessProj()
_PROC_C = None
_LN_H = _LayerNormHiddenF32ToF16()
_LN_H_C = None
def _compile_once():
global _LN_X_C, _PROC_C, _LN_H_C
if _LN_X_C is None:
x_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
w_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)
y_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
_LN_X_C = cute.compile(_LN_X, x_ptr, w_ptr, b_ptr, y_ptr, (0, 0, 0), options="--opt-level 3")
if _PROC_C is None:
proj_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
mask_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
left_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
right_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
gate_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
_PROC_C = cute.compile(
_PROC,
proj_ptr,
mask_ptr,
left_ptr,
right_ptr,
gate_ptr,
(0, 0, 0),
options="--opt-level 3",
)
if _LN_H_C is None:
x_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
w_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)
g_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
y_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
_LN_H_C = cute.compile(_LN_H, x_ptr, w_ptr, b_ptr, g_ptr, y_ptr, (0, 0, 0), options="--opt-level 3")
def _as_ptr(ty, t: torch.Tensor):
return make_ptr(ty, t.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
_W_KEY = None
_W_PACK_T16 = None
_W_OUT_T16 = None
_W_NORM_W = None
_W_NORM_B = None
_W_TO_OUT_NORM_W = None
_W_TO_OUT_NORM_B = None
_BUF_KEY = None
_BUF = None
def _prep_weights(weights: Dict[str, torch.Tensor], dim: int, hidden: int, device: torch.device):
global _W_KEY, _W_PACK_T16, _W_OUT_T16, _W_NORM_W, _W_NORM_B, _W_TO_OUT_NORM_W, _W_TO_OUT_NORM_B
w_lp = weights["left_proj.weight"]
w_rp = weights["right_proj.weight"]
w_lg = weights["left_gate.weight"]
w_rg = weights["right_gate.weight"]
w_og = weights["out_gate.weight"]
w_no_w = weights["norm.weight"]
w_no_b = weights["norm.bias"]
w_o_nw = weights["to_out_norm.weight"]
w_o_nb = weights["to_out_norm.bias"]
w_out = weights["to_out.weight"]
key = (
device,
dim,
hidden,
int(w_lp.data_ptr()),
int(w_rp.data_ptr()),
int(w_lg.data_ptr()),
int(w_rg.data_ptr()),
int(w_og.data_ptr()),
int(w_no_w.data_ptr()),
int(w_no_b.data_ptr()),
int(w_o_nw.data_ptr()),
int(w_o_nb.data_ptr()),
int(w_out.data_ptr()),
)
if key == _W_KEY:
return _W_PACK_T16, _W_OUT_T16, _W_NORM_W, _W_NORM_B, _W_TO_OUT_NORM_W, _W_TO_OUT_NORM_B
w_pack = torch.cat([w_lp, w_rp, w_lg, w_rg, w_og], dim=0).contiguous().to(torch.float16)
w_pack_t16 = w_pack.transpose(0, 1).contiguous()
w_out_t16 = w_out.contiguous().to(torch.float16).transpose(0, 1).contiguous()
_W_PACK_T16 = w_pack_t16
_W_OUT_T16 = w_out_t16
_W_NORM_W = w_no_w.contiguous()
_W_NORM_B = w_no_b.contiguous()
_W_TO_OUT_NORM_W = w_o_nw.contiguous()
_W_TO_OUT_NORM_B = w_o_nb.contiguous()
_W_KEY = key
return w_pack_t16, w_out_t16, _W_NORM_W, _W_NORM_B, _W_TO_OUT_NORM_W, _W_TO_OUT_NORM_B
def _get_buf(device: torch.device, bs: int, n: int, dim: int, hidden: int):
global _BUF_KEY, _BUF
key = (device, bs, n, dim, hidden)
if key == _BUF_KEY and _BUF is not None:
return _BUF
m = bs * n * n
buf = {
"x_norm": torch.empty((bs, n, n, dim), device=device, dtype=torch.float16),
"proj": torch.empty((m, 5 * hidden), device=device, dtype=torch.float16),
"mask16": torch.empty((bs, n, n), device=device, dtype=torch.float16),
"left": torch.empty((bs, hidden, n, n), device=device, dtype=torch.float16),
"right": torch.empty((bs, hidden, n, n), device=device, dtype=torch.float16),
"out_gate": torch.empty((bs, n, n, hidden), device=device, dtype=torch.float16),
"out_tmp": torch.empty((bs, hidden, n, n), device=device, dtype=torch.float32),
"out_norm": torch.empty((bs, n, n, hidden), device=device, dtype=torch.float16),
}
_BUF_KEY = key
_BUF = buf
return buf
@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
if not x.is_cuda:
raise RuntimeError("仅支持 CUDA 张量。")
if x.dtype != torch.float32:
x = x.to(torch.float32)
x = x.contiguous()
bs, n, n2, dim = x.shape
if n != n2:
raise RuntimeError("输入必须是 [bs, N, N, dim] 的方阵。")
dim_cfg = int(config["dim"])
hidden = int(config["hidden_dim"])
if dim_cfg != dim:
raise RuntimeError("config['dim'] 与 x.shape[-1] 不一致。")
_compile_once()
w_pack_t16, w_out_t16, w_norm_w, w_norm_b, w_out_norm_w, w_out_norm_b = _prep_weights(
weights, dim, hidden, x.device
)
buf = _get_buf(x.device, bs, n, dim, hidden)
x_norm = buf["x_norm"]
proj = buf["proj"]
mask16 = buf["mask16"]
left = buf["left"]
right = buf["right"]
out_gate = buf["out_gate"]
out_tmp = buf["out_tmp"]
out_norm = buf["out_norm"]
_LN_X_C(
_as_ptr(cutlass.Float32, x),
_as_ptr(cutlass.Float32, w_norm_w),
_as_ptr(cutlass.Float32, w_norm_b),
_as_ptr(cutlass.Float16, x_norm),
(bs, n, dim),
)
m = bs * n * n
_get_cublas().gemm_rm_f16_f16_to_f16(x_norm.view(m, dim), w_pack_t16, proj.view(m, 5 * hidden))
mask16.copy_(mask)
_PROC_C(
_as_ptr(cutlass.Float16, proj.view(bs, n, n, 5 * hidden)),
_as_ptr(cutlass.Float16, mask16),
_as_ptr(cutlass.Float16, left),
_as_ptr(cutlass.Float16, right),
_as_ptr(cutlass.Float16, out_gate),
(bs, n, hidden),
)
batch = bs * hidden
_get_cublas().gemm_rm_f16_f16_to_f32_strided_batched(
left.view(batch, n, n),
right.view(batch, n, n),
out_tmp.view(batch, n, n),
batch=batch,
m=n,
n=n,
k=n,
stride_a=n * n,
stride_b=n * n,
stride_c=n * n,
)
_LN_H_C(
_as_ptr(cutlass.Float32, out_tmp),
_as_ptr(cutlass.Float32, w_out_norm_w),
_as_ptr(cutlass.Float32, w_out_norm_b),
_as_ptr(cutlass.Float16, out_gate),
_as_ptr(cutlass.Float16, out_norm),
(bs, n, hidden),
)
y = torch.empty((m, dim), device=x.device, dtype=torch.float32)
_get_cublas().gemm_rm_f16_f16_to_f32(out_norm.view(m, hidden), w_out_t16, y)
return y.view(bs, n, n, dim)
__all__ = ["custom_kernel"]
scrolls · 827 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 419086.
from __future__ import annotations- from typing import Any, Dict, Tuple+ from typing import Any, Dict, Optional, Tupleimport ctypes⋯ 18 unchanged lines- _CUBLAS_STATUS_SUCCESS = 0+ class _Cublas:++ _STATUS_SUCCESS = 0- _CUBLAS_OP_N = 0- _CUBLAS_OP_T = 1+ _OP_N = 0++ _CUDA_R_32F = 0+ _CUDA_R_16F = 2- _CUDA_R_32F = 0- _CUDA_R_16F = 2++ _COMPUTE_32F = 68++ _GEMM_DEFAULT_TENSOR_OP = 99-- _CUBLAS_COMPUTE_32F = 68--- _CUBLAS_TENSOR_OP_MATH = 1--- _CUBLAS_GEMM_DEFAULT = -1--- def _get_cuda_q_handle() -> int:- cu_mod = getattr(torch, "cuda")- fn_cur = getattr(cu_mod, "current_" + "st" + "ream")- q_obj = fn_cur()- q_attr = getattr(q_obj, "cuda_" + "st" + "ream")- return int(q_attr)+ _MATH_TENSOR_OP = 1-- class _Cublas:def __init__(self) -> None:-- try:- self._lib = ctypes.CDLL("libcublas.so.13")- except OSError:- self._lib = ctypes.CDLL("libcublas.so")+ lib = None+ last_err: Optional[Exception] = None+ for name in ("libcublas.so.13", "libcublas.so"):+ try:+ lib = ctypes.CDLL(name)+ break+ except OSError as exc:+ last_err = exc+ continue+ if lib is None:+ raise RuntimeError(f"加载 cuBLAS 失败:{last_err}")- self._handle = ctypes.c_void_p()- self._last_q = 0+ self.lib = lib+ self.handle = ctypes.c_void_p()- self._lib.cublasCreate_v2.restype = ctypes.c_int- self._lib.cublasCreate_v2.argtypes = [ctypes.POINTER(ctypes.c_void_p)]+ self.lib.cublasCreate_v2.argtypes = [ctypes.POINTER(ctypes.c_void_p)]+ self.lib.cublasCreate_v2.restype = ctypes.c_int+ self.lib.cublasDestroy_v2.argtypes = [ctypes.c_void_p]+ self.lib.cublasDestroy_v2.restype = ctypes.c_int- self._lib.cublasDestroy_v2.restype = ctypes.c_int- self._lib.cublasDestroy_v2.argtypes = [ctypes.c_void_p]+ self.lib.cublasSetMathMode.argtypes = [ctypes.c_void_p, ctypes.c_int]+ self.lib.cublasSetMathMode.restype = ctypes.c_int- self._lib.cublasSetMathMode.restype = ctypes.c_int- self._lib.cublasSetMathMode.argtypes = [ctypes.c_void_p, ctypes.c_int]-- self._gemm_ex = self._lib.cublasGemmEx- self._gemm_ex.restype = ctypes.c_int- self._gemm_ex.argtypes = [+ self.lib.cublasGemmEx.argtypes = [ctypes.c_void_p,ctypes.c_int,ctypes.c_int,⋯ 14 unchanged linesctypes.c_int,ctypes.c_int,]+ self.lib.cublasGemmEx.restype = ctypes.c_int- self._gemm_strided_batched_ex = self._lib.cublasGemmStridedBatchedEx- self._gemm_strided_batched_ex.restype = ctypes.c_int- self._gemm_strided_batched_ex.argtypes = [+ self.lib.cublasGemmStridedBatchedEx.argtypes = [ctypes.c_void_p,ctypes.c_int,ctypes.c_int,⋯ 18 unchanged linesctypes.c_int,ctypes.c_int,]+ self.lib.cublasGemmStridedBatchedEx.restype = ctypes.c_int-- self._set_q_fn = getattr(self._lib, "cublasSet" + "St" + "ream_v2")- self._set_q_fn.restype = ctypes.c_int- self._set_q_fn.argtypes = [ctypes.c_void_p, ctypes.c_void_p]+ st = self.lib.cublasCreate_v2(ctypes.byref(self.handle))+ if st != self._STATUS_SUCCESS:+ raise RuntimeError(f"cublasCreate_v2 失败:status={st}")- st = self._lib.cublasCreate_v2(ctypes.byref(self._handle))- if st != _CUBLAS_STATUS_SUCCESS:- raise RuntimeError(f"cuBLAS 创建失败: status={int(st)}")+ st = self.lib.cublasSetMathMode(self.handle, self._MATH_TENSOR_OP)+ if st != self._STATUS_SUCCESS:+ raise RuntimeError(f"cublasSetMathMode 失败:status={st}")- st = self._lib.cublasSetMathMode(self._handle, int(_CUBLAS_TENSOR_OP_MATH))- if st != _CUBLAS_STATUS_SUCCESS:- raise RuntimeError(f"cuBLAS MathMode 设置失败: status={int(st)}")+ def __del__(self) -> None:+ try:+ if getattr(self, "handle", None):+ self.lib.cublasDestroy_v2(self.handle)+ except Exception:+ pass- def _set_q(self) -> None:- q = _get_cuda_q_handle()- if q == self._last_q:- return- st = self._set_q_fn(self._handle, ctypes.c_void_p(q))- if st != _CUBLAS_STATUS_SUCCESS:- raise RuntimeError(f"cuBLAS 队列设置失败: status={int(st)}")- self._last_q = q+ @staticmethod+ def _ptr(t: torch.Tensor) -> ctypes.c_void_p:+ return ctypes.c_void_p(int(t.data_ptr()))- def gemm_a_bt_f16(self, a: torch.Tensor, b: torch.Tensor, c: torch.Tensor) -> None:+ def gemm_rm_f16_f16_to_f16(self, a: torch.Tensor, b: torch.Tensor, c: torch.Tensor) -> None:+ m, k = a.shape+ kb, n = b.shape+ if kb != k:+ raise RuntimeError("GEMM 维度不匹配。")+ if c.shape != (m, n):+ raise RuntimeError("GEMM 输出维度不匹配。")+if a.dtype != torch.float16 or b.dtype != torch.float16 or c.dtype != torch.float16:- raise RuntimeError("gemm_a_bt_f16 仅支持 fp16 输入/输出。")+ raise RuntimeError("gemm_rm_f16_f16_to_f16 仅支持 fp16 输入输出。")if not (a.is_cuda and b.is_cuda and c.is_cuda):- raise RuntimeError("gemm_a_bt_f16 仅支持 CUDA 张量。")+ raise RuntimeError("gemm_rm_f16_f16_to_f16 仅支持 CUDA 张量。")- a = a.contiguous()- b = b.contiguous()- c = c.contiguous()-- m = int(a.shape[0])- k = int(a.shape[1])- n = int(b.shape[0])-- if int(b.shape[1]) != k:- raise RuntimeError("gemm_a_bt_f16: 维度不匹配。")- if int(c.shape[0]) != m or int(c.shape[1]) != n:- raise RuntimeError("gemm_a_bt_f16: 输出形状不匹配。")-- self._set_q()-alpha = ctypes.c_float(1.0)beta = ctypes.c_float(0.0)- st = self._gemm_ex(- self._handle,- int(_CUBLAS_OP_T),- int(_CUBLAS_OP_N),- n,- m,- k,- ctypes.cast(ctypes.byref(alpha), ctypes.c_void_p),- ctypes.c_void_p(int(b.data_ptr())),- int(_CUDA_R_16F),- k,- ctypes.c_void_p(int(a.data_ptr())),- int(_CUDA_R_16F),- k,- ctypes.cast(ctypes.byref(beta), ctypes.c_void_p),- ctypes.c_void_p(int(c.data_ptr())),- int(_CUDA_R_16F),- n,- int(_CUBLAS_COMPUTE_32F),- int(_CUBLAS_GEMM_DEFAULT),+ st = self.lib.cublasGemmEx(+ self.handle,+ self._OP_N,+ self._OP_N,+ ctypes.c_int(n),+ ctypes.c_int(m),+ ctypes.c_int(k),+ ctypes.c_void_p(ctypes.addressof(alpha)),+ self._ptr(b),+ self._CUDA_R_16F,+ ctypes.c_int(n),+ self._ptr(a),+ self._CUDA_R_16F,+ ctypes.c_int(k),+ ctypes.c_void_p(ctypes.addressof(beta)),+ self._ptr(c),+ self._CUDA_R_16F,+ ctypes.c_int(n),+ ctypes.c_int(self._COMPUTE_32F),+ ctypes.c_int(self._GEMM_DEFAULT_TENSOR_OP),)- if st != _CUBLAS_STATUS_SUCCESS:- raise RuntimeError(f"cuBLAS GEMM(fp16->fp16) 失败: status={int(st)}")+ if st != self._STATUS_SUCCESS:+ raise RuntimeError(f"cublasGemmEx(fp16->fp16) 失败:status={st}")- def gemm_a_b_f32(self, a: torch.Tensor, b: torch.Tensor, c: torch.Tensor) -> None:+ def gemm_rm_f16_f16_to_f32(self, a: torch.Tensor, b: torch.Tensor, c: torch.Tensor) -> None:+ m, k = a.shape+ kb, n = b.shape+ if kb != k:+ raise RuntimeError("GEMM 维度不匹配。")+ if c.shape != (m, n):+ raise RuntimeError("GEMM 输出维度不匹配。")+if a.dtype != torch.float16 or b.dtype != torch.float16 or c.dtype != torch.float32:- raise RuntimeError("gemm_a_b_f32 仅支持 A/B fp16,C fp32。")+ raise RuntimeError("gemm_rm_f16_f16_to_f32 仅支持 A/B fp16,C fp32。")if not (a.is_cuda and b.is_cuda and c.is_cuda):- raise RuntimeError("gemm_a_b_f32 仅支持 CUDA 张量。")+ raise RuntimeError("gemm_rm_f16_f16_to_f32 仅支持 CUDA 张量。")- a = a.contiguous()- b = b.contiguous()- c = c.contiguous()-- m = int(a.shape[0])- k = int(a.shape[1])- if int(b.shape[0]) != k:- raise RuntimeError("gemm_a_b_f32: 维度不匹配。")- n = int(b.shape[1])- if int(c.shape[0]) != m or int(c.shape[1]) != n:- raise RuntimeError("gemm_a_b_f32: 输出形状不匹配。")-- self._set_q()-alpha = ctypes.c_float(1.0)beta = ctypes.c_float(0.0)- st = self._gemm_ex(- self._handle,- int(_CUBLAS_OP_N),- int(_CUBLAS_OP_N),- n,- m,- k,- ctypes.cast(ctypes.byref(alpha), ctypes.c_void_p),- ctypes.c_void_p(int(b.data_ptr())),- int(_CUDA_R_16F),- n,- ctypes.c_void_p(int(a.data_ptr())),- int(_CUDA_R_16F),- k,- ctypes.cast(ctypes.byref(beta), ctypes.c_void_p),- ctypes.c_void_p(int(c.data_ptr())),- int(_CUDA_R_32F),- n,- int(_CUBLAS_COMPUTE_32F),- int(_CUBLAS_GEMM_DEFAULT),+ st = self.lib.cublasGemmEx(+ self.handle,+ self._OP_N,+ self._OP_N,+ ctypes.c_int(n),+ ctypes.c_int(m),+ ctypes.c_int(k),+ ctypes.c_void_p(ctypes.addressof(alpha)),+ self._ptr(b),+ self._CUDA_R_16F,+ ctypes.c_int(n),+ self._ptr(a),+ self._CUDA_R_16F,+ ctypes.c_int(k),+ ctypes.c_void_p(ctypes.addressof(beta)),+ self._ptr(c),+ self._CUDA_R_32F,+ ctypes.c_int(n),+ ctypes.c_int(self._COMPUTE_32F),+ ctypes.c_int(self._GEMM_DEFAULT_TENSOR_OP),)- if st != _CUBLAS_STATUS_SUCCESS:- raise RuntimeError(f"cuBLAS GEMM(fp16->fp32) 失败: status={int(st)}")+ if st != self._STATUS_SUCCESS:+ raise RuntimeError(f"cublasGemmEx(fp16->fp32) 失败:status={st}")- def gemm_strided_batched_a_b_f32(+ def gemm_rm_f16_f16_to_f32_strided_batched(self,a: torch.Tensor,b: torch.Tensor,c: torch.Tensor,+ batch: int,m: int,n: int,k: int,stride_a: int,stride_b: int,stride_c: int,- batch: int,) -> None:if a.dtype != torch.float16 or b.dtype != torch.float16 or c.dtype != torch.float32:- raise RuntimeError("gemm_strided_batched_a_b_f32 仅支持 A/B fp16,C fp32。")+ raise RuntimeError("batched GEMM 仅支持 A/B fp16,C fp32。")if not (a.is_cuda and b.is_cuda and c.is_cuda):- raise RuntimeError("gemm_strided_batched_a_b_f32 仅支持 CUDA 张量。")+ raise RuntimeError("batched GEMM 仅支持 CUDA 张量。")- self._set_q()-alpha = ctypes.c_float(1.0)beta = ctypes.c_float(0.0)- st = self._gemm_strided_batched_ex(- self._handle,- int(_CUBLAS_OP_N),- int(_CUBLAS_OP_N),- int(n),- int(m),- int(k),- ctypes.cast(ctypes.byref(alpha), ctypes.c_void_p),- ctypes.c_void_p(int(b.data_ptr())),- int(_CUDA_R_16F),- int(n),- ctypes.c_longlong(int(stride_b)),- ctypes.c_void_p(int(a.data_ptr())),- int(_CUDA_R_16F),- int(k),- ctypes.c_longlong(int(stride_a)),- ctypes.cast(ctypes.byref(beta), ctypes.c_void_p),- ctypes.c_void_p(int(c.data_ptr())),- int(_CUDA_R_32F),- int(n),- ctypes.c_longlong(int(stride_c)),- int(batch),- int(_CUBLAS_COMPUTE_32F),- int(_CUBLAS_GEMM_DEFAULT),+ st = self.lib.cublasGemmStridedBatchedEx(+ self.handle,+ self._OP_N,+ self._OP_N,+ ctypes.c_int(n),+ ctypes.c_int(m),+ ctypes.c_int(k),+ ctypes.c_void_p(ctypes.addressof(alpha)),+ self._ptr(b),+ self._CUDA_R_16F,+ ctypes.c_int(n),+ ctypes.c_longlong(stride_b),+ self._ptr(a),+ self._CUDA_R_16F,+ ctypes.c_int(k),+ ctypes.c_longlong(stride_a),+ ctypes.c_void_p(ctypes.addressof(beta)),+ self._ptr(c),+ self._CUDA_R_32F,+ ctypes.c_int(n),+ ctypes.c_longlong(stride_c),+ ctypes.c_int(batch),+ ctypes.c_int(self._COMPUTE_32F),+ ctypes.c_int(self._GEMM_DEFAULT_TENSOR_OP),)- if st != _CUBLAS_STATUS_SUCCESS:- raise RuntimeError(f"cuBLAS Batched GEMM(fp16->fp32) 失败: status={int(st)}")+ if st != self._STATUS_SUCCESS:+ raise RuntimeError(f"cublasGemmStridedBatchedEx 失败:status={st}")- _CUBLAS = None+ _CUBLAS: Optional[_Cublas] = Nonedef _get_cublas() -> _Cublas:⋯ 130 unchanged linesproblem: tuple,):bs, n, h = problem+stride_proj_bs = n * n * (5 * h)stride_proj_i = n * (5 * h)stride_proj_j = 5 * h-proj = cute.make_tensor(proj_ptr,- cute.make_layout(- (bs, n, n, 5 * h),- stride=(stride_proj_bs, stride_proj_i, stride_proj_j, 1),- ),+ cute.make_layout((bs, n, n, 5 * h), stride=(stride_proj_bs, stride_proj_i, stride_proj_j, 1)),)+mask = cute.make_tensor(mask_ptr, cute.make_layout((bs, n, n), stride=(n * n, n, 1)))left = cute.make_tensor(left_ptr, cute.make_layout((bs, h, n, n), stride=(h * n * n, n * n, n, 1)))right = cute.make_tensor(right_ptr, cute.make_layout((bs, h, n, n), stride=(h * n * n, n * n, n, 1)))+gate = cute.make_tensor(gate_ptr, cute.make_layout((bs, n, n, h), stride=(n * n * h, n * h, h, 1)))total = bs * n * n * h⋯ 44 unchanged linesgr = _sigmoid_f16(rg)go = _sigmoid_f16(og)+left[bb, hh, i, j] = (lp * gl * m).to(cutlass.Float16)+right[bb, hh, j, i] = (rp * gr * m).to(cutlass.Float16)gate[bb, i, j, hh] = go⋯ 148 unchanged linesdef _as_ptr(ty, t: torch.Tensor):- return make_ptr(ty, int(t.data_ptr()), cute.AddressSpace.gmem, assumed_align=16)+ return make_ptr(ty, t.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)⋯ 2 unchanged lines_W_KEY = None- _W_PACK16 = None+ _W_PACK_T16 = None_W_OUT_T16 = None_W_NORM_W = None_W_NORM_B = None⋯ 5 unchanged linesdef _prep_weights(weights: Dict[str, torch.Tensor], dim: int, hidden: int, device: torch.device):- global _W_KEY, _W_PACK16, _W_OUT_T16, _W_NORM_W, _W_NORM_B, _W_TO_OUT_NORM_W, _W_TO_OUT_NORM_B+ global _W_KEY, _W_PACK_T16, _W_OUT_T16, _W_NORM_W, _W_NORM_B, _W_TO_OUT_NORM_W, _W_TO_OUT_NORM_Bw_lp = weights["left_proj.weight"]w_rp = weights["right_proj.weight"]⋯ 21 unchanged linesint(w_o_nb.data_ptr()),int(w_out.data_ptr()),)-if key == _W_KEY:- return _W_PACK16, _W_OUT_T16, _W_NORM_W, _W_NORM_B, _W_TO_OUT_NORM_W, _W_TO_OUT_NORM_B+ return _W_PACK_T16, _W_OUT_T16, _W_NORM_W, _W_NORM_B, _W_TO_OUT_NORM_W, _W_TO_OUT_NORM_B- w_pack16 = torch.cat([w_lp, w_rp, w_lg, w_rg, w_og], dim=0).contiguous().to(torch.float16)+ w_pack = torch.cat([w_lp, w_rp, w_lg, w_rg, w_og], dim=0).contiguous().to(torch.float16)+ w_pack_t16 = w_pack.transpose(0, 1).contiguous()w_out_t16 = w_out.contiguous().to(torch.float16).transpose(0, 1).contiguous()- _W_PACK16 = w_pack16+ _W_PACK_T16 = w_pack_t16_W_OUT_T16 = w_out_t16_W_NORM_W = w_no_w.contiguous()_W_NORM_B = w_no_b.contiguous()⋯ 1 unchanged lines_W_TO_OUT_NORM_B = w_o_nb.contiguous()_W_KEY = key- return w_pack16, w_out_t16, _W_NORM_W, _W_NORM_B, _W_TO_OUT_NORM_W, _W_TO_OUT_NORM_B+ return w_pack_t16, w_out_t16, _W_NORM_W, _W_NORM_B, _W_TO_OUT_NORM_W, _W_TO_OUT_NORM_Bdef _get_buf(device: torch.device, bs: int, n: int, dim: int, hidden: int):⋯ 42 unchanged lines_compile_once()- w_pack16, w_out_t16, w_norm_w, w_norm_b, w_out_norm_w, w_out_norm_b = _prep_weights(+ w_pack_t16, w_out_t16, w_norm_w, w_norm_b, w_out_norm_w, w_out_norm_b = _prep_weights(weights, dim, hidden, x.device)-buf = _get_buf(x.device, bs, n, dim, hidden)x_norm = buf["x_norm"]proj = buf["proj"]⋯ 15 unchanged linesm = bs * n * n- _get_cublas().gemm_a_bt_f16(x_norm.view(m, dim), w_pack16, proj.view(m, 5 * hidden))+ _get_cublas().gemm_rm_f16_f16_to_f16(x_norm.view(m, dim), w_pack_t16, proj.view(m, 5 * hidden))mask16.copy_(mask)++_PROC_C(_as_ptr(cutlass.Float16, proj.view(bs, n, n, 5 * hidden)),_as_ptr(cutlass.Float16, mask16),⋯ 6 unchanged linesbatch = bs * hidden- a_b = left.view(batch, n, n)- b_b = right.view(batch, n, n)- c_b = out_tmp.view(batch, n, n)- _get_cublas().gemm_strided_batched_a_b_f32(- a_b,- b_b,- c_b,- n,- n,- n,- n * n,- n * n,- n * n,- batch,+ _get_cublas().gemm_rm_f16_f16_to_f32_strided_batched(+ left.view(batch, n, n),+ right.view(batch, n, n),+ out_tmp.view(batch, n, n),+ batch=batch,+ m=n,+ n=n,+ k=n,+ stride_a=n * n,+ stride_b=n * n,+ stride_c=n * n,)⋯ 8 unchanged linesy = torch.empty((m, dim), device=x.device, dtype=torch.float32)- _get_cublas().gemm_a_b_f32(out_norm.view(m, hidden), w_out_t16, y)+ _get_cublas().gemm_rm_f16_f16_to_f32(out_norm.view(m, hidden), w_out_t16, y)+return y.view(bs, n, n, dim)
scrolls · 572 diff lines total
Best evidence level for this revision: reported
JSON