Skip to content
KernelIndex
Search⌘K

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
NVIDIA H100
9.04ms
#54 of 71
2026-01-31

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 lines
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, 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 = 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)
⋯ 5 unchanged lines
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.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 lines
weights["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