Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
3.90ms
#21 of 43
2026-01-29

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, Tuple
import torch
import 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 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 · 122 diff lines total

Best evidence level for this revision: reported

JSON