Skip to content
KernelIndex
Search⌘K

submission 456739

jackkhuu · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-456739?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
39.4ms
#65 of 71
2026-02-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9aed629d1da286ab7836ae27996d2b4741e6f9b79fc98ed14df1a9443ae6b92d
license declaredunknown
license concludedunknown
authorsjackkhuu
imported2026-08-15

Techniques

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

fused-epilogue- Kernel 3 (epilogue): LayerNorm over hidden_dim on out[b,i,j,:], multiply by out_gate[b,i,j,:],
num-warps = 4num_warps=4, num_stages=2
stages = 2num_warps=4, num_stages=2

Kernel source

submission.py449 lines
import math
import torch
import triton
import triton.language as tl

"""
Triangle Multiplicative Update (TriMul, "outgoing") implemented with Triton.

What is fused:
- Kernel 1 (precompute): LayerNorm(x) over dim, five linear projections (left/right and three gates),
  masking and sigmoid gates application for left/right, and sigmoid(out_gate). Stores:
  left[b,i,k,h], right_tmp[b,i,k,h], out_gate[b,i,j,h].
  This fuses LN + 5 matvecs + mask + 3 sigmoids + elementwise gating.
- Kernel 2 (contract): Performs out[b,i,j,h] = sum_k left[b,i,k,h] * right_tmp[b,j,k,h].
  This is the "triangle multiplicative" contraction over k with per-channel outer products.
- Kernel 3 (epilogue): LayerNorm over hidden_dim on out[b,i,j,:], multiply by out_gate[b,i,j,:],
  and final linear projection to dim. This fuses LN + gating + final matvec.

Wrapper (kernel_function):
- Validates inputs, allocates intermediates/outputs, computes launch grids and launches kernels.
- No PyTorch math is used in the wrapper; all computation occurs in Triton kernels.

Notes:
- Uses eps=1e-5 for both LayerNorms to match torch.nn.LayerNorm default.
- All math is in float32 for numerical stability and to match the test.
"""


