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
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 = 4
num_warps=4, num_stages=2stages = 2
num_warps=4, num_stages=2tile-m = 8
BM = 8tile-n = 8
BN = 8Kernel 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