submission 407551
Zeyu Shen · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 163 lines, June 9 Researcher Reciprocity License v1.0.
fused_prologue_epilogue.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-407551?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:46a76860cb928c2b0ed40757ed44364ceb4f8af064b565a894d3bdf5a3387212
license declaredunknown
license concludedunknown
authorsZeyu Shen
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fused-epilogue
def _fused_epilogue_kernel(mma
acc_l += tl.dot(x_n, tl.load(W_L_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16))stages = 1
B, N, C, D, BLOCK_SIZE_N=32, BLOCK_SIZE_C=32, num_stages=1tile-n = 32
B, N, C, D, BLOCK_SIZE_N=32, BLOCK_SIZE_C=32, num_stages=1Kernel source
fused_prologue_epilogue.py163 lines
import torch
import triton
import triton.language as tl
@triton.jit
def _fused_prologue_kernel(
X_ptr, M_ptr, L_ptr, R_ptr, OG_ptr,
W_norm_ptr, B_norm_ptr,
W_L_ptr, W_R_ptr, W_LG_ptr, W_RG_ptr, W_OG_ptr,
stride_xb, stride_xi, stride_xj, stride_xc,
stride_mb, stride_mi, stride_mj,
B, N, C, D: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_C: tl.constexpr
):
pid_b = tl.program_id(0)
pid_i = tl.program_id(1)
pid_j_start = tl.program_id(2) * BLOCK_SIZE_N
offsets_j = pid_j_start + tl.arange(0, BLOCK_SIZE_N)
mask_j = offsets_j < N
# 1. LayerNorm statistics
acc_sum = tl.zeros([BLOCK_SIZE_N], dtype=tl.float32)
acc_sum_sq = tl.zeros([BLOCK_SIZE_N], dtype=tl.float32)
for c_offset in range(0, C, BLOCK_SIZE_C):
cols = c_offset + tl.arange(0, BLOCK_SIZE_C)
c_mask = cols < C
x_ptr = X_ptr + pid_b * stride_xb + pid_i * stride_xi + offsets_j[:, None] * stride_xj + cols[None, :]
x_chunk = tl.load(x_ptr, mask=(mask_j[:, None] & c_mask[None, :]), other=0.0).to(tl.float32)
acc_sum += tl.sum(x_chunk, axis=1)
acc_sum_sq += tl.sum(x_chunk * x_chunk, axis=1)
mean = acc_sum / C
var = (acc_sum_sq / C) - (mean * mean)
rstd = 1.0 / tl.sqrt(var + 1e-5)
# 2. Projections
acc_l = tl.zeros([BLOCK_SIZE_N, D], dtype=tl.float32)
acc_lg = tl.zeros([BLOCK_SIZE_N, D], dtype=tl.float32)
acc_r = tl.zeros([BLOCK_SIZE_N, D], dtype=tl.float32)
acc_rg = tl.zeros([BLOCK_SIZE_N, D], dtype=tl.float32)
acc_og = tl.zeros([BLOCK_SIZE_N, D], dtype=tl.float32)
off_d = tl.arange(0, D)
for c_offset in range(0, C, BLOCK_SIZE_C):
cols = c_offset + tl.arange(0, BLOCK_SIZE_C)
c_mask = cols < C
x_ptr = X_ptr + pid_b * stride_xb + pid_i * stride_xi + offsets_j[:, None] * stride_xj + cols[None, :]
x_chunk = tl.load(x_ptr, mask=(mask_j[:, None] & c_mask[None, :]), other=0.0).to(tl.float32)
w_n = tl.load(W_norm_ptr + cols, mask=c_mask, other=0.0)
b_n = tl.load(B_norm_ptr + cols, mask=c_mask, other=0.0)
x_n = (x_chunk - mean[:, None]) * rstd[:, None] * w_n[None, :] + b_n[None, :]
x_n = x_n.to(tl.float16)
# Load weights and perform dots
acc_l += tl.dot(x_n, tl.load(W_L_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16))
acc_lg += tl.dot(x_n, tl.load(W_LG_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16))
acc_r += tl.dot(x_n, tl.load(W_R_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16))
acc_rg += tl.dot(x_n, tl.load(W_RG_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16))
acc_og += tl.dot(x_n, tl.load(W_OG_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16))
m_ptr = M_ptr + pid_b * stride_mb + pid_i * stride_mi + offsets_j
mask_val = tl.load(m_ptr, mask=mask_j, other=0.0).to(tl.float32)
l_final = acc_l * tl.sigmoid(acc_lg) * mask_val[:, None]
r_final = acc_r * tl.sigmoid(acc_rg) * mask_val[:, None]
og_final = tl.sigmoid(acc_og)
# Store L, R in [B, D, N, N] for BMM, OG in [B, N, N, D]
idx_nn = pid_i * N + offsets_j
off_l_r = pid_b * D * N * N + off_d[None, :] * N * N + idx_nn[:, None]
tl.store(L_ptr + off_l_r, l_final.to(tl.float16), mask=(mask_j[:, None] & (off_d[None, :] < D)))
tl.store(R_ptr + off_l_r, r_final.to(tl.float16), mask=(mask_j[:, None] & (off_d[None, :] < D)))
off_og = pid_b * N * N * D + idx_nn[:, None] * D + off_d[None, :]
tl.store(OG_ptr + off_og, og_final.to(tl.float16), mask=(mask_j[:, None] & (off_d[None, :] < D)))
@triton.jit
def _fused_epilogue_kernel(
BMM_OUT_ptr, OG_ptr, OUT_ptr,
W_TN_ptr, B_TN_ptr,
stride_bmm_b, stride_bmm_d, stride_bmm_i, stride_bmm_j,
stride_og_b, stride_og_i, stride_og_j, stride_og_d,
stride_out_b, stride_out_i, stride_out_j, stride_out_d,
B, N, D: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr
):
pid_b = tl.program_id(0)
pid_i = tl.program_id(1)
pid_j_start = tl.program_id(2) * BLOCK_SIZE_N
offsets_j = pid_j_start + tl.arange(0, BLOCK_SIZE_N)
mask_j = offsets_j < N
off_d = tl.arange(0, D)
# Load BMM output [BLOCK_SIZE_N, D] from [B, D, N, N]
bmm_ptr = BMM_OUT_ptr + pid_b * stride_bmm_b + off_d[None, :] * stride_bmm_d + pid_i * stride_bmm_i + offsets_j[:, None] * stride_bmm_j
val = tl.load(bmm_ptr, mask=(mask_j[:, None] & (off_d[None, :] < D)), other=0.0).to(tl.float32)
# Load OG [BLOCK_SIZE_N, D] from [B, N, N, D]
og_ptr = OG_ptr + pid_b * stride_og_b + pid_i * stride_og_i + offsets_j[:, None] * stride_og_j + off_d[None, :]
og = tl.load(og_ptr, mask=(mask_j[:, None] & (off_d[None, :] < D)), other=0.0).to(tl.float32)
# LayerNorm over D
mean = tl.sum(val, axis=1) / D
var = (tl.sum(val * val, axis=1) / D) - (mean * mean)
rstd = 1.0 / tl.sqrt(var + 1e-5)
w_tn = tl.load(W_TN_ptr + off_d)
b_tn = tl.load(B_TN_ptr + off_d)
val = (val - mean[:, None]) * rstd[:, None] * w_tn[None, :] + b_tn[None, :]
val = val * og
# Store in [B, N, N, D] format for final matmul
out_ptr = OUT_ptr + pid_b * stride_out_b + pid_i * stride_out_i + offsets_j[:, None] * stride_out_j + off_d[None, :]
tl.store(out_ptr, val.to(tl.float16), mask=(mask_j[:, None] & (off_d[None, :] < D)))
def custom_kernel(data):
x, mask, weights, config = data
B, N, _, C = x.shape
D = config["hidden_dim"]
device = x.device
w_fp16 = {k: v.to(torch.float16) for k, v in weights.items()}
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)
# Prologue: LN + 5 Projections + Gating. Reduced BLOCK_SIZE_C and num_stages=1 to fit SRAM.
grid_pre = (B, N, (N + 32 - 1) // 32)
_fused_prologue_kernel[grid_pre](
x, mask, L, R, OG,
w_fp16["norm.weight"], w_fp16["norm.bias"],
w_fp16["left_proj.weight"].t().contiguous(), w_fp16["right_proj.weight"].t().contiguous(),
w_fp16["left_gate.weight"].t().contiguous(), w_fp16["right_gate.weight"].t().contiguous(),
w_fp16["out_gate.weight"].t().contiguous(),
x.stride(0), x.stride(1), x.stride(2), x.stride(3),
mask.stride(0), mask.stride(1), mask.stride(2),
B, N, C, D, BLOCK_SIZE_N=32, BLOCK_SIZE_C=32, num_stages=1
)
# Contraction: [B*D, N, N] @ [B*D, N, N].T -> [B, D, N, N]
bmm_out = torch.bmm(L.view(B * D, N, N), R.view(B * D, N, N).transpose(-1, -2)).view(B, D, N, N)
# Epilogue Part 1: LN + Gating. Fused to handle the [B, D, N, N] -> [B, N, N, D] layout change.
epi_inter = torch.empty((B, N, N, D), device=device, dtype=torch.float16)
grid_epi = (B, N, (N + 32 - 1) // 32)
_fused_epilogue_kernel[grid_epi](
bmm_out, OG, epi_inter,
w_fp16["to_out_norm.weight"], w_fp16["to_out_norm.bias"],
bmm_out.stride(0), bmm_out.stride(1), bmm_out.stride(2), bmm_out.stride(3),
OG.stride(0), OG.stride(1), OG.stride(2), OG.stride(3),
epi_inter.stride(0), epi_inter.stride(1), epi_inter.stride(2), epi_inter.stride(3),
B, N, D, BLOCK_SIZE_N=32
)
# Epilogue Part 2: Final Projection to C using cuBLAS.
return (epi_inter @ w_fp16["to_out.weight"].t()).to(torch.float32)
scrolls · 163 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 407546.
⋯ 2 unchanged linesimport triton.language as tl@triton.jit- def _fused_preprocess_kernel(+ def _fused_prologue_kernel(X_ptr, M_ptr, L_ptr, R_ptr, OG_ptr,W_norm_ptr, B_norm_ptr,W_L_ptr, W_R_ptr, W_LG_ptr, W_RG_ptr, W_OG_ptr,⋯ 10 unchanged linesoffsets_j = pid_j_start + tl.arange(0, BLOCK_SIZE_N)mask_j = offsets_j < N- # 1. Compute LayerNorm for the block [BLOCK_SIZE_N, C]+ # 1. LayerNorm statisticsacc_sum = tl.zeros([BLOCK_SIZE_N], dtype=tl.float32)acc_sum_sq = tl.zeros([BLOCK_SIZE_N], dtype=tl.float32)⋯ 9 unchanged linesvar = (acc_sum_sq / C) - (mean * mean)rstd = 1.0 / tl.sqrt(var + 1e-5)- # 2. Compute Projections+ # 2. Projectionsacc_l = tl.zeros([BLOCK_SIZE_N, D], dtype=tl.float32)acc_lg = tl.zeros([BLOCK_SIZE_N, D], dtype=tl.float32)acc_r = tl.zeros([BLOCK_SIZE_N, D], dtype=tl.float32)⋯ 13 unchanged linesx_n = (x_chunk - mean[:, None]) * rstd[:, None] * w_n[None, :] + b_n[None, :]x_n = x_n.to(tl.float16)- # Projection weights [C, D]- w_l = tl.load(W_L_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16)- acc_l += tl.dot(x_n, w_l)-- w_lg = tl.load(W_LG_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16)- acc_lg += tl.dot(x_n, w_lg)-- w_r = tl.load(W_R_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16)- acc_r += tl.dot(x_n, w_r)-- w_rg = tl.load(W_RG_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16)- acc_rg += tl.dot(x_n, w_rg)-- w_og = tl.load(W_OG_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16)- acc_og += tl.dot(x_n, w_og)+ # Load weights and perform dots+ acc_l += tl.dot(x_n, tl.load(W_L_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16))+ acc_lg += tl.dot(x_n, tl.load(W_LG_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16))+ acc_r += tl.dot(x_n, tl.load(W_R_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16))+ acc_rg += tl.dot(x_n, tl.load(W_RG_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16))+ acc_og += tl.dot(x_n, tl.load(W_OG_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16))- # 3. Apply Gating and Maskingm_ptr = M_ptr + pid_b * stride_mb + pid_i * stride_mi + offsets_jmask_val = tl.load(m_ptr, mask=mask_j, other=0.0).to(tl.float32)⋯ 1 unchanged linesr_final = acc_r * tl.sigmoid(acc_rg) * mask_val[:, None]og_final = tl.sigmoid(acc_og)- # 4. Store results+ # Store L, R in [B, D, N, N] for BMM, OG in [B, N, N, D]idx_nn = pid_i * N + offsets_j- # L and R stored in [B, D, N, N] for BMM efficiencyoff_l_r = pid_b * D * N * N + off_d[None, :] * N * N + idx_nn[:, None]tl.store(L_ptr + off_l_r, l_final.to(tl.float16), mask=(mask_j[:, None] & (off_d[None, :] < D)))tl.store(R_ptr + off_l_r, r_final.to(tl.float16), mask=(mask_j[:, None] & (off_d[None, :] < D)))- # OG stored in [B, N, N, D]off_og = pid_b * N * N * D + idx_nn[:, None] * D + off_d[None, :]tl.store(OG_ptr + off_og, og_final.to(tl.float16), mask=(mask_j[:, None] & (off_d[None, :] < D)))+ @triton.jit+ def _fused_epilogue_kernel(+ BMM_OUT_ptr, OG_ptr, OUT_ptr,+ W_TN_ptr, B_TN_ptr,+ stride_bmm_b, stride_bmm_d, stride_bmm_i, stride_bmm_j,+ stride_og_b, stride_og_i, stride_og_j, stride_og_d,+ stride_out_b, stride_out_i, stride_out_j, stride_out_d,+ B, N, D: tl.constexpr,+ BLOCK_SIZE_N: tl.constexpr+ ):+ pid_b = tl.program_id(0)+ pid_i = tl.program_id(1)+ pid_j_start = tl.program_id(2) * BLOCK_SIZE_N++ offsets_j = pid_j_start + tl.arange(0, BLOCK_SIZE_N)+ mask_j = offsets_j < N+ off_d = tl.arange(0, D)++ # Load BMM output [BLOCK_SIZE_N, D] from [B, D, N, N]+ bmm_ptr = BMM_OUT_ptr + pid_b * stride_bmm_b + off_d[None, :] * stride_bmm_d + pid_i * stride_bmm_i + offsets_j[:, None] * stride_bmm_j+ val = tl.load(bmm_ptr, mask=(mask_j[:, None] & (off_d[None, :] < D)), other=0.0).to(tl.float32)++ # Load OG [BLOCK_SIZE_N, D] from [B, N, N, D]+ og_ptr = OG_ptr + pid_b * stride_og_b + pid_i * stride_og_i + offsets_j[:, None] * stride_og_j + off_d[None, :]+ og = tl.load(og_ptr, mask=(mask_j[:, None] & (off_d[None, :] < D)), other=0.0).to(tl.float32)++ # LayerNorm over D+ mean = tl.sum(val, axis=1) / D+ var = (tl.sum(val * val, axis=1) / D) - (mean * mean)+ rstd = 1.0 / tl.sqrt(var + 1e-5)++ w_tn = tl.load(W_TN_ptr + off_d)+ b_tn = tl.load(B_TN_ptr + off_d)+ val = (val - mean[:, None]) * rstd[:, None] * w_tn[None, :] + b_tn[None, :]+ val = val * og++ # Store in [B, N, N, D] format for final matmul+ out_ptr = OUT_ptr + pid_b * stride_out_b + pid_i * stride_out_i + offsets_j[:, None] * stride_out_j + off_d[None, :]+ tl.store(out_ptr, val.to(tl.float16), mask=(mask_j[:, None] & (off_d[None, :] < D)))+def custom_kernel(data):x, mask, weights, config = dataB, N, _, C = x.shapeD = config["hidden_dim"]device = x.device-- # Cast weights to FP16 for the fused kernel to avoid dtype mismatch in tl.dot- # and for the final linear layer.w_fp16 = {k: v.to(torch.float16) for k, v in weights.items()}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)- BLOCK_SIZE_N = 32- BLOCK_SIZE_C = 128-- grid = (B, N, (N + BLOCK_SIZE_N - 1) // BLOCK_SIZE_N)- _fused_preprocess_kernel[grid](+ # Prologue: LN + 5 Projections + Gating. Reduced BLOCK_SIZE_C and num_stages=1 to fit SRAM.+ grid_pre = (B, N, (N + 32 - 1) // 32)+ _fused_prologue_kernel[grid_pre](x, mask, L, R, OG,w_fp16["norm.weight"], w_fp16["norm.bias"],w_fp16["left_proj.weight"].t().contiguous(), w_fp16["right_proj.weight"].t().contiguous(),⋯ 1 unchanged linesw_fp16["out_gate.weight"].t().contiguous(),x.stride(0), x.stride(1), x.stride(2), x.stride(3),mask.stride(0), mask.stride(1), mask.stride(2),- B, N, C, D,- BLOCK_SIZE_N=BLOCK_SIZE_N,- BLOCK_SIZE_C=BLOCK_SIZE_C,- num_stages=2+ B, N, C, D, BLOCK_SIZE_N=32, BLOCK_SIZE_C=32, num_stages=1)- # Contraction: einsum('... i k d, ... j k d -> ... i j d')- # L: [B, D, N, K], R: [B, D, N, K]. We want [B, D, N, N] where out[b, d, i, j] = sum_k L[b, d, i, k] * R[b, d, j, k]- # This is L @ R.T- bmm_out = torch.bmm(L.view(B * D, N, N), R.view(B * D, N, N).transpose(-1, -2))- bmm_out = bmm_out.view(B, D, N, N).permute(0, 2, 3, 1) # [B, N, N, D]-- # Final layers: LayerNorm -> Gating -> Linear- # Use FP32 for LayerNorm stability if needed, but here we use FP16 for speed- out = torch.nn.functional.layer_norm(bmm_out, (D,), w_fp16["to_out_norm.weight"], w_fp16["to_out_norm.bias"])- out = out * OG+ # Contraction: [B*D, N, N] @ [B*D, N, N].T -> [B, D, N, N]+ bmm_out = torch.bmm(L.view(B * D, N, N), R.view(B * D, N, N).transpose(-1, -2)).view(B, D, N, N)- # Final projection to C- return (out @ w_fp16["to_out.weight"].t()).to(torch.float32)+ # Epilogue Part 1: LN + Gating. Fused to handle the [B, D, N, N] -> [B, N, N, D] layout change.+ epi_inter = torch.empty((B, N, N, D), device=device, dtype=torch.float16)+ grid_epi = (B, N, (N + 32 - 1) // 32)+ _fused_epilogue_kernel[grid_epi](+ bmm_out, OG, epi_inter,+ w_fp16["to_out_norm.weight"], w_fp16["to_out_norm.bias"],+ bmm_out.stride(0), bmm_out.stride(1), bmm_out.stride(2), bmm_out.stride(3),+ OG.stride(0), OG.stride(1), OG.stride(2), OG.stride(3),+ epi_inter.stride(0), epi_inter.stride(1), epi_inter.stride(2), epi_inter.stride(3),+ B, N, D, BLOCK_SIZE_N=32+ )++ # Epilogue Part 2: Final Projection to C using cuBLAS.+ return (epi_inter @ w_fp16["to_out.weight"].t()).to(torch.float32)
scrolls · 178 diff lines total
Best evidence level for this revision: reported
JSON