# --------------------------
# Kernel 1: Precompute projections, gates and gating
# --------------------------
@triton.jit
def _precompute_lr_gates(
    x_ptr,           # *float32 [B, I, J, D]
    mask_ptr,        # *float32 [B, I, J]
    # LayerNorm (over D)
    ln_w_ptr,        # *float32 [D]
    ln_b_ptr,        # *float32 [D]
    # Weights: [H, D] (row-major: out x in)
    w_left_ptr,      # *float32 [H, D]
    w_right_ptr,     # *float32 [H, D]
    w_lg_ptr,        # *float32 [H, D]
    w_rg_ptr,        # *float32 [H, D]
    w_og_ptr,        # *float32 [H, D]
    # Outputs
    left_ptr,        # *float32 [B, I, J, H]
    right_tmp_ptr,   # *float32 [B, I, J, H] (will be read as [B, J, K, H] with J as first pair dim)
    og_ptr,          # *float32 [B, I, J, H]
    # Shapes
    B: tl.constexpr, I: tl.constexpr, J: tl.constexpr, D: tl.constexpr, H: tl.constexpr,
    # Strides
    sxB, sxI, sxJ, sxD,
    smB, smI, smJ,
    soB, soI, soJ, soH,
    # Weight strides
    swLH, swLD, swRH, swRD, swLGH, swLGD, swRGH, swRGD, swOGH, swOGD,
    # Meta
    BLOCK_D: tl.constexpr,
    BLOCK_H: tl.constexpr,
    EPS: tl.constexpr,
):
    # Program IDs: map to one (b, i, j) triple per program
    pid_j = tl.program_id(0)
    pid_i = tl.program_id(1)
    pid_b = tl.program_id(2)

    # Bounds mask for out-of-bounds (in case grid is larger)
    in_bounds = (pid_b < B) & (pid_i < I) & (pid_j < J)

    # Early exit if OOB
    if not in_bounds:
        return

    # Base pointers to current (b,i,j) row vector (over D)
    x_row_ptr = x_ptr + pid_b * sxB + pid_i * sxI + pid_j * sxJ
    # Load mask scalar
    m_val = tl.load(mask_ptr + pid_b * smB + pid_i * smI + pid_j * smJ)
    # First pass: compute mean and variance over D
    offs_d = tl.arange(0, BLOCK_D)
    sum1 = 0.0
    sum2 = 0.0
    for d0 in range(0, D, BLOCK_D):
        d_idx = d0 + offs_d
        d_mask = d_idx < D
        x_chunk = tl.load(x_row_ptr + d_idx * sxD, mask=d_mask, other=0.0)
        sum1 += tl.sum(x_chunk, axis=0)
        sum2 += tl.sum(x_chunk * x_chunk, axis=0)
    Df = tl.full((), D, dtype=tl.float32)
    mean = sum1 / Df
    var = sum2 / Df - mean * mean
    inv_std = 1.0 / tl.sqrt(var + EPS)

    # Prepare LN weights pointers
    ln_w_base = ln_w_ptr
    ln_b_base = ln_b_ptr

    # Accumulate and store tiles over H
    offs_h = tl.arange(0, BLOCK_H)
    for h0 in range(0, H, BLOCK_H):
        h_idx = h0 + offs_h
        h_mask = h_idx < H

        # Initialize accumulators for the five projections
        accL = tl.zeros([BLOCK_H], dtype=tl.float32)
        accR = tl.zeros([BLOCK_H], dtype=tl.float32)
        accLG = tl.zeros([BLOCK_H], dtype=tl.float32)
        accRG = tl.zeros([BLOCK_H], dtype=tl.float32)
        accOG = tl.zeros([BLOCK_H], dtype=tl.float32)

        # Accumulate across D
        for d0 in range(0, D, BLOCK_D):
            d_idx = d0 + offs_d
            d_mask = d_idx < D

            # Load x chunk and apply LayerNorm (affine)
            x_vals = tl.load(x_row_ptr + d_idx * sxD, mask=d_mask, other=0.0)
            w_ln = tl.load(ln_w_base + d_idx, mask=d_mask, other=0.0)
            b_ln = tl.load(ln_b_base + d_idx, mask=d_mask, other=0.0)
            y = (x_vals - mean) * inv_std
            y = y * w_ln + b_ln  # LN output chunk [BLOCK_D]

            # Load weight tiles [BLOCK_H, BLOCK_D] for each projection
            # Pointers shaped as: base + h_idx[:,None]*stride_h + d_idx[None,:]*stride_d
            # left
            wL = tl.load(w_left_ptr + h_idx[:, None] * swLH + d_idx[None, :] * swLD, mask=h_mask[:, None] & d_mask[None, :], other=0.0)
            # right
            wR = tl.load(w_right_ptr + h_idx[:, None] * swRH + d_idx[None, :] * swRD, mask=h_mask[:, None] & d_mask[None, :], other=0.0)
            # left gate
            wLG = tl.load(w_lg_ptr + h_idx[:, None] * swLGH + d_idx[None, :] * swLGD, mask=h_mask[:, None] & d_mask[None, :], other=0.0)
            # right gate
            wRG = tl.load(w_rg_ptr + h_idx[:, None] * swRGH + d_idx[None, :] * swRGD, mask=h_mask[:, None] & d_mask[None, :], other=0.0)
            # out gate
            wOG = tl.load(w_og_ptr + h_idx[:, None] * swOGH + d_idx[None, :] * swOGD, mask=h_mask[:, None] & d_mask[None, :], other=0.0)

            # Row-wise dot: sum over D tile
            accL += tl.sum(wL * y[None, :], axis=1)
            accR += tl.sum(wR * y[None, :], axis=1)
            accLG += tl.sum(wLG * y[None, :], axis=1)
            accRG += tl.sum(wRG * y[None, :], axis=1)
            accOG += tl.sum(wOG * y[None, :], axis=1)

        # Apply sigmoids for gates
        # sigmoid(x) = 1 / (1 + exp(-x))
        lg = 1.0 / (1.0 + tl.exp(-accLG))
        rg = 1.0 / (1.0 + tl.exp(-accRG))
        og = 1.0 / (1.0 + tl.exp(-accOG))

        # Apply mask and gates to left/right (mask expands across H)
        accL = accL * m_val * lg
        accR = accR * m_val * rg

        # Store left[b, i, j, h], right_tmp[b, i, j, h], og[b, i, j, h]
        base_out = pid_b * soB + pid_i * soI + pid_j * soJ
        tl.store(left_ptr + base_out + h_idx * soH, accL, mask=h_mask)
        tl.store(right_tmp_ptr + base_out + h_idx * soH, accR, mask=h_mask)
        tl.store(og_ptr + base_out + h_idx * soH, og, mask=h_mask)


