submission 408244
Zeyu Shen · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 141 lines, June 9 Researcher Reciprocity License v1.0.
fused_frontend_backend_v108.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-408244?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:87ac8d379d021709f1eafabeb2d67c346a59a1e54e767490cd1c70cefb655f0c
license declaredunknown
license concludedunknown
authorsZeyu Shen
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
l_acc += tl.dot(xn, tl.load(W_ptr + 0*w_stride + w_off, mask=c_mask[:, None], other=0.0).to(tl.float16))num-warps = 8
fused_frontend_v108[(B, N, (N+64-1)//64)](x, mask, w_5, weights["norm.weight"], weights["norm.bias"], l, r, og, B, N, C, D, *x.stride(), *mask.stride(), 1e-5, 64, 64, num_warps=8, num_stages=2)stages = 2
fused_frontend_v108[(B, N, (N+64-1)//64)](x, mask, w_5, weights["norm.weight"], weights["norm.bias"], l, r, og, B, N, C, D, *x.stride(), *mask.stride(), 1e-5, 64, 64, num_warps=8, num_stages=2)Kernel source
fused_frontend_backend_v108.py141 lines
import torch
import triton
import triton.language as tl
@triton.jit
def fused_frontend_v108(
X_ptr, Mask_ptr,
W_ptr, NW_ptr, NB_ptr,
L_ptr, R_ptr, OG_ptr,
B, N, C, D: tl.constexpr,
stride_xb, stride_xn1, stride_xn2, stride_xc,
stride_mb, stride_mn1, stride_mn2,
eps, BLOCK_N: tl.constexpr, BLOCK_C: tl.constexpr
):
pid_b = tl.program_id(0)
pid_n1 = tl.program_id(1)
pid_n2_start = tl.program_id(2) * BLOCK_N
n2 = pid_n2_start + tl.arange(0, BLOCK_N)
mask_n2 = n2 < N
# Single-pass LayerNorm statistics to minimize global memory reads
s1 = tl.zeros([BLOCK_N], dtype=tl.float32)
s2 = tl.zeros([BLOCK_N], dtype=tl.float32)
for c_off in range(0, C, BLOCK_C):
rc = c_off + tl.arange(0, BLOCK_C)
c_mask = rc < C
x = tl.load(X_ptr + pid_b*stride_xb + pid_n1*stride_xn1 + n2[:, None]*stride_xn2 + rc[None, :]*stride_xc, mask=mask_n2[:, None] & c_mask[None, :], other=0.0).to(tl.float32)
s1 += tl.sum(x, axis=1)
s2 += tl.sum(x*x, axis=1)
mean = (s1 / C)[:, None]
var = tl.maximum(0.0, (s2 / C)[:, None] - mean*mean)
rstd = 1.0 / tl.sqrt(var + eps)
# Projection Pass
BLOCK_D: tl.constexpr = 128
l_acc = tl.zeros([BLOCK_N, BLOCK_D], dtype=tl.float32)
r_acc = tl.zeros([BLOCK_N, BLOCK_D], dtype=tl.float32)
lg_acc = tl.zeros([BLOCK_N, BLOCK_D], dtype=tl.float32)
rg_acc = tl.zeros([BLOCK_N, BLOCK_D], dtype=tl.float32)
og_acc = tl.zeros([BLOCK_N, BLOCK_D], dtype=tl.float32)
rd = tl.arange(0, BLOCK_D)
# Pre-calculate weight offsets to reduce loop overhead
w_stride = BLOCK_D * C
for c_off in range(0, C, BLOCK_C):
rc = c_off + tl.arange(0, BLOCK_C)
c_mask = rc < C
x = tl.load(X_ptr + pid_b*stride_xb + pid_n1*stride_xn1 + n2[:, None]*stride_xn2 + rc[None, :]*stride_xc, mask=mask_n2[:, None] & c_mask[None, :], other=0.0).to(tl.float32)
nw = tl.load(NW_ptr + rc, mask=c_mask, other=0.0)
nb = tl.load(NB_ptr + rc, mask=c_mask, other=0.0)
xn = ((x - mean) * rstd * nw[None, :] + nb[None, :]).to(tl.float16)
w_off = rd[None, :] * C + rc[:, None]
# Batch load weights if possible or keep separate to manage register pressure
l_acc += tl.dot(xn, tl.load(W_ptr + 0*w_stride + w_off, mask=c_mask[:, None], other=0.0).to(tl.float16))
r_acc += tl.dot(xn, tl.load(W_ptr + 1*w_stride + w_off, mask=c_mask[:, None], other=0.0).to(tl.float16))
lg_acc += tl.dot(xn, tl.load(W_ptr + 2*w_stride + w_off, mask=c_mask[:, None], other=0.0).to(tl.float16))
rg_acc += tl.dot(xn, tl.load(W_ptr + 3*w_stride + w_off, mask=c_mask[:, None], other=0.0).to(tl.float16))
og_acc += tl.dot(xn, tl.load(W_ptr + 4*w_stride + w_off, mask=c_mask[:, None], other=0.0).to(tl.float16))
m = tl.load(Mask_ptr + pid_b*stride_mb + pid_n1*stride_mn1 + n2, mask=mask_n2, other=0.0)[:, None]
l = (l_acc * m * tl.sigmoid(lg_acc)).to(tl.float16)
r = (r_acc * m * tl.sigmoid(rg_acc)).to(tl.float16)
og = tl.sigmoid(og_acc).to(tl.float16)
# Store L/R in [B, D, N, N] layout for BMM (strided write)
l_base = pid_b * (D*N*N) + rd[None, :] * (N*N) + pid_n1 * N + n2[:, None]
tl.store(L_ptr + l_base, l, mask=mask_n2[:, None])
tl.store(R_ptr + l_base, r, mask=mask_n2[:, None])
# Store OG in [B, N, N, D] layout for backend (coalesced write)
og_off = pid_b*(N*N*D) + pid_n1*(N*D) + n2[:, None]*D + rd[None, :]
tl.store(OG_ptr + og_off, og, mask=mask_n2[:, None])
@triton.jit
def fused_backend_v108(
BMM_ptr, OG_ptr, NW_ptr, NB_ptr, W_ptr, Out_ptr,
B, N, D: tl.constexpr, C,
stride_ob, stride_on1, stride_on2, stride_oc,
eps, BLOCK_N: tl.constexpr, BLOCK_C: tl.constexpr
):
pid_b, pid_n1 = tl.program_id(0), tl.program_id(1)
n2_start = tl.program_id(2) * BLOCK_N
n2 = n2_start + tl.arange(0, BLOCK_N)
mask_n2 = n2 < N
rd = tl.arange(0, D)
# Load from [B, D, N, N] (BMM output layout)
# Optimization: Load entire D dimension into registers for LayerNorm
x = tl.load(BMM_ptr + pid_b*(D*N*N) + rd[:, None]*N*N + pid_n1*N + n2[None, :], mask=mask_n2[None, :], other=0.0).to(tl.float32)
mean = (tl.sum(x, axis=0) / D)[None, :]
diff = x - mean
var = (tl.sum(diff * diff, axis=0) / D)[None, :]
rstd = 1.0 / tl.sqrt(var + eps)
nw = tl.load(NW_ptr + rd)[:, None]
nb = tl.load(NB_ptr + rd)[:, None]
xn = (diff * rstd * nw + nb).to(tl.float16)
# Load OG from [B, N, N, D] (coalesced)
og = tl.load(OG_ptr + pid_b*(N*N*D) + pid_n1*(N*D) + n2[:, None]*D + rd[None, :], mask=mask_n2[:, None], other=0.0).to(tl.float16)
xf = (tl.trans(xn) * og).to(tl.float16)
for c_off in range(0, C, BLOCK_C):
rc = c_off + tl.arange(0, BLOCK_C)
c_mask = rc < C
# Weight matrix for to_out is [C, D]
w = tl.load(W_ptr + rc[:, None]*D + rd[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16)
res = tl.dot(xf, tl.trans(w))
tl.store(Out_ptr + pid_b*stride_ob + pid_n1*stride_on1 + n2[:, None]*stride_on2 + rc[None, :]*stride_oc, res.to(tl.float32), mask=mask_n2[:, None] & c_mask[None, :])
def custom_kernel(data):
x, mask, weights, config = data
B, N, _, C = x.shape
D = config["hidden_dim"]
device = x.device
# Ensure weights are contiguous and in FP16
w_5 = torch.stack([weights[k] for k in ["left_proj.weight", "right_proj.weight", "left_gate.weight", "right_gate.weight", "out_gate.weight"]]).to(device, torch.float16).contiguous()
to_out_w = weights["to_out.weight"].to(device, torch.float16).contiguous()
l = torch.empty((B, D, N, N), device=device, dtype=torch.float16)
r = torch.empty((B, D, N, N), device=device, dtype=torch.float16)
og = torch.empty((B, N, N, D), device=device, dtype=torch.float16)
# Frontend: BLOCK_N=64, BLOCK_C=64, 8 warps, 2 stages
fused_frontend_v108[(B, N, (N+64-1)//64)](x, mask, w_5, weights["norm.weight"], weights["norm.bias"], l, r, og, B, N, C, D, *x.stride(), *mask.stride(), 1e-5, 64, 64, num_warps=8, num_stages=2)
# BMM: [B*D, N, N] @ [B*D, N, N] -> [B*D, N, N]
bmm_out = torch.bmm(l.view(-1, N, N), r.view(-1, N, N).transpose(-1, -2))
out = torch.empty((B, N, N, C), device=device, dtype=torch.float32)
# Backend: BLOCK_N=128, BLOCK_C=64, 8 warps, 2 stages
fused_backend_v108[(B, N, (N+128-1)//128)](bmm_out, og, weights["to_out_norm.weight"], weights["to_out_norm.bias"], to_out_w, out, B, N, D, C, *out.stride(), 1e-5, 128, 64, num_warps=8, num_stages=2)
return out
scrolls · 141 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 408231.
⋯ 3 unchanged lines@triton.jit- def fused_frontend_v102(+ def fused_frontend_v108(X_ptr, Mask_ptr,W_ptr, NW_ptr, NB_ptr,L_ptr, R_ptr, OG_ptr,⋯ 8 unchanged linesn2 = pid_n2_start + tl.arange(0, BLOCK_N)mask_n2 = n2 < N- # Vectorized Statistics Pass (Single-pass to reduce global memory reads)+ # Single-pass LayerNorm statistics to minimize global memory readss1 = tl.zeros([BLOCK_N], dtype=tl.float32)s2 = tl.zeros([BLOCK_N], dtype=tl.float32)-for c_off in range(0, C, BLOCK_C):rc = c_off + tl.arange(0, BLOCK_C)c_mask = rc < C⋯ 12 unchanged lineslg_acc = tl.zeros([BLOCK_N, BLOCK_D], dtype=tl.float32)rg_acc = tl.zeros([BLOCK_N, BLOCK_D], dtype=tl.float32)og_acc = tl.zeros([BLOCK_N, BLOCK_D], dtype=tl.float32)-rd = tl.arange(0, BLOCK_D)- w_stride_type = D * C-++ # Pre-calculate weight offsets to reduce loop overhead+ w_stride = BLOCK_D * C+for c_off in range(0, C, BLOCK_C):rc = c_off + tl.arange(0, BLOCK_C)c_mask = rc < C⋯ 3 unchanged linesxn = ((x - mean) * rstd * nw[None, :] + nb[None, :]).to(tl.float16)w_off = rd[None, :] * C + rc[:, None]- l_acc += tl.dot(xn, tl.load(W_ptr + 0*w_stride_type + w_off, mask=c_mask[:, None], other=0.0).to(tl.float16))- r_acc += tl.dot(xn, tl.load(W_ptr + 1*w_stride_type + w_off, mask=c_mask[:, None], other=0.0).to(tl.float16))- lg_acc += tl.dot(xn, tl.load(W_ptr + 2*w_stride_type + w_off, mask=c_mask[:, None], other=0.0).to(tl.float16))- rg_acc += tl.dot(xn, tl.load(W_ptr + 3*w_stride_type + w_off, mask=c_mask[:, None], other=0.0).to(tl.float16))- og_acc += tl.dot(xn, tl.load(W_ptr + 4*w_stride_type + w_off, mask=c_mask[:, None], other=0.0).to(tl.float16))+ # Batch load weights if possible or keep separate to manage register pressure+ l_acc += tl.dot(xn, tl.load(W_ptr + 0*w_stride + w_off, mask=c_mask[:, None], other=0.0).to(tl.float16))+ r_acc += tl.dot(xn, tl.load(W_ptr + 1*w_stride + w_off, mask=c_mask[:, None], other=0.0).to(tl.float16))+ lg_acc += tl.dot(xn, tl.load(W_ptr + 2*w_stride + w_off, mask=c_mask[:, None], other=0.0).to(tl.float16))+ rg_acc += tl.dot(xn, tl.load(W_ptr + 3*w_stride + w_off, mask=c_mask[:, None], other=0.0).to(tl.float16))+ og_acc += tl.dot(xn, tl.load(W_ptr + 4*w_stride + w_off, mask=c_mask[:, None], other=0.0).to(tl.float16))m = tl.load(Mask_ptr + pid_b*stride_mb + pid_n1*stride_mn1 + n2, mask=mask_n2, other=0.0)[:, None]l = (l_acc * m * tl.sigmoid(lg_acc)).to(tl.float16)r = (r_acc * m * tl.sigmoid(rg_acc)).to(tl.float16)og = tl.sigmoid(og_acc).to(tl.float16)- # Store L/R in (B, D, N, N) layout for torch.bmm- l_base = pid_b * D * N * N + rd[None, :] * N * N + pid_n1 * N + n2[:, None]+ # Store L/R in [B, D, N, N] layout for BMM (strided write)+ l_base = pid_b * (D*N*N) + rd[None, :] * (N*N) + pid_n1 * N + n2[:, None]tl.store(L_ptr + l_base, l, mask=mask_n2[:, None])tl.store(R_ptr + l_base, r, mask=mask_n2[:, None])- # Store OG in (B, N, N, D) layout for backend+ # Store OG in [B, N, N, D] layout for backend (coalesced write)og_off = pid_b*(N*N*D) + pid_n1*(N*D) + n2[:, None]*D + rd[None, :]tl.store(OG_ptr + og_off, og, mask=mask_n2[:, None])@triton.jit- def fused_backend_v102(+ def fused_backend_v108(BMM_ptr, OG_ptr, NW_ptr, NB_ptr, W_ptr, Out_ptr,B, N, D: tl.constexpr, C,stride_ob, stride_on1, stride_on2, stride_oc,⋯ 5 unchanged linesmask_n2 = n2 < Nrd = tl.arange(0, D)+ # Load from [B, D, N, N] (BMM output layout)+ # Optimization: Load entire D dimension into registers for LayerNormx = tl.load(BMM_ptr + pid_b*(D*N*N) + rd[:, None]*N*N + pid_n1*N + n2[None, :], mask=mask_n2[None, :], other=0.0).to(tl.float32)mean = (tl.sum(x, axis=0) / D)[None, :]diff = x - mean⋯ 4 unchanged linesnb = tl.load(NB_ptr + rd)[:, None]xn = (diff * rstd * nw + nb).to(tl.float16)+ # Load OG from [B, N, N, D] (coalesced)og = tl.load(OG_ptr + pid_b*(N*N*D) + pid_n1*(N*D) + n2[:, None]*D + rd[None, :], mask=mask_n2[:, None], other=0.0).to(tl.float16)xf = (tl.trans(xn) * og).to(tl.float16)for c_off in range(0, C, BLOCK_C):rc = c_off + tl.arange(0, BLOCK_C)c_mask = rc < C+ # Weight matrix for to_out is [C, D]w = tl.load(W_ptr + rc[:, None]*D + rd[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16)res = tl.dot(xf, tl.trans(w))tl.store(Out_ptr + pid_b*stride_ob + pid_n1*stride_on1 + n2[:, None]*stride_on2 + rc[None, :]*stride_oc, res.to(tl.float32), mask=mask_n2[:, None] & c_mask[None, :])⋯ 5 unchanged linesD = config["hidden_dim"]device = x.device- w_5 = torch.stack([- weights["left_proj.weight"], weights["right_proj.weight"],- weights["left_gate.weight"], weights["right_gate.weight"], weights["out_gate.weight"]- ]).to(device, torch.float16).contiguous()+ # Ensure weights are contiguous and in FP16+ w_5 = torch.stack([weights[k] for k in ["left_proj.weight", "right_proj.weight", "left_gate.weight", "right_gate.weight", "out_gate.weight"]]).to(device, torch.float16).contiguous()to_out_w = weights["to_out.weight"].to(device, torch.float16).contiguous()l = torch.empty((B, D, N, N), device=device, dtype=torch.float16)r = torch.empty((B, D, N, N), device=device, dtype=torch.float16)og = torch.empty((B, N, N, D), device=device, dtype=torch.float16)- grid_f = (B, N, (N + 64 - 1) // 64)- # Reverting num_stages to 2 to avoid shared memory overflow on H100- fused_frontend_v102[grid_f](x, mask, w_5, weights["norm.weight"], weights["norm.bias"], l, r, og, B, N, C, D, *x.stride(), *mask.stride(), 1e-5, BLOCK_N=64, BLOCK_C=64, num_warps=8, num_stages=2)+ # Frontend: BLOCK_N=64, BLOCK_C=64, 8 warps, 2 stages+ fused_frontend_v108[(B, N, (N+64-1)//64)](x, mask, w_5, weights["norm.weight"], weights["norm.bias"], l, r, og, B, N, C, D, *x.stride(), *mask.stride(), 1e-5, 64, 64, num_warps=8, num_stages=2)+ # BMM: [B*D, N, N] @ [B*D, N, N] -> [B*D, N, N]bmm_out = torch.bmm(l.view(-1, N, N), r.view(-1, N, N).transpose(-1, -2))out = torch.empty((B, N, N, C), device=device, dtype=torch.float32)- grid_b = (B, N, (N + 128 - 1) // 128)- fused_backend_v102[grid_b](bmm_out, og, weights["to_out_norm.weight"], weights["to_out_norm.bias"], to_out_w, out, B, N, D, C, *out.stride(), 1e-5, BLOCK_N=128, BLOCK_C=64, num_warps=8, num_stages=2)+ # Backend: BLOCK_N=128, BLOCK_C=64, 8 warps, 2 stages+ fused_backend_v108[(B, N, (N+128-1)//128)](bmm_out, og, weights["to_out_norm.weight"], weights["to_out_norm.bias"], to_out_w, out, B, N, D, C, *out.stride(), 1e-5, 128, 64, num_warps=8, num_stages=2)return out
scrolls · 131 diff lines total
Best evidence level for this revision: reported
JSON