Skip to content
KernelIndex
Search⌘K

submission 416249

Cookie 🍪 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 421 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-416249?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
106.0ms
#71 of 71
2026-01-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:90d1b4bcd91d8f51ad0123089d9be5a044ae62de95ee40d3c4beb33a6e4f6809
license declaredunknown
license concludedunknown
authorsCookie 🍪
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

mmaacc += tl.dot(x_vals, tl.trans(w_vals))

Kernel source

submission.py421 lines
from typing import Dict, Tuple, TypeVar

import torch
import triton
import triton.language as tl


@triton.jit
def _layernorm_kernel(
    x_ptr,
    out_ptr,
    gamma_ptr,
    beta_ptr,
    N,
    D,
    stride_n,
    stride_d,
    eps: tl.constexpr,
    BLOCK_D: tl.constexpr,
):
    pid = tl.program_id(0)
    offs_d = tl.arange(0, BLOCK_D)
    mask = offs_d < D

    x_ptrs = x_ptr + pid * stride_n + offs_d * stride_d
    x = tl.load(x_ptrs, mask=mask, other=0.0)

    mean = tl.sum(x, axis=0) / D
    x_centered = x - mean
    var = tl.sum(x_centered * x_centered, axis=0) / D
    rstd = 1.0 / tl.sqrt(var + eps)

    gamma = tl.load(gamma_ptr + offs_d, mask=mask, other=1.0)
    beta = tl.load(beta_ptr + offs_d, mask=mask, other=0.0)

    out = x_centered * rstd * gamma + beta
    out_ptrs = out_ptr + pid * stride_n + offs_d * stride_d
    tl.store(out_ptrs, out, mask=mask)