# --------------------------
# Kernel 2: TriMul contraction over k
# out[b, i, j, h] = sum_k left[b, i, k, h] * right_tmp[b, j, k, h]
# --------------------------
@triton.jit
def _trimul_contract(
    left_ptr,        # *float32 [B, I, K, H]
    right_tmp_ptr,   # *float32 [B, J, K, H] but laid out [B, I, J, H], we index first dim as j
    out_ptr,         # *float32 [B, I, J, H]
    B: tl.constexpr, I: tl.constexpr, J: tl.constexpr, K: tl.constexpr, H: tl.constexpr,
    # Strides for left/right/out (all shaped [B, I, J, H])
    sLB, sLI, sLJ, sLH,
    sRB, sRI, sRJ, sRH,   # for right_tmp (same layout as left: [B, I, J, H])
    sOB, sOI, sOJ, sOH,
    # Meta
    BLOCK_M: tl.constexpr,  # tile over i
    BLOCK_N: tl.constexpr,  # tile over j
    BLOCK_H: tl.constexpr,  # tile over h
):
    pid_j_tile = tl.program_id(0)
    pid_i_tile = tl.program_id(1)
    pid_b = tl.program_id(2)

    offs_i = pid_i_tile * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_j = pid_j_tile * BLOCK_N + tl.arange(0, BLOCK_N)
    mask_i = offs_i < I
    mask_j = offs_j < J

    offs_h = tl.arange(0, BLOCK_H)

    # Loop over H tiles
    for h0 in range(0, H, BLOCK_H):
        h_idx = h0 + offs_h
        mask_h = h_idx < H

        # Initialize accumulator [BM, BN, BH]
        acc = tl.zeros([BLOCK_M, BLOCK_N, BLOCK_H], dtype=tl.float32)

        # Sum over k
        for k in range(0, K):
            # Load L[b, i, k, h] -> [BM, BH]
            l_ptrs = left_ptr + pid_b * sLB + offs_i[:, None] * sLI + k * sLJ + h_idx[None, :] * sLH
            l_vals = tl.load(l_ptrs, mask=mask_i[:, None] & mask_h[None, :], other=0.0)

            # Load Rtmp[b, j, k, h] -> [BN, BH] reading from right_tmp as [B, I, J, H] with j as 'I' index
            r_ptrs = right_tmp_ptr + pid_b * sRB + offs_j[:, None] * sRI + k * sRJ + h_idx[None, :] * sRH
            r_vals = tl.load(r_ptrs, mask=mask_j[:, None] & mask_h[None, :], other=0.0)

            # Outer product across (i,j) per-channel h: acc += L[:,None,:] * R[None,:, :]
            acc += l_vals[:, None, :] * r_vals[None, :, :]

        # Store out[b, i, j, h] tile
        out_ptrs = out_ptr + pid_b * sOB + offs_i[:, None, None] * sOI + offs_j[None, :, None] * sOJ + h_idx[None, None, :] * sOH
        tl.store(out_ptrs, acc, mask=mask_i[:, None, None] & mask_j[None, :, None] & mask_h[None, None, :])


