submission 417301
shiyegao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 250 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-417301?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:823e51b4cf2fe1553d4ae257152f8c7465e88a85175f43da99a20bac59fb0d9e
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
smem_stride = _TILE_K + 1tile-k = 32
_TILE_K = 32Kernel source
submission.py250 lines
from __future__ import annotations
from typing import Any, Dict, Tuple
import torch
import torch.nn.functional as F
import cutlass
import cutlass.cute as cute
from cutlass.cute.runtime import make_ptr
from cutlass.cutlass_dsl import for_generate, yield_out
_TILE_MN = 32
_TILE_K = 32
_THREADS = 256
class _OutgoingContractBatchedKernel:
def __init__(self) -> None:
self.threads = _THREADS
@cute.jit
def __call__(
self,
a_ptr: "cute.Pointer",
b_ptr: "cute.Pointer",
c_ptr: "cute.Pointer",
problem: tuple,
):
bh, n = problem
stride_bh = n * n
stride_m = n
stride_n = 1
a = cute.make_tensor(
a_ptr,
cute.make_layout((bh, n, n), stride=(stride_bh, stride_m, stride_n)),
)
b = cute.make_tensor(
b_ptr,
cute.make_layout((bh, n, n), stride=(stride_bh, stride_m, stride_n)),
)
c = cute.make_tensor(
c_ptr,
cute.make_layout((bh, n, n), stride=(stride_bh, stride_m, stride_n)),
)
grid_x = n // _TILE_MN
grid_y = n // _TILE_MN
grid_z = bh
self.kernel(a, b, c, n).launch(
grid=[grid_x, grid_y, grid_z],
block=[self.threads, 1, 1],
)
return
@cute.kernel
def kernel(
self,
a: "cute.Tensor",
b: "cute.Tensor",
c: "cute.Tensor",
n: int,
):
tx, _, _ = cute.arch.thread_idx()
bx, by, bz = cute.arch.block_idx()
tid = tx
lane_m = tid >> 4
lane_n = tid & 15
base_m = by * _TILE_MN
base_n = bx * _TILE_MN
i0 = base_m + lane_m
j0 = base_n + lane_n
i1 = i0 + 16
j1 = j0 + 16
acc00 = cutlass.Float32(0.0)
acc01 = cutlass.Float32(0.0)
acc10 = cutlass.Float32(0.0)
acc11 = cutlass.Float32(0.0)
smem_stride = _TILE_K + 1
smem_a_elems = _TILE_MN * smem_stride
smem_b_elems = _TILE_MN * smem_stride
smem_ptr = cute.arch.alloc_smem(cutlass.Float32, smem_a_elems + smem_b_elems)
sA = cute.make_tensor(
smem_ptr,
cute.make_layout((_TILE_MN, _TILE_K), stride=(smem_stride, 1)),
)
sB = cute.make_tensor(
smem_ptr + smem_a_elems,
cute.make_layout((_TILE_MN, _TILE_K), stride=(smem_stride, 1)),
)
for k0, (acc00, acc01, acc10, acc11), (acc00_out, acc01_out, acc10_out, acc11_out) in for_generate(
0,
n,
_TILE_K,
iter_args=[acc00, acc01, acc10, acc11],
):
base = tid << 2
for t in range(4):
idx = base + t
row = idx >> 5
col = idx & 31
sA[row, col] = a[bz, base_m + row, k0 + col]
sB[row, col] = b[bz, base_n + row, k0 + col]
cute.arch.sync_threads()
for kk in range(_TILE_K):
a0 = sA[lane_m, kk]
a1 = sA[lane_m + 16, kk]
b0 = sB[lane_n, kk]
b1 = sB[lane_n + 16, kk]
acc00 = acc00 + a0 * b0
acc01 = acc01 + a0 * b1
acc10 = acc10 + a1 * b0
acc11 = acc11 + a1 * b1
cute.arch.sync_threads()
yield_out([acc00, acc01, acc10, acc11])
c[bz, i0, j0] = acc00_out
c[bz, i0, j1] = acc01_out
c[bz, i1, j0] = acc10_out
c[bz, i1, j1] = acc11_out
_CONTRACT = _OutgoingContractBatchedKernel()
_CONTRACT_COMPILED = None
def _get_contract_compiled():
global _CONTRACT_COMPILED
if _CONTRACT_COMPILED is not None:
return _CONTRACT_COMPILED
a_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
c_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
_CONTRACT_COMPILED = cute.compile(
_CONTRACT,
a_ptr,
b_ptr,
c_ptr,
(0, 0),
options="--opt-level 3",
)
return _CONTRACT_COMPILED
def _contract_outgoing(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:
bs, n, _, hidden = left.shape
if not (left.is_cuda and right.is_cuda):
raise RuntimeError("该实现仅支持 CUDA 张量。")
if left.dtype != torch.float32 or right.dtype != torch.float32:
raise RuntimeError("该实现期望 left/right 为 float32。")
if hidden != 128:
raise RuntimeError("仅支持 hidden_dim=128 的特化路径。")
if (n & (_TILE_MN - 1)) != 0:
raise RuntimeError("仅支持 N 为 32 的倍数的特化路径。")
left_t = left.permute(0, 3, 1, 2).contiguous()
right_t = right.permute(0, 3, 1, 2).contiguous()
bh = bs * hidden
a = left_t.view(bh, n, n)
b = right_t.view(bh, n, n)
c = torch.empty((bh, n, n), device=left.device, dtype=torch.float32)
compiled = _get_contract_compiled()
a_ptr = make_ptr(cutlass.Float32, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(cutlass.Float32, b.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
c_ptr = make_ptr(cutlass.Float32, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
compiled(a_ptr, b_ptr, c_ptr, (bh, n))
out = c.view(bs, hidden, n, n).permute(0, 2, 3, 1).contiguous()
return out
@torch.inference_mode()
def custom_kernel(data: Tuple[torch.Tensor, torch.Tensor, Dict[str, torch.Tensor], Dict[str, Any]]) -> torch.Tensor:
x, mask, weights, config = data
dim = int(config["dim"])
hidden_dim = int(config["hidden_dim"])
if x.dtype != torch.float32:
x = x.to(dtype=torch.float32)
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
x = F.layer_norm(x, (dim,), weights["norm.weight"], weights["norm.bias"], 1e-5)
left = F.linear(x, weights["left_proj.weight"], None)
right = F.linear(x, weights["right_proj.weight"], None)
mask_f = mask.unsqueeze(-1)
if mask_f.dtype != left.dtype:
mask_f = mask_f.to(dtype=left.dtype)
left = left * mask_f
right = right * mask_f
left_gate = torch.sigmoid(F.linear(x, weights["left_gate.weight"], None))
right_gate = torch.sigmoid(F.linear(x, weights["right_gate.weight"], None))
out_gate = torch.sigmoid(F.linear(x, weights["out_gate.weight"], None))
left = left * left_gate
right = right * right_gate
out = _contract_outgoing(left, right)
out = F.layer_norm(
out,
(hidden_dim,),
weights["to_out_norm.weight"],
weights["to_out_norm.bias"],
1e-5,
)
out = out * out_gate
out = F.linear(out, weights["to_out.weight"], None)
return out
__all__ = ["custom_kernel"]
scrolls · 250 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 417295.
⋯ 1 unchanged linesfrom typing import Any, Dict, Tuple+ import torch+ import torch.nn.functional as F- _EXT_MOD = None+ import cutlass+ import cutlass.cute as cute+ from cutlass.cute.runtime import make_ptr+ from cutlass.cutlass_dsl import for_generate, yield_out- def _load_ext():- global _EXT_MOD- if _EXT_MOD is not None:- return _EXT_MOD- import inspect- import torch- from torch.utils.cpp_extension import load_inline-- cpp = (- r"""- #include <torch/extension.h>- #include <ATen/cuda/CUDABlas.h>- #include <cublas_v2.h>- static inline void _check(bool ok, const char* msg) {- if (!ok) {- throw std::runtime_error(msg);- }- }- static inline void _cublas_ok(cublasStatus_t st, const char* msg) {- if (st != CUBLAS_STATUS_SUCCESS) {- throw std::runtime_error(msg);- }- }- torch::Tensor contract_bf16(torch::Tensor a, torch::Tensor b) {- _check(a.is_cuda(), "a must be CUDA tensor");- _check(b.is_cuda(), "b must be CUDA tensor");- _check(a.is_contiguous(), "a must be contiguous");- _check(b.is_contiguous(), "b must be contiguous");- _check(a.scalar_type() == at::kBFloat16, "a must be bfloat16");- _check(b.scalar_type() == at::kBFloat16, "b must be bfloat16");- _check(a.dim() == 3, "a must be [batch, n, n]");- _check(b.dim() == 3, "b must be [batch, n, n]");- _check(a.sizes() == b.sizes(), "a and b must have same shape");+ _TILE_MN = 32+ _TILE_K = 32+ _THREADS = 256- const auto batch = (int)a.size(0);- const auto n0 = (int)a.size(1);- const auto n1 = (int)a.size(2);- _check(n0 == n1, "a must be square");- _check(n0 > 0, "n must be > 0");- auto out = torch::empty({batch, n0, n0}, a.options().dtype(at::kFloat));+ class _OutgoingContractBatchedKernel:+ def __init__(self) -> None:+ self.threads = _THREADS- const float alpha = 1.0f;- const float beta = 0.0f;+ @cute.jit+ def __call__(+ self,+ a_ptr: "cute.Pointer",+ b_ptr: "cute.Pointer",+ c_ptr: "cute.Pointer",+ problem: tuple,+ ):+ bh, n = problem- // 目标:C_rm = A_rm @ B_rm^T- // 通过列主序计算 C_cm = B_rm @ A_rm^T 写回到 C_rm 的内存(等价的转置视图技巧)。- const int m = n0;- const int n = n0;- const int k = n0;- const int lda = k;- const int ldb = k;- const int ldc = m;- const long long stride_a = (long long)k * (long long)m;- const long long stride_b = (long long)k * (long long)n;- const long long stride_c = (long long)m * (long long)n;+ stride_bh = n * n+ stride_m = n+ stride_n = 1- cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();- _cublas_ok(cublasSetMathMode(handle, CUBLAS_TENSOR_OP_MATH), "cublasSetMathMode failed");+ a = cute.make_tensor(+ a_ptr,+ cute.make_layout((bh, n, n), stride=(stride_bh, stride_m, stride_n)),+ )+ b = cute.make_tensor(+ b_ptr,+ cute.make_layout((bh, n, n), stride=(stride_bh, stride_m, stride_n)),+ )+ c = cute.make_tensor(+ c_ptr,+ cute.make_layout((bh, n, n), stride=(stride_bh, stride_m, stride_n)),+ )- _cublas_ok(- cublasGemmStridedBatchedEx(- handle,- CUBLAS_OP_T,- CUBLAS_OP_N,- m,- n,- k,- &alpha,- /*A*/ b.data_ptr(),- CUDA_R_16BF,- lda,- stride_a,- /*B*/ a.data_ptr(),- CUDA_R_16BF,- ldb,- stride_b,- &beta,- /*C*/ out.data_ptr(),- CUDA_R_32F,- ldc,- stride_c,- batch,- CUBLAS_COMPUTE_32F,- CUBLAS_GEMM_DEFAULT_TENSOR_OP),- "cublasGemmStridedBatchedEx failed");+ grid_x = n // _TILE_MN+ grid_y = n // _TILE_MN+ grid_z = bh+ self.kernel(a, b, c, n).launch(+ grid=[grid_x, grid_y, grid_z],+ block=[self.threads, 1, 1],+ )+ return- return out;- }+ @cute.kernel+ def kernel(+ self,+ a: "cute.Tensor",+ b: "cute.Tensor",+ c: "cute.Tensor",+ n: int,+ ):+ tx, _, _ = cute.arch.thread_idx()+ bx, by, bz = cute.arch.block_idx()- PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {- m.def("contract_bf16", &contract_bf16, "contract_bf16");- }- """- )+ tid = tx+ lane_m = tid >> 4+ lane_n = tid & 15-- kwargs = {"with_cuda": True, "verbose": False}- sig = inspect.signature(load_inline)- if "extra_cflags" in sig.parameters:- kwargs["extra_cflags"] = ["-O3"]- if "extra_ldflags" in sig.parameters:- kwargs["extra_ldflags"] = ["-lcublas"]- if "functions" in sig.parameters:- kwargs["functions"] = None+ base_m = by * _TILE_MN+ base_n = bx * _TILE_MN- _EXT_MOD = load_inline(- name="trimul_cublas_ext",- cpp_sources=cpp,- **kwargs,+ i0 = base_m + lane_m+ j0 = base_n + lane_n+ i1 = i0 + 16+ j1 = j0 + 16++ acc00 = cutlass.Float32(0.0)+ acc01 = cutlass.Float32(0.0)+ acc10 = cutlass.Float32(0.0)+ acc11 = cutlass.Float32(0.0)+++ smem_stride = _TILE_K + 1+ smem_a_elems = _TILE_MN * smem_stride+ smem_b_elems = _TILE_MN * smem_stride+ smem_ptr = cute.arch.alloc_smem(cutlass.Float32, smem_a_elems + smem_b_elems)++ sA = cute.make_tensor(+ smem_ptr,+ cute.make_layout((_TILE_MN, _TILE_K), stride=(smem_stride, 1)),+ )+ sB = cute.make_tensor(+ smem_ptr + smem_a_elems,+ cute.make_layout((_TILE_MN, _TILE_K), stride=(smem_stride, 1)),+ )++ for k0, (acc00, acc01, acc10, acc11), (acc00_out, acc01_out, acc10_out, acc11_out) in for_generate(+ 0,+ n,+ _TILE_K,+ iter_args=[acc00, acc01, acc10, acc11],+ ):++ base = tid << 2+ for t in range(4):+ idx = base + t+ row = idx >> 5+ col = idx & 31+ sA[row, col] = a[bz, base_m + row, k0 + col]+ sB[row, col] = b[bz, base_n + row, k0 + col]++ cute.arch.sync_threads()+++ for kk in range(_TILE_K):+ a0 = sA[lane_m, kk]+ a1 = sA[lane_m + 16, kk]+ b0 = sB[lane_n, kk]+ b1 = sB[lane_n + 16, kk]++ acc00 = acc00 + a0 * b0+ acc01 = acc01 + a0 * b1+ acc10 = acc10 + a1 * b0+ acc11 = acc11 + a1 * b1++ cute.arch.sync_threads()+ yield_out([acc00, acc01, acc10, acc11])++ c[bz, i0, j0] = acc00_out+ c[bz, i0, j1] = acc01_out+ c[bz, i1, j0] = acc10_out+ c[bz, i1, j1] = acc11_out+++ _CONTRACT = _OutgoingContractBatchedKernel()+ _CONTRACT_COMPILED = None+++ def _get_contract_compiled():+ global _CONTRACT_COMPILED+ if _CONTRACT_COMPILED is not None:+ return _CONTRACT_COMPILED++ a_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)+ b_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)+ c_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)+ _CONTRACT_COMPILED = cute.compile(+ _CONTRACT,+ a_ptr,+ b_ptr,+ c_ptr,+ (0, 0),+ options="--opt-level 3",)- return _EXT_MOD+ return _CONTRACT_COMPILED- def _contract_outgoing(left: "torch.Tensor", right: "torch.Tensor") -> "torch.Tensor":- import torch+ def _contract_outgoing(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:+ bs, n, _, hidden = left.shapeif not (left.is_cuda and right.is_cuda):- raise RuntimeError("CUDA only")- if left.dtype != torch.bfloat16:- left = left.to(dtype=torch.bfloat16)- if right.dtype != torch.bfloat16:- right = right.to(dtype=torch.bfloat16)+ raise RuntimeError("该实现仅支持 CUDA 张量。")+ if left.dtype != torch.float32 or right.dtype != torch.float32:+ raise RuntimeError("该实现期望 left/right 为 float32。")- bs, n, _, h = left.shape- left_m = left.permute(0, 3, 1, 2).contiguous().view(bs * h, n, n)- right_m = right.permute(0, 3, 1, 2).contiguous().view(bs * h, n, n)+ if hidden != 128:+ raise RuntimeError("仅支持 hidden_dim=128 的特化路径。")- ext = _load_ext()- out_m = ext.contract_bf16(left_m, right_m)- out = out_m.view(bs, h, n, n).permute(0, 2, 3, 1).contiguous()- return out+ if (n & (_TILE_MN - 1)) != 0:+ raise RuntimeError("仅支持 N 为 32 的倍数的特化路径。")++ left_t = left.permute(0, 3, 1, 2).contiguous()+ right_t = right.permute(0, 3, 1, 2).contiguous()+ bh = bs * hidden+ a = left_t.view(bh, n, n)+ b = right_t.view(bh, n, n)- def custom_kernel(- data: Tuple["torch.Tensor", "torch.Tensor", Dict[str, "torch.Tensor"], Dict[str, Any]],- ) -> "torch.Tensor":- import torch- import torch.nn.functional as F+ c = torch.empty((bh, n, n), device=left.device, dtype=torch.float32)+ compiled = _get_contract_compiled()+ a_ptr = make_ptr(cutlass.Float32, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)+ b_ptr = make_ptr(cutlass.Float32, b.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)+ c_ptr = make_ptr(cutlass.Float32, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)+ compiled(a_ptr, b_ptr, c_ptr, (bh, n))++ out = c.view(bs, hidden, n, n).permute(0, 2, 3, 1).contiguous()+ return out+++ @torch.inference_mode()+ def custom_kernel(data: Tuple[torch.Tensor, torch.Tensor, Dict[str, torch.Tensor], Dict[str, Any]]) -> torch.Tensor:x, mask, weights, config = data++dim = int(config["dim"])hidden_dim = int(config["hidden_dim"])- if not x.is_cuda:- raise RuntimeError("CUDA only")-if x.dtype != torch.float32:x = x.to(dtype=torch.float32)++ torch.backends.cuda.matmul.allow_tf32 = True+ torch.backends.cudnn.allow_tf32 = True+x = F.layer_norm(x, (dim,), weights["norm.weight"], weights["norm.bias"], 1e-5)left = F.linear(x, weights["left_proj.weight"], None)⋯ 2 unchanged linesmask_f = mask.unsqueeze(-1)if mask_f.dtype != left.dtype:mask_f = mask_f.to(dtype=left.dtype)- left.mul_(mask_f)- right.mul_(mask_f)+ left = left * mask_f+ right = right * mask_fleft_gate = torch.sigmoid(F.linear(x, weights["left_gate.weight"], None))right_gate = torch.sigmoid(F.linear(x, weights["right_gate.weight"], None))out_gate = torch.sigmoid(F.linear(x, weights["out_gate.weight"], None))- left.mul_(left_gate)- right.mul_(right_gate)+ left = left * left_gate+ right = right * right_gateout = _contract_outgoing(left, right)out = F.layer_norm(⋯ 3 unchanged linesweights["to_out_norm.bias"],1e-5,)- out.mul_(out_gate)+ out = out * out_gateout = F.linear(out, weights["to_out.weight"], None)return out__all__ = ["custom_kernel"]-
scrolls · 376 diff lines total
Best evidence level for this revision: reported
JSON