submission 417247
shiyegao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 179 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-417247?include=source"interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
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:5ffaffe9d33ee71dddf7458dd9e4bbbc5f3cb59d9f82070da9376b2b007ad99b
license declaredunknown
license concludedunknown
authorsshiyegao
imported2026-08-15
Kernel source
submission.py179 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, if_generate, yield_out
class _OutgoingContractKernel:
def __init__(self, threads: int = 256) -> None:
self.threads = int(threads)
@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
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),
),
)
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
@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()
idx = bx * bdx + tx
total = bs * n * n * hidden
def _do_one():
h = idx % hidden
t0 = idx // hidden
j = t0 % n
t1 = t0 // n
i = t1 % n
b = t1 // n
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
if_generate(idx < total, _do_one)
_CONTRACT = _OutgoingContractKernel()
_CONTRACT_COMPILED = None
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",
)
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("This kernel requires CUDA tensors.")
if left.dtype != torch.float32 or right.dtype != torch.float32:
raise RuntimeError("This kernel expects float32 tensors.")
left = left.contiguous()
right = right.contiguous()
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))
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)
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 · 179 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON