Skip to content
KernelIndex
Search⌘K

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
NVIDIA A100
20.5ms
#58 of 69
2026-01-31

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-memorysmem_stride = _TILE_K + 1
tile-k = 32_TILE_K = 32

Kernel 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 lines
from 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.shape
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)
+ 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 lines
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 = 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.mul_(left_gate)
- right.mul_(right_gate)
+ left = left * left_gate
+ right = right * right_gate
out = _contract_outgoing(left, right)
out = F.layer_norm(
⋯ 3 unchanged lines
weights["to_out_norm.bias"],
1e-5,
)
- out.mul_(out_gate)
+ out = out * out_gate
out = 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