submission 417295
shiyegao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 199 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-417295?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:0ab6a2335f8dc2a0d76fd7c2a192fb022439c38bac489f67b0a066ff0683a625
license declaredunknown
license concludedunknown
authorsshiyegao
imported2026-08-15
Kernel source
submission.py199 lines
from __future__ import annotations
from typing import Any, Dict, Tuple
_EXT_MOD = None
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");
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));
const float alpha = 1.0f;
const float beta = 0.0f;
// 目标: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;
cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();
_cublas_ok(cublasSetMathMode(handle, CUBLAS_TENSOR_OP_MATH), "cublasSetMathMode failed");
_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");
return out;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("contract_bf16", &contract_bf16, "contract_bf16");
}
"""
)
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
_EXT_MOD = load_inline(
name="trimul_cublas_ext",
cpp_sources=cpp,
**kwargs,
)
return _EXT_MOD
def _contract_outgoing(left: "torch.Tensor", right: "torch.Tensor") -> "torch.Tensor":
import torch
if 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)
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)
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
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
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)
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.mul_(mask_f)
right.mul_(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.mul_(left_gate)
right.mul_(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.mul_(out_gate)
out = F.linear(out, weights["to_out.weight"], None)
return out
__all__ = ["custom_kernel"]
scrolls · 199 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 417249.
⋯ 1 unchanged linesfrom 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, if_generate, yield_out+ _EXT_MOD = None- class _OutgoingContractKernel:- def __init__(self, threads: int = 256) -> None:- self.threads = int(threads)+ def _load_ext():+ global _EXT_MOD+ if _EXT_MOD is not None:+ return _EXT_MOD- @cute.jit- def __call__(- self,- left_ptr: "cute.Pointer",- right_ptr: "cute.Pointer",- out_ptr: "cute.Pointer",- problem: tuple,- ):- bs, n, hidden = problem- stride_bs = n * n * hidden- stride_i = n * hidden- stride_k = hidden+ import inspect- left = cute.make_tensor(- left_ptr,- cute.make_layout(- (bs, n, n, hidden),- stride=(stride_bs, stride_i, stride_k, 1),- ),- )- right = cute.make_tensor(- right_ptr,- cute.make_layout(- (bs, n, n, hidden),- stride=(stride_bs, stride_i, stride_k, 1),- ),- )- out = cute.make_tensor(- out_ptr,- cute.make_layout(- (bs, n, n, hidden),- stride=(stride_bs, stride_i, stride_k, 1),- ),- )+ import torch+ from torch.utils.cpp_extension import load_inline- total = bs * n * n * hidden- grid = (total + self.threads - 1) // self.threads- self.kernel(left, right, out, bs, n, hidden).launch(- grid=[grid, 1, 1],- block=[self.threads, 1, 1],- )- return++ cpp = (+ r"""+ #include <torch/extension.h>+ #include <ATen/cuda/CUDABlas.h>+ #include <cublas_v2.h>- @cute.kernel- def kernel(- self,- left: "cute.Tensor",- right: "cute.Tensor",- out: "cute.Tensor",- bs: int,- n: int,- hidden: int,- ):- tx, _, _ = cute.arch.thread_idx()- bx, _, _ = cute.arch.block_idx()- bdx, _, _ = cute.arch.block_dim()+ static inline void _check(bool ok, const char* msg) {+ if (!ok) {+ throw std::runtime_error(msg);+ }+ }- idx = bx * bdx + tx- total = bs * n * n * hidden+ static inline void _cublas_ok(cublasStatus_t st, const char* msg) {+ if (st != CUBLAS_STATUS_SUCCESS) {+ throw std::runtime_error(msg);+ }+ }- def _do_one():- h = idx % hidden- t0 = idx // hidden- j = t0 % n- t1 = t0 // n- i = t1 % n- b = t1 // n+ 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");- acc0 = cutlass.Float32(0.0)- for k, acc, acc_out in for_generate(0, n, iter_args=[acc0]):- acc = acc + left[b, i, k, h] * right[b, j, k, h]- yield_out([acc])- out[b, i, j, h] = acc_out+ 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");- if_generate(idx < total, _do_one)+ auto out = torch::empty({batch, n0, n0}, a.options().dtype(at::kFloat));+ const float alpha = 1.0f;+ const float beta = 0.0f;- _CONTRACT = _OutgoingContractKernel()- _CONTRACT_COMPILED = None+ // 目标: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;+ cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();+ _cublas_ok(cublasSetMathMode(handle, CUBLAS_TENSOR_OP_MATH), "cublasSetMathMode failed");- def _get_contract_compiled():- global _CONTRACT_COMPILED- if _CONTRACT_COMPILED is not None:- return _CONTRACT_COMPILED- left_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)- right_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)- out_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)- _CONTRACT_COMPILED = cute.compile(- _CONTRACT,- left_ptr,- right_ptr,- out_ptr,- (0, 0, 0),- options="--opt-level 3",+ _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");++ return out;+ }++ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {+ m.def("contract_bf16", &contract_bf16, "contract_bf16");+ }+ """)- return _CONTRACT_COMPILED++ 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- def _contract_outgoing(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:- bs, n, _, hidden = left.shape+ _EXT_MOD = load_inline(+ name="trimul_cublas_ext",+ cpp_sources=cpp,+ **kwargs,+ )+ return _EXT_MOD++ def _contract_outgoing(left: "torch.Tensor", right: "torch.Tensor") -> "torch.Tensor":+ import torch+if not (left.is_cuda and right.is_cuda):- raise RuntimeError("This kernel requires CUDA tensors.")- if left.dtype != torch.float32 or right.dtype != torch.float32:- raise RuntimeError("This kernel expects float32 tensors.")+ 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)- left = left.contiguous()- right = right.contiguous()++ 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)- out = torch.empty((bs, n, n, hidden), device=left.device, dtype=torch.float32)- compiled = _get_contract_compiled()- left_ptr = make_ptr(cutlass.Float32, left.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)- right_ptr = make_ptr(cutlass.Float32, right.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)- out_ptr = make_ptr(cutlass.Float32, out.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)- compiled(left_ptr, right_ptr, out_ptr, (bs, n, hidden))+ 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- @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+ 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+ x, mask, weights, config = datadim = 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)⋯ 5 unchanged linesmask_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.mul_(mask_f)+ right.mul_(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+ left.mul_(left_gate)+ right.mul_(right_gate)out = _contract_outgoing(left, right)out = F.layer_norm(⋯ 3 unchanged linesweights["to_out_norm.bias"],1e-5,)- out = out * out_gate+ out.mul_(out_gate)out = F.linear(out, weights["to_out.weight"], None)return out__all__ = ["custom_kernel"]+
scrolls · 307 diff lines total
Best evidence level for this revision: reported
JSON