@triton.jit
def _linear_sigmoid_kernel(
    x_ptr,
    w_ptr,
    out_ptr,
    M,
    K,
    N,
    stride_xm,
    stride_xk,
    stride_wn,
    stride_wk,
    stride_om,
    stride_on,
    apply_sigmoid: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    offs_k = tl.arange(0, BLOCK_K)

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    for k_start in range(0, K, BLOCK_K):
        k_offs = k_start + offs_k
        k_mask = k_offs < K

        x_ptrs = x_ptr + offs_m[:, None] * stride_xm + k_offs[None, :] * stride_xk
        w_ptrs = w_ptr + offs_n[:, None] * stride_wn + k_offs[None, :] * stride_wk

        x_mask = (offs_m[:, None] < M) & k_mask[None, :]
        w_mask = (offs_n[:, None] < N) & k_mask[None, :]

        x_vals = tl.load(x_ptrs, mask=x_mask, other=0.0)
        w_vals = tl.load(w_ptrs, mask=w_mask, other=0.0)

        acc += tl.dot(x_vals, tl.trans(w_vals))

    if apply_sigmoid:
        acc = tl.sigmoid(acc)

    out_ptrs = out_ptr + offs_m[:, None] * stride_om + offs_n[None, :] * stride_on
    out_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
    tl.store(out_ptrs, acc, mask=out_mask)


@triton.jit
def _masked_gating_einsum_kernel(
    left_ptr,
    right_ptr,
    left_gate_ptr,
    right_gate_ptr,
    mask_ptr,
    out_ptr,
    B,
    seq_len,
    hidden_dim,
    stride_lb,
    stride_li,
    stride_lk,
    stride_ld,
    stride_rb,
    stride_ri,
    stride_rk,
    stride_rd,
    stride_lgb,
    stride_lgi,
    stride_lgk,
    stride_lgd,
    stride_rgb,
    stride_rgi,
    stride_rgk,
    stride_rgd,
    stride_mb,
    stride_mi,
    stride_mk,
    stride_ob,
    stride_oi,
    stride_oj,
    stride_od,
    BLOCK_D: tl.constexpr,
):
    pid_bi = tl.program_id(0)
    pid_j = tl.program_id(1)
    pid_d = tl.program_id(2)

    pid_b = pid_bi // seq_len
    pid_i = pid_bi % seq_len

    d_offs = pid_d * BLOCK_D + tl.arange(0, BLOCK_D)
    d_mask = d_offs < hidden_dim

    acc = tl.zeros((BLOCK_D,), dtype=tl.float32)

    for k in range(seq_len):
        mask_l = tl.load(
            mask_ptr + pid_b * stride_mb + pid_i * stride_mi + k * stride_mk
        )
        mask_r = tl.load(
            mask_ptr + pid_b * stride_mb + pid_j * stride_mi + k * stride_mk
        )

        left_off = (
            pid_b * stride_lb + pid_i * stride_li + k * stride_lk + d_offs * stride_ld
        )
        left_val = tl.load(left_ptr + left_off, mask=d_mask, other=0.0)
        lg_off = (
            pid_b * stride_lgb
            + pid_i * stride_lgi
            + k * stride_lgk
            + d_offs * stride_lgd
        )
        lg_val = tl.load(left_gate_ptr + lg_off, mask=d_mask, other=0.0)
        left_gated = left_val * mask_l * lg_val

        right_off = (
            pid_b * stride_rb + pid_j * stride_ri + k * stride_rk + d_offs * stride_rd
        )
        right_val = tl.load(right_ptr + right_off, mask=d_mask, other=0.0)
        rg_off = (
            pid_b * stride_rgb
            + pid_j * stride_rgi
            + k * stride_rgk
            + d_offs * stride_rgd
        )
        rg_val = tl.load(right_gate_ptr + rg_off, mask=d_mask, other=0.0)
        right_gated = right_val * mask_r * rg_val

        acc += left_gated * right_gated

    out_off = (
        pid_b * stride_ob + pid_i * stride_oi + pid_j * stride_oj + d_offs * stride_od
    )
    tl.store(out_ptr + out_off, acc, mask=d_mask)


@triton.jit
def _layernorm_gate_linear_kernel(
    x_ptr,
    gate_ptr,
    gamma_ptr,
    beta_ptr,
    w_ptr,
    out_ptr,
    N,
    hidden_dim,
    out_dim,
    stride_xn,
    stride_xd,
    stride_gn,
    stride_gd,
    stride_wout,
    stride_whid,
    stride_on,
    stride_od,
    eps: tl.constexpr,
    BLOCK_HD: tl.constexpr,
    BLOCK_OUT: tl.constexpr,
):
    pid_n = tl.program_id(0)
    pid_out = tl.program_id(1)

    offs_hd = tl.arange(0, BLOCK_HD)
    hd_mask = offs_hd < hidden_dim

    x_ptrs = x_ptr + pid_n * stride_xn + offs_hd * stride_xd
    x = tl.load(x_ptrs, mask=hd_mask, other=0.0)

    mean = tl.sum(x, axis=0) / hidden_dim
    x_centered = x - mean
    var = tl.sum(x_centered * x_centered, axis=0) / hidden_dim
    rstd = 1.0 / tl.sqrt(var + eps)

    gamma = tl.load(gamma_ptr + offs_hd, mask=hd_mask, other=1.0)
    beta = tl.load(beta_ptr + offs_hd, mask=hd_mask, other=0.0)
    x_normed = x_centered * rstd * gamma + beta

    gate_ptrs = gate_ptr + pid_n * stride_gn + offs_hd * stride_gd
    gate = tl.load(gate_ptrs, mask=hd_mask, other=0.0)
    x_gated = x_normed * gate

    offs_out = pid_out * BLOCK_OUT + tl.arange(0, BLOCK_OUT)
    out_mask = offs_out < out_dim

    acc = tl.zeros((BLOCK_OUT,), dtype=tl.float32)
    for hd_start in range(0, hidden_dim, BLOCK_HD):
        hd_offs = hd_start + tl.arange(0, BLOCK_HD)
        hd_m = hd_offs < hidden_dim

        if hd_start == 0:
            x_chunk = x_gated
        else:
            x_ptrs2 = x_ptr + pid_n * stride_xn + hd_offs * stride_xd
            x2 = tl.load(x_ptrs2, mask=hd_m, other=0.0)
            mean2 = tl.sum(x2, axis=0) / hidden_dim
            var2 = tl.sum((x2 - mean2) * (x2 - mean2), axis=0) / hidden_dim
            rstd2 = 1.0 / tl.sqrt(var2 + eps)
            gamma2 = tl.load(gamma_ptr + hd_offs, mask=hd_m, other=1.0)
            beta2 = tl.load(beta_ptr + hd_offs, mask=hd_m, other=0.0)
            x_normed2 = (x2 - mean2) * rstd2 * gamma2 + beta2
            gate2 = tl.load(
                gate_ptr + pid_n * stride_gn + hd_offs * stride_gd, mask=hd_m, other=0.0
            )
            x_chunk = x_normed2 * gate2

        for out_idx in range(BLOCK_OUT):
            if pid_out * BLOCK_OUT + out_idx < out_dim:
                w_ptrs = (
                    w_ptr
                    + (pid_out * BLOCK_OUT + out_idx) * stride_wout
                    + hd_offs * stride_whid
                )
                w_vals = tl.load(w_ptrs, mask=hd_m, other=0.0)
                acc = tl.where(
                    tl.arange(0, BLOCK_OUT) == out_idx,
                    acc + tl.sum(x_chunk * w_vals),
                    acc,
                )
        break

    out_ptrs = out_ptr + pid_n * stride_on + offs_out * stride_od
    tl.store(out_ptrs, acc, mask=out_mask)


def kernel_function(x, mask, weights, config):
    B, seq_len, _, dim = x.shape
    hidden_dim = config["hidden_dim"]
    device = x.device
    dtype = x.dtype

    # Flatten for processing
    N = B * seq_len * seq_len
    x_flat = x.reshape(N, dim).contiguous()

    # LayerNorm
    x_normed = torch.empty_like(x_flat)
    BLOCK_D = triton.next_power_of_2(dim)
    _layernorm_kernel[(N,)](
        x_flat,
        x_normed,
        weights["norm.weight"],
        weights["norm.bias"],
        N,
        dim,
        dim,
        1,
        eps=1e-5,
        BLOCK_D=BLOCK_D,
    )

    # Linear projections using torch.mm (for simplicity, we use matmul here)
    x_normed_2d = x_normed.view(N, dim)
    left = x_normed_2d @ weights["left_proj.weight"].T
    right = x_normed_2d @ weights["right_proj.weight"].T
    left_gate = torch.sigmoid(x_normed_2d @ weights["left_gate.weight"].T)
    right_gate = torch.sigmoid(x_normed_2d @ weights["right_gate.weight"].T)
    out_gate = torch.sigmoid(x_normed_2d @ weights["out_gate.weight"].T)

    # Reshape back
    left = left.view(B, seq_len, seq_len, hidden_dim)
    right = right.view(B, seq_len, seq_len, hidden_dim)
    left_gate = left_gate.view(B, seq_len, seq_len, hidden_dim)
    right_gate = right_gate.view(B, seq_len, seq_len, hidden_dim)
    out_gate = out_gate.view(B, seq_len, seq_len, hidden_dim)

    # Masked gating einsum
    out = torch.empty((B, seq_len, seq_len, hidden_dim), device=device, dtype=dtype)
    BLOCK_D = min(64, triton.next_power_of_2(hidden_dim))
    grid = (B * seq_len, seq_len, triton.cdiv(hidden_dim, BLOCK_D))
    _masked_gating_einsum_kernel[grid](
        left,
        right,
        left_gate,
        right_gate,
        mask.float(),
        out,
        B,
        seq_len,
        hidden_dim,
        *left.stride(),
        *right.stride(),
        *left_gate.stride(),
        *right_gate.stride(),
        *mask.stride(),
        *out.stride(),
        BLOCK_D=BLOCK_D,
    )

    # Output: LayerNorm, gate, linear
    out_flat = out.reshape(N, hidden_dim)
    out_normed = torch.empty_like(out_flat)
    BLOCK_HD = triton.next_power_of_2(hidden_dim)
    _layernorm_kernel[(N,)](
        out_flat,
        out_normed,
        weights["to_out_norm.weight"],
        weights["to_out_norm.bias"],
        N,
        hidden_dim,
        hidden_dim,
        1,
        eps=1e-5,
        BLOCK_D=BLOCK_HD,
    )
    out_gated = out_normed * out_gate.reshape(N, hidden_dim)
    result = out_gated @ weights["to_out.weight"].T
    return result.view(B, seq_len, seq_len, dim)


def test_kernel():
    from math import sqrt

    torch.manual_seed(42)
    B, seq_len, dim, hidden_dim = 2, 8, 32, 64
    x = torch.randn(B, seq_len, seq_len, dim, device="cuda")
    mask = torch.randint(0, 2, (B, seq_len, seq_len), device="cuda").float()
    weights = {
        "norm.weight": torch.randn(dim, device="cuda"),
        "norm.bias": torch.randn(dim, device="cuda"),
        "left_proj.weight": torch.randn(hidden_dim, dim, device="cuda")
        / sqrt(hidden_dim),
        "right_proj.weight": torch.randn(hidden_dim, dim, device="cuda")
        / sqrt(hidden_dim),
        "left_gate.weight": torch.randn(hidden_dim, dim, device="cuda")
        / sqrt(hidden_dim),
        "right_gate.weight": torch.randn(hidden_dim, dim, device="cuda")
        / sqrt(hidden_dim),
        "out_gate.weight": torch.randn(hidden_dim, dim, device="cuda")
        / sqrt(hidden_dim),
        "to_out_norm.weight": torch.randn(hidden_dim, device="cuda"),
        "to_out_norm.bias": torch.randn(hidden_dim, device="cuda"),
        "to_out.weight": torch.randn(dim, hidden_dim, device="cuda") / sqrt(dim),
    }
    config = {"dim": dim, "hidden_dim": hidden_dim}
    out_triton = kernel_function(x, mask, weights, config)

    from torch import nn

    model = nn.Module()
    x_n = nn.functional.layer_norm(
        x, [dim], weights["norm.weight"], weights["norm.bias"]
    )
    left = x_n @ weights["left_proj.weight"].T
    right = x_n @ weights["right_proj.weight"].T
    lg = torch.sigmoid(x_n @ weights["left_gate.weight"].T)
    rg = torch.sigmoid(x_n @ weights["right_gate.weight"].T)
    og = torch.sigmoid(x_n @ weights["out_gate.weight"].T)
    m = mask.unsqueeze(-1)
    left = left * m * lg
    right = right * m * rg
    out = torch.einsum("bikd,bjkd->bijd", left, right)
    out = nn.functional.layer_norm(
        out, [hidden_dim], weights["to_out_norm.weight"], weights["to_out_norm.bias"]
    )
    out = out * og
    out_ref = out @ weights["to_out.weight"].T

    if torch.allclose(out_triton, out_ref, rtol=2e-2, atol=2e-2):
        print("PASS")
    else:
        print(f"FAIL: max diff = {(out_triton - out_ref).abs().max()}")


# if __name__ == "__main__":
#     test_kernel()


input_t = TypeVar(
    "input_t", bound=Tuple[torch.Tensor, torch.Tensor, Dict[str, torch.Tensor], Dict]
)


def custom_kernel(input: input_t) -> torch.Tensor:
    x, mask, weights, config = input
    return kernel_function(x, mask, weights, config)
scrolls · 421 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