submission 409251
novo_force · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 100 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-409251?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:cb931d4d7525d3f2430291ad49085a2b0e4445b1170e9ac8d86595c7345f922d
license declaredunknown
license concludedunknown
authorsnovo_force
imported2026-08-15
Kernel source
submission.py100 lines
from __future__ import annotations
from typing import Any, Dict, Tuple
import torch
import torch.nn.functional as F
_FUSED_W_5X: torch.Tensor | None = None
_FUSED_W_5X_META: tuple[int, int, int, int, int, int] | None = None
def _get_fused_w_5x(weights: Dict[str, torch.Tensor], *, device: torch.device) -> torch.Tensor:
global _FUSED_W_5X, _FUSED_W_5X_META
w0 = weights["left_proj.weight"]
w1 = weights["right_proj.weight"]
w2 = weights["left_gate.weight"]
w3 = weights["right_gate.weight"]
w4 = weights["out_gate.weight"]
meta = (
int(device.index) if device.type == "cuda" else -1,
int(w0.data_ptr()),
int(w1.data_ptr()),
int(w2.data_ptr()),
int(w3.data_ptr()),
int(w4.data_ptr()),
)
if _FUSED_W_5X is not None and _FUSED_W_5X_META == meta:
return _FUSED_W_5X
fused = torch.cat((w0, w1, w2, w3, w4), dim=0).contiguous()
_FUSED_W_5X = fused
_FUSED_W_5X_META = meta
return fused
def _contract_outgoing_bmm(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:
bs, n, _, hidden = left.shape
a = left.permute(0, 3, 1, 2).contiguous().reshape(bs * hidden, n, n)
b = right.permute(0, 3, 1, 2).contiguous().reshape(bs * hidden, n, n)
out = torch.bmm(a, b.transpose(1, 2))
return out.reshape(bs, hidden, n, n).permute(0, 2, 3, 1).contiguous()
@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
if not x.is_cuda:
raise RuntimeError("CUDA tensors required")
if x.dtype != torch.float32:
x = x.to(dtype=torch.float32)
dim = int(config["dim"])
hidden_dim = int(config["hidden_dim"])
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
x = x.contiguous()
x = F.layer_norm(x, (dim,), weights["norm.weight"], weights["norm.bias"], 1e-5)
fused_w = _get_fused_w_5x(weights, device=x.device)
proj = F.linear(x, fused_w, None)
left, right, left_gate, right_gate, out_gate = proj.split(hidden_dim, dim=-1)
left_gate.sigmoid_()
right_gate.sigmoid_()
out_gate.sigmoid_()
mask_f = mask.unsqueeze(-1)
if mask_f.dtype != left.dtype:
mask_f = mask_f.to(dtype=left.dtype)
left.mul_(mask_f).mul_(left_gate)
right.mul_(mask_f).mul_(right_gate)
out = _contract_outgoing_bmm(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 · 100 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 409172.
from __future__ import annotations- from typing import Dict, Tuple, Any+ from typing import Any, Dict, Tupleimport torchimport torch.nn.functional as F- def _outgoing_core(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:+ _FUSED_W_5X: torch.Tensor | None = None+ _FUSED_W_5X_META: tuple[int, int, int, int, int, int] | None = None+++ def _get_fused_w_5x(weights: Dict[str, torch.Tensor], *, device: torch.device) -> torch.Tensor:+ global _FUSED_W_5X, _FUSED_W_5X_META++ w0 = weights["left_proj.weight"]+ w1 = weights["right_proj.weight"]+ w2 = weights["left_gate.weight"]+ w3 = weights["right_gate.weight"]+ w4 = weights["out_gate.weight"]++ meta = (+ int(device.index) if device.type == "cuda" else -1,+ int(w0.data_ptr()),+ int(w1.data_ptr()),+ int(w2.data_ptr()),+ int(w3.data_ptr()),+ int(w4.data_ptr()),+ )+ if _FUSED_W_5X is not None and _FUSED_W_5X_META == meta:+ return _FUSED_W_5X++ fused = torch.cat((w0, w1, w2, w3, w4), dim=0).contiguous()+ _FUSED_W_5X = fused+ _FUSED_W_5X_META = meta+ return fused+++ def _contract_outgoing_bmm(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:+ bs, n, _, hidden = left.shape+-- bs, i, k, hidden = left.shape- j = right.shape[1]+ a = left.permute(0, 3, 1, 2).contiguous().reshape(bs * hidden, n, n)+ b = right.permute(0, 3, 1, 2).contiguous().reshape(bs * hidden, n, n)+ out = torch.bmm(a, b.transpose(1, 2))+ return out.reshape(bs, hidden, n, n).permute(0, 2, 3, 1).contiguous()- left_bd = left.permute(0, 3, 1, 2).contiguous().view(bs * hidden, i, k)- right_bd = right.permute(0, 3, 1, 2).contiguous().view(bs * hidden, j, k)- out_bd = torch.bmm(left_bd, right_bd.transpose(1, 2))- return out_bd.view(bs, hidden, i, j).permute(0, 2, 3, 1).contiguous()-@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+ if not x.is_cuda:+ raise RuntimeError("CUDA tensors required")+ if x.dtype != torch.float32:+ x = x.to(dtype=torch.float32)+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 = x.contiguous()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)+ fused_w = _get_fused_w_5x(weights, device=x.device)+ proj = F.linear(x, fused_w, None)+ left, right, left_gate, right_gate, out_gate = proj.split(hidden_dim, dim=-1)+ left_gate.sigmoid_()+ right_gate.sigmoid_()+ out_gate.sigmoid_()+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.mul_(mask_f).mul_(left_gate)+ right.mul_(mask_f).mul_(right_gate)- left = left * left_gate- right = right * right_gate+ out = _contract_outgoing_bmm(left, right)- out = _outgoing_core(left, right)out = F.layer_norm(out,(hidden_dim,),⋯ 1 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 · 122 diff lines total
Best evidence level for this revision: reported
JSON