# --------------------------
# Kernel 3: LN over H, apply out_gate, and final linear to dim
# --------------------------
@triton.jit
def _epilogue_ln_gate_linear(
    out_h_ptr,     # *float32 [B, I, J, H]  (input from contraction)
    og_ptr,        # *float32 [B, I, J, H]  (sigmoid(out_gate(x)))
    # LayerNorm over H
    ln2_w_ptr,     # *float32 [H]
    ln2_b_ptr,     # *float32 [H]
    # Final projection weights
    w_out_ptr,     # *float32 [D, H]
    # Output
    y_ptr,         # *float32 [B, I, J, D]
    # Shapes
    B: tl.constexpr, I: tl.constexpr, J: tl.constexpr, H: tl.constexpr, D: tl.constexpr,
    # Strides
    sOHB, sOHI, sOHJ, sOHH,    # out_h strides
    sOGB, sOGI, sOGJ, sOGH,    # og strides
    sL2W, sL2B,                # LN2 gamma/beta strides (contiguous along H)
    sWOD, sWOH,                # w_out strides: [D, H]
    sYB, sYI, sYJ, sYD,        # y strides
    # Meta
    BLOCK_H: tl.constexpr,
    BLOCK_P: tl.constexpr,     # tile for output dim D
    EPS: tl.constexpr,
):
    pid_j = tl.program_id(0)
    pid_i = tl.program_id(1)
    pid_b = tl.program_id(2)

    in_bounds = (pid_b < B) & (pid_i < I) & (pid_j < J)
    if not in_bounds:
        return

    # Base pointers
    base_h = out_h_ptr + pid_b * sOHB + pid_i * sOHI + pid_j * sOHJ
    base_og = og_ptr + pid_b * sOGB + pid_i * sOGI + pid_j * sOGJ

    # Pass 1: compute mean/var over H
    offs_h = tl.arange(0, BLOCK_H)
    sum1 = 0.0
    sum2 = 0.0
    for h0 in range(0, H, BLOCK_H):
        h_idx = h0 + offs_h
        h_mask = h_idx < H
        vals = tl.load(base_h + h_idx * sOHH, mask=h_mask, other=0.0)
        sum1 += tl.sum(vals, axis=0)
        sum2 += tl.sum(vals * vals, axis=0)
    Hf = tl.full((), H, dtype=tl.float32)
    mean = sum1 / Hf
    var = sum2 / Hf - mean * mean
    inv_std = 1.0 / tl.sqrt(var + EPS)

    # Tiling over output dim D (projection)
    offs_p = tl.arange(0, BLOCK_P)
    for p0 in range(0, D, BLOCK_P):
        p_idx = p0 + offs_p
        p_mask = p_idx < D

        acc = tl.zeros([BLOCK_P], dtype=tl.float32)

        # Accumulate across H
        for h0 in range(0, H, BLOCK_H):
            h_idx = h0 + offs_h
            h_mask = h_idx < H

            # Load out_h, apply second LN affine, then gate with og
            o = tl.load(base_h + h_idx * sOHH, mask=h_mask, other=0.0)
            og = tl.load(base_og + h_idx * sOGH, mask=h_mask, other=0.0)

            # LN2 gamma/beta over H
            gamma = tl.load(ln2_w_ptr + h_idx * sL2W, mask=h_mask, other=0.0)
            beta = tl.load(ln2_b_ptr + h_idx * sL2B, mask=h_mask, other=0.0)

            normed = ((o - mean) * inv_std) * gamma + beta
            post = normed * og  # apply out_gate after LN2

            # Load weight tile [BLOCK_P, BLOCK_H]
            w_tile = tl.load(w_out_ptr + p_idx[:, None] * sWOD + h_idx[None, :] * sWOH,
                             mask=p_mask[:, None] & h_mask[None, :], other=0.0)

            # Row-wise dot
            acc += tl.sum(w_tile * post[None, :], axis=1)

        # Store y[b, i, j, p_idx]
        y_base = y_ptr + pid_b * sYB + pid_i * sYI + pid_j * sYJ
        tl.store(y_base + p_idx * sYD, acc, mask=p_mask)


