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
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.
mma
acc += 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