def kernel_function(x, mask, weights, config):
    """
    Triton implementation of Triangle Multiplicative Update (outgoing).

    Args:
      x: torch.Tensor [B, N, N, D], float32, CUDA
      mask: torch.Tensor [B, N, N], float32, CUDA
      weights: dict with the following keys and shapes:
        - "norm.weight": [D], "norm.bias": [D]
        - "left_proj.weight":  [H, D]
        - "right_proj.weight": [H, D]
        - "left_gate.weight":  [H, D]
        - "right_gate.weight": [H, D]
        - "out_gate.weight":   [H, D]
        - "to_out_norm.weight": [H], "to_out_norm.bias": [H]
        - "to_out.weight": [D, H]
      config: dict with {"dim": D, "hidden_dim": H}

    Returns:
      y: torch.Tensor [B, N, N, D], float32, CUDA
    """
    assert isinstance(x, torch.Tensor) and isinstance(mask, torch.Tensor)
    assert x.is_cuda and mask.is_cuda, "Inputs must be CUDA tensors"
    assert x.dtype == torch.float32 and mask.dtype == torch.float32, "Expect float32 inputs"

    B, N1, N2, D = x.shape
    assert N1 == N2, "Second and third dims must be equal (square pair matrix)"
    N = N1
    assert config is not None and isinstance(config, dict)
    H = int(config["hidden_dim"])
    D_cfg = int(config["dim"])
    assert D_cfg == D, "config['dim'] must match x.size(-1)"
    device = x.device

    # Validate weights and shapes/dtypes/devices
    req_keys = [
        "norm.weight", "norm.bias",
        "left_proj.weight", "right_proj.weight",
        "left_gate.weight", "right_gate.weight", "out_gate.weight",
        "to_out_norm.weight", "to_out_norm.bias",
        "to_out.weight",
    ]
    for k in req_keys:
        assert k in weights, f"Missing weight: {k}"
        assert isinstance(weights[k], torch.Tensor)
        assert weights[k].is_cuda and weights[k].dtype == torch.float32 and weights[k].device == device

    assert tuple(weights["norm.weight"].shape) == (D,)
    assert tuple(weights["norm.bias"].shape) == (D,)
    assert tuple(weights["left_proj.weight"].shape) == (H, D)
    assert tuple(weights["right_proj.weight"].shape) == (H, D)
    assert tuple(weights["left_gate.weight"].shape) == (H, D)
    assert tuple(weights["right_gate.weight"].shape) == (H, D)
    assert tuple(weights["out_gate.weight"].shape) == (H, D)
    assert tuple(weights["to_out_norm.weight"].shape) == (H,)
    assert tuple(weights["to_out_norm.bias"].shape) == (H,)
    assert tuple(weights["to_out.weight"].shape) == (D, H)

    # Allocate intermediates/results
    left = torch.empty((B, N, N, H), device=device, dtype=torch.float32)
    right_tmp = torch.empty((B, N, N, H), device=device, dtype=torch.float32)
    ogate = torch.empty((B, N, N, H), device=device, dtype=torch.float32)
    out_h = torch.empty((B, N, N, H), device=device, dtype=torch.float32)
    y = torch.empty((B, N, N, D), device=device, dtype=torch.float32)

    # Common constants
    EPS = 1e-5

    # --------------------------
    # Launch Kernel 1: Precompute projections and gates
    # Grid: (J, I, B) -> one (b,i,j) row per program
    # --------------------------
    BLOCK_D = 64  # tile for input dim D
    BLOCK_H = 64  # tile for hidden dim H
    grid1 = (N, N, B)

    _precompute_lr_gates[grid1](
        x, mask,
        weights["norm.weight"], weights["norm.bias"],
        weights["left_proj.weight"], weights["right_proj.weight"],
        weights["left_gate.weight"], weights["right_gate.weight"],
        weights["out_gate.weight"],
        left, right_tmp, ogate,
        B, N, N, D, H,
        x.stride(0), x.stride(1), x.stride(2), x.stride(3),
        mask.stride(0), mask.stride(1), mask.stride(2),
        left.stride(0), left.stride(1), left.stride(2), left.stride(3),
        # Weight strides (H, D)
        weights["left_proj.weight"].stride(0), weights["left_proj.weight"].stride(1),
        weights["right_proj.weight"].stride(0), weights["right_proj.weight"].stride(1),
        weights["left_gate.weight"].stride(0), weights["left_gate.weight"].stride(1),
        weights["right_gate.weight"].stride(0), weights["right_gate.weight"].stride(1),
        weights["out_gate.weight"].stride(0), weights["out_gate.weight"].stride(1),
        BLOCK_D=BLOCK_D, BLOCK_H=BLOCK_H, EPS=EPS,
        num_warps=4, num_stages=2
    )

    # --------------------------
    # Launch Kernel 2: TriMul contraction
    # Grid tiles over (j, i, b)
    # --------------------------
    # Choose tiles
    BM = 8
    BN = 8
    BH = 32
    grid2 = (triton.cdiv(N, BN), triton.cdiv(N, BM), B)
    _trimul_contract[grid2](
        left, right_tmp, out_h,
        B, N, N, N, H,
        left.stride(0), left.stride(1), left.stride(2), left.stride(3),
        right_tmp.stride(0), right_tmp.stride(1), right_tmp.stride(2), right_tmp.stride(3),
        out_h.stride(0), out_h.stride(1), out_h.stride(2), out_h.stride(3),
        BLOCK_M=BM, BLOCK_N=BN, BLOCK_H=BH,
        num_warps=4, num_stages=2
    )

    # --------------------------
    # Launch Kernel 3: LN over H, gate with out_gate, final linear to D
    # Grid: (J, I, B) per row
    # --------------------------
    BH2 = 64
    BP = 64  # tile for D
    grid3 = (N, N, B)
    _epilogue_ln_gate_linear[grid3](
        out_h, ogate,
        weights["to_out_norm.weight"], weights["to_out_norm.bias"],
        weights["to_out.weight"],
        y,
        B, N, N, H, D,
        out_h.stride(0), out_h.stride(1), out_h.stride(2), out_h.stride(3),
        ogate.stride(0), ogate.stride(1), ogate.stride(2), ogate.stride(3),
        # LN2 gamma/beta strides (contiguous along last dim)
        weights["to_out_norm.weight"].stride(0), weights["to_out_norm.bias"].stride(0),
        # w_out strides: [D, H]
        weights["to_out.weight"].stride(0), weights["to_out.weight"].stride(1),
        y.stride(0), y.stride(1), y.stride(2), y.stride(3),
        BLOCK_H=BH2, BLOCK_P=BP, EPS=EPS,
        num_warps=4, num_stages=2
    )

    return y

def custom_kernel(input):
    return kernel_function(*input)
scrolls · 449 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