submission 408278
Zeyu Shen · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 198 lines, June 9 Researcher Reciprocity License v1.0.
modular_kernels_v11.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-408278?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:e77645cfbcac8766252c44de619bed09628731b3674b779149199cc76f80ab28
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
acc_l += tl.dot(x, tl.trans(w_l))num-warps = 4
B, N, C, BLOCK_N=32, BLOCK_C=128, num_warps=4stages = 2
B=B, N=N, C=C, D=D, BLOCK_N=BLOCK_N_PROJ, BLOCK_C=64, num_warps=8, num_stages=2tile-n = 32
B, N, C, BLOCK_N=32, BLOCK_C=128, num_warps=4Kernel source
modular_kernels_v11.py198 lines
import torch
import triton
import triton.language as tl
@triton.jit
def layernorm_kernel_v11(
X, LN_W, LN_B, Out,
stride_xb, stride_xn1, stride_xn2, stride_xc,
B, N, C,
BLOCK_N: tl.constexpr, BLOCK_C: tl.constexpr
):
pid_b, pid_n1 = tl.program_id(0), tl.program_id(1)
pid_n2_start = tl.program_id(2) * BLOCK_N
offs_n2 = pid_n2_start + tl.arange(0, BLOCK_N)
mask_n2 = offs_n2 < N
m1 = tl.zeros([BLOCK_N], dtype=tl.float32)
m2 = tl.zeros([BLOCK_N], dtype=tl.float32)
for c_start in range(0, C, BLOCK_C):
offs_c = c_start + tl.arange(0, BLOCK_C)
mask_c = offs_c < C
x_ptr = X + pid_b * stride_xb + pid_n1 * stride_xn1 + offs_n2[:, None] * stride_xn2 + offs_c[None, :]
x = tl.load(x_ptr, mask=(mask_n2[:, None] & mask_c[None, :]), other=0.0).to(tl.float32)
m1 += tl.sum(x, axis=1)
m2 += tl.sum(x * x, axis=1)
mean = m1 / C
var = tl.maximum(0.0, (m2 / C) - (mean * mean))
rstd = 1.0 / tl.sqrt(var + 1e-5)
for c_start in range(0, C, BLOCK_C):
offs_c = c_start + tl.arange(0, BLOCK_C)
mask_c = offs_c < C
x_ptr = X + pid_b * stride_xb + pid_n1 * stride_xn1 + offs_n2[:, None] * stride_xn2 + offs_c[None, :]
x = tl.load(x_ptr, mask=(mask_n2[:, None] & mask_c[None, :]), other=0.0).to(tl.float32)
ln_w = tl.load(LN_W + offs_c, mask=mask_c, other=0.0)
ln_b = tl.load(LN_B + offs_c, mask=mask_c, other=0.0)
x_hat = (x - mean[:, None]) * rstd[:, None] * ln_w[None, :] + ln_b[None, :]
out_ptr = Out + pid_b * stride_xb + pid_n1 * stride_xn1 + offs_n2[:, None] * stride_xn2 + offs_c[None, :]
tl.store(out_ptr, x_hat.to(tl.float16), mask=(mask_n2[:, None] & mask_c[None, :]))
@triton.jit
def projection_kernel_v7(
X_norm, Mask, W_concat,
L_out, R_out, OG_out,
stride_xb, stride_xn1, stride_xn2, stride_xc,
stride_mb, stride_mn1, stride_mn2,
stride_ob, stride_od, stride_on1, stride_on2,
B, N, C, D: tl.constexpr,
BLOCK_N: tl.constexpr, BLOCK_C: tl.constexpr
):
pid_b, pid_n1 = tl.program_id(0), tl.program_id(1)
pid_n2_start = tl.program_id(2) * BLOCK_N
offs_n2 = pid_n2_start + tl.arange(0, BLOCK_N)
mask_n2 = offs_n2 < N
acc_l = tl.zeros([BLOCK_N, D], dtype=tl.float32)
acc_r = tl.zeros([BLOCK_N, D], dtype=tl.float32)
acc_lg = tl.zeros([BLOCK_N, D], dtype=tl.float32)
acc_rg = tl.zeros([BLOCK_N, D], dtype=tl.float32)
acc_og = tl.zeros([BLOCK_N, D], dtype=tl.float32)
offs_d = tl.arange(0, D)
# Use block pointers for X_norm to improve memory access efficiency
x_block_ptr = tl.make_block_ptr(
base=X_norm + pid_b * stride_xb + pid_n1 * stride_xn1,
shape=(N, C),
strides=(stride_xn2, stride_xc),
offsets=(pid_n2_start, 0),
block_shape=(BLOCK_N, BLOCK_C),
order=(1, 0)
)
for c_start in range(0, C, BLOCK_C):
x = tl.load(x_block_ptr, boundary_check=(0, 1)).to(tl.float16)
offs_c = c_start + tl.arange(0, BLOCK_C)
mask_c = offs_c < C
# Weights are [5*D, C]. We load chunks of [D, BLOCK_C] and transpose for dot.
# This is more efficient than standard pointer indexing.
w_l = tl.load(W_concat + (0*D + offs_d)[:, None] * C + offs_c[None, :], mask=mask_c[None, :], other=0.0).to(tl.float16)
w_r = tl.load(W_concat + (1*D + offs_d)[:, None] * C + offs_c[None, :], mask=mask_c[None, :], other=0.0).to(tl.float16)
w_lg = tl.load(W_concat + (2*D + offs_d)[:, None] * C + offs_c[None, :], mask=mask_c[None, :], other=0.0).to(tl.float16)
w_rg = tl.load(W_concat + (3*D + offs_d)[:, None] * C + offs_c[None, :], mask=mask_c[None, :], other=0.0).to(tl.float16)
w_og = tl.load(W_concat + (4*D + offs_d)[:, None] * C + offs_c[None, :], mask=mask_c[None, :], other=0.0).to(tl.float16)
acc_l += tl.dot(x, tl.trans(w_l))
acc_r += tl.dot(x, tl.trans(w_r))
acc_lg += tl.dot(x, tl.trans(w_lg))
acc_rg += tl.dot(x, tl.trans(w_rg))
acc_og += tl.dot(x, tl.trans(w_og))
x_block_ptr = tl.advance(x_block_ptr, [0, BLOCK_C])
m_ptr = Mask + pid_b * stride_mb + pid_n1 * stride_mn1 + offs_n2 * stride_mn2
mask_val = tl.load(m_ptr, mask=mask_n2, other=0.0).to(tl.float32)
l_final = acc_l * (mask_val[:, None] * tl.sigmoid(acc_lg))
r_final = acc_r * (mask_val[:, None] * tl.sigmoid(acc_rg))
og_final = tl.sigmoid(acc_og)
out_off = pid_b * stride_ob + offs_d[:, None] * stride_od + pid_n1 * stride_on1 + offs_n2[None, :] * stride_on2
tl.store(L_out + out_off, tl.trans(l_final).to(tl.float16), mask=mask_n2[None, :])
tl.store(R_out + out_off, tl.trans(r_final).to(tl.float16), mask=mask_n2[None, :])
tl.store(OG_out + out_off, tl.trans(og_final).to(tl.float16), mask=mask_n2[None, :])
@triton.jit
def post_process_kernel_v62(
matmul_out, OG_p, LN_W, LN_B, TO_OUT_W_T, Out,
stride_mb, stride_md, stride_mn, stride_mm,
stride_ob, stride_on1, stride_on2, stride_oc,
B, N, D: tl.constexpr, C: tl.constexpr,
BLOCK_N: tl.constexpr, BLOCK_C: tl.constexpr
):
pid_b, pid_n1 = tl.program_id(0), tl.program_id(1)
pid_n2_start = tl.program_id(2) * BLOCK_N
offs_n2 = pid_n2_start + tl.arange(0, BLOCK_N)
mask_n2 = offs_n2 < N
offs_d = tl.arange(0, D)
m_off = pid_b * stride_mb + offs_d[None, :] * stride_md + pid_n1 * stride_mn + offs_n2[:, None] * stride_mm
x = tl.load(matmul_out + m_off, mask=mask_n2[:, None], other=0.0).to(tl.float32)
og = tl.load(OG_p + m_off, mask=mask_n2[:, None], other=0.0).to(tl.float32)
mean = tl.sum(x, axis=1) / D
diff = x - mean[:, None]
var = tl.sum(diff * diff, axis=1) / D
rstd = 1.0 / tl.sqrt(var + 1e-5)
ln_w = tl.load(LN_W + offs_d)
ln_b = tl.load(LN_B + offs_d)
x_norm = ((diff * rstd[:, None]) * ln_w[None, :] + ln_b[None, :]) * og
x_norm_f16 = x_norm.to(tl.float16)
# Using pre-transposed weights for the final projection
for c_start in range(0, C, BLOCK_C):
offs_c = c_start + tl.arange(0, BLOCK_C)
mask_c = offs_c < C
# TO_OUT_W_T is [D, C]
w = tl.load(TO_OUT_W_T + offs_d[:, None] * C + offs_c[None, :], mask=mask_c[None, :], other=0.0).to(tl.float16)
out_chunk = tl.dot(x_norm_f16, w)
tl.store(Out + pid_b * stride_ob + pid_n1 * stride_on1 + offs_n2[:, None] * stride_on2 + offs_c[None, :] * stride_oc, out_chunk.to(tl.float32), mask=(mask_n2[:, None] & mask_c[None, :]))
def custom_kernel(data):
x, mask, weights, config = data
B, N, _, C = x.shape
D = config["hidden_dim"]
device = x.device
x_norm = torch.empty_like(x, dtype=torch.float16)
layernorm_kernel_v11[(B, N, (N + 31) // 32)](
x, weights["norm.weight"], weights["norm.bias"], x_norm,
x.stride(0), x.stride(1), x.stride(2), x.stride(3),
B, N, C, BLOCK_N=32, BLOCK_C=128, num_warps=4
)
W_concat = torch.cat([
weights["left_proj.weight"], weights["right_proj.weight"],
weights["left_gate.weight"], weights["right_gate.weight"], weights["out_gate.weight"]
], dim=0).to(device=device, dtype=torch.float16)
L_p = torch.empty((B, D, N, N), device=device, dtype=torch.float16)
R_p = torch.empty((B, D, N, N), device=device, dtype=torch.float16)
OG_p = torch.empty((B, D, N, N), device=device, dtype=torch.float16)
BLOCK_N_PROJ = 64
# Use num_stages=2 to avoid OutOfResources error
projection_kernel_v7[(B, N, (N + BLOCK_N_PROJ - 1) // BLOCK_N_PROJ)](
x_norm, mask, W_concat,
L_out=L_p, R_out=R_p, OG_out=OG_p,
stride_xb=x_norm.stride(0), stride_xn1=x_norm.stride(1), stride_xn2=x_norm.stride(2), stride_xc=x_norm.stride(3),
stride_mb=mask.stride(0), stride_mn1=mask.stride(1), stride_mn2=mask.stride(2),
stride_ob=L_p.stride(0), stride_od=L_p.stride(1), stride_on1=L_p.stride(2), stride_on2=L_p.stride(3),
B=B, N=N, C=C, D=D, BLOCK_N=BLOCK_N_PROJ, BLOCK_C=64, num_warps=8, num_stages=2
)
matmul_out = torch.matmul(L_p, R_p.transpose(-1, -2))
# Pre-transpose final weight for better loading in kernel
to_out_w_t = weights["to_out.weight"].t().contiguous().to(device=device, dtype=torch.float16)
out = torch.empty((B, N, N, C), device=device, dtype=torch.float32)
BLOCK_N_POST = 64
post_process_kernel_v62[(B, N, (N + BLOCK_N_POST - 1) // BLOCK_N_POST)](
matmul_out, OG_p, weights["to_out_norm.weight"], weights["to_out_norm.bias"], to_out_w_t, out,
matmul_out.stride(0), matmul_out.stride(1), matmul_out.stride(2), matmul_out.stride(3),
out.stride(0), out.stride(1), out.stride(2), out.stride(3),
B, N, D, C, BLOCK_N=BLOCK_N_POST, BLOCK_C=64, num_warps=8
)
return out
scrolls · 198 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 408244.
⋯ 3 unchanged lines@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,+ def layernorm_kernel_v11(+ X, LN_W, LN_B, Out,stride_xb, stride_xn1, stride_xn2, stride_xc,- stride_mb, stride_mn1, stride_mn2,- eps, BLOCK_N: tl.constexpr, BLOCK_C: tl.constexpr+ B, N, C,+ BLOCK_N: tl.constexpr, BLOCK_C: tl.constexpr):- pid_b = tl.program_id(0)- pid_n1 = tl.program_id(1)+ pid_b, pid_n1 = tl.program_id(0), 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+ offs_n2 = pid_n2_start + tl.arange(0, BLOCK_N)+ mask_n2 = offs_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)+ m1 = tl.zeros([BLOCK_N], dtype=tl.float32)+ m2 = tl.zeros([BLOCK_N], dtype=tl.float32)+ for c_start in range(0, C, BLOCK_C):+ offs_c = c_start + tl.arange(0, BLOCK_C)+ mask_c = offs_c < C+ x_ptr = X + pid_b * stride_xb + pid_n1 * stride_xn1 + offs_n2[:, None] * stride_xn2 + offs_c[None, :]+ x = tl.load(x_ptr, mask=(mask_n2[:, None] & mask_c[None, :]), other=0.0).to(tl.float32)+ m1 += tl.sum(x, axis=1)+ m2 += 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)+ mean = m1 / C+ var = tl.maximum(0.0, (m2 / C) - (mean * mean))+ rstd = 1.0 / tl.sqrt(var + 1e-5)- # 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)+ for c_start in range(0, C, BLOCK_C):+ offs_c = c_start + tl.arange(0, BLOCK_C)+ mask_c = offs_c < C+ x_ptr = X + pid_b * stride_xb + pid_n1 * stride_xn1 + offs_n2[:, None] * stride_xn2 + offs_c[None, :]+ x = tl.load(x_ptr, mask=(mask_n2[:, None] & mask_c[None, :]), other=0.0).to(tl.float32)+ ln_w = tl.load(LN_W + offs_c, mask=mask_c, other=0.0)+ ln_b = tl.load(LN_B + offs_c, mask=mask_c, other=0.0)+ x_hat = (x - mean[:, None]) * rstd[:, None] * ln_w[None, :] + ln_b[None, :]+ out_ptr = Out + pid_b * stride_xb + pid_n1 * stride_xn1 + offs_n2[:, None] * stride_xn2 + offs_c[None, :]+ tl.store(out_ptr, x_hat.to(tl.float16), mask=(mask_n2[:, None] & mask_c[None, :]))+++ @triton.jit+ def projection_kernel_v7(+ X_norm, Mask, W_concat,+ L_out, R_out, OG_out,+ stride_xb, stride_xn1, stride_xn2, stride_xc,+ stride_mb, stride_mn1, stride_mn2,+ stride_ob, stride_od, stride_on1, stride_on2,+ B, N, C, D: tl.constexpr,+ BLOCK_N: tl.constexpr, BLOCK_C: tl.constexpr+ ):+ pid_b, pid_n1 = tl.program_id(0), tl.program_id(1)+ pid_n2_start = tl.program_id(2) * BLOCK_N+ offs_n2 = pid_n2_start + tl.arange(0, BLOCK_N)+ mask_n2 = offs_n2 < N- # 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)+ acc_l = tl.zeros([BLOCK_N, D], dtype=tl.float32)+ acc_r = tl.zeros([BLOCK_N, D], dtype=tl.float32)+ acc_lg = tl.zeros([BLOCK_N, D], dtype=tl.float32)+ acc_rg = tl.zeros([BLOCK_N, D], dtype=tl.float32)+ acc_og = tl.zeros([BLOCK_N, D], dtype=tl.float32)++ offs_d = tl.arange(0, D)++ # Use block pointers for X_norm to improve memory access efficiency+ x_block_ptr = tl.make_block_ptr(+ base=X_norm + pid_b * stride_xb + pid_n1 * stride_xn1,+ shape=(N, C),+ strides=(stride_xn2, stride_xc),+ offsets=(pid_n2_start, 0),+ block_shape=(BLOCK_N, BLOCK_C),+ order=(1, 0)+ )++ for c_start in range(0, C, BLOCK_C):+ x = tl.load(x_block_ptr, boundary_check=(0, 1)).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))+ offs_c = c_start + tl.arange(0, BLOCK_C)+ mask_c = offs_c < C- 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)+ # Weights are [5*D, C]. We load chunks of [D, BLOCK_C] and transpose for dot.+ # This is more efficient than standard pointer indexing.+ w_l = tl.load(W_concat + (0*D + offs_d)[:, None] * C + offs_c[None, :], mask=mask_c[None, :], other=0.0).to(tl.float16)+ w_r = tl.load(W_concat + (1*D + offs_d)[:, None] * C + offs_c[None, :], mask=mask_c[None, :], other=0.0).to(tl.float16)+ w_lg = tl.load(W_concat + (2*D + offs_d)[:, None] * C + offs_c[None, :], mask=mask_c[None, :], other=0.0).to(tl.float16)+ w_rg = tl.load(W_concat + (3*D + offs_d)[:, None] * C + offs_c[None, :], mask=mask_c[None, :], other=0.0).to(tl.float16)+ w_og = tl.load(W_concat + (4*D + offs_d)[:, None] * C + offs_c[None, :], mask=mask_c[None, :], other=0.0).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])+ acc_l += tl.dot(x, tl.trans(w_l))+ acc_r += tl.dot(x, tl.trans(w_r))+ acc_lg += tl.dot(x, tl.trans(w_lg))+ acc_rg += tl.dot(x, tl.trans(w_rg))+ acc_og += tl.dot(x, tl.trans(w_og))++ x_block_ptr = tl.advance(x_block_ptr, [0, BLOCK_C])+ m_ptr = Mask + pid_b * stride_mb + pid_n1 * stride_mn1 + offs_n2 * stride_mn2+ mask_val = tl.load(m_ptr, mask=mask_n2, other=0.0).to(tl.float32)+ l_final = acc_l * (mask_val[:, None] * tl.sigmoid(acc_lg))+ r_final = acc_r * (mask_val[:, None] * tl.sigmoid(acc_rg))+ og_final = tl.sigmoid(acc_og)++ out_off = pid_b * stride_ob + offs_d[:, None] * stride_od + pid_n1 * stride_on1 + offs_n2[None, :] * stride_on2+ tl.store(L_out + out_off, tl.trans(l_final).to(tl.float16), mask=mask_n2[None, :])+ tl.store(R_out + out_off, tl.trans(r_final).to(tl.float16), mask=mask_n2[None, :])+ tl.store(OG_out + out_off, tl.trans(og_final).to(tl.float16), 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,+ def post_process_kernel_v62(+ matmul_out, OG_p, LN_W, LN_B, TO_OUT_W_T, Out,+ stride_mb, stride_md, stride_mn, stride_mm,stride_ob, stride_on1, stride_on2, stride_oc,- eps, BLOCK_N: tl.constexpr, BLOCK_C: tl.constexpr+ B, N, D: tl.constexpr, C: tl.constexpr,+ 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)+ pid_n2_start = tl.program_id(2) * BLOCK_N+ offs_n2 = pid_n2_start + tl.arange(0, BLOCK_N)+ mask_n2 = offs_n2 < N+ offs_d = tl.arange(0, D)++ m_off = pid_b * stride_mb + offs_d[None, :] * stride_md + pid_n1 * stride_mn + offs_n2[:, None] * stride_mm+ x = tl.load(matmul_out + m_off, mask=mask_n2[:, None], other=0.0).to(tl.float32)+ og = tl.load(OG_p + m_off, mask=mask_n2[:, None], other=0.0).to(tl.float32)++ mean = tl.sum(x, axis=1) / D+ diff = x - mean[:, None]+ var = tl.sum(diff * diff, axis=1) / D+ rstd = 1.0 / tl.sqrt(var + 1e-5)- # 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)+ ln_w = tl.load(LN_W + offs_d)+ ln_b = tl.load(LN_B + offs_d)+ x_norm = ((diff * rstd[:, None]) * ln_w[None, :] + ln_b[None, :]) * og+ x_norm_f16 = x_norm.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, :])+ # Using pre-transposed weights for the final projection+ for c_start in range(0, C, BLOCK_C):+ offs_c = c_start + tl.arange(0, BLOCK_C)+ mask_c = offs_c < C+ # TO_OUT_W_T is [D, C]+ w = tl.load(TO_OUT_W_T + offs_d[:, None] * C + offs_c[None, :], mask=mask_c[None, :], other=0.0).to(tl.float16)+ out_chunk = tl.dot(x_norm_f16, w)+ tl.store(Out + pid_b * stride_ob + pid_n1 * stride_on1 + offs_n2[:, None] * stride_on2 + offs_c[None, :] * stride_oc, out_chunk.to(tl.float32), mask=(mask_n2[:, None] & mask_c[None, :]))def custom_kernel(data):⋯ 1 unchanged linesB, N, _, C = x.shapeD = config["hidden_dim"]device = x.device++ x_norm = torch.empty_like(x, dtype=torch.float16)+ layernorm_kernel_v11[(B, N, (N + 31) // 32)](+ x, weights["norm.weight"], weights["norm.bias"], x_norm,+ x.stride(0), x.stride(1), x.stride(2), x.stride(3),+ B, N, C, BLOCK_N=32, BLOCK_C=128, num_warps=4+ )++ W_concat = torch.cat([+ weights["left_proj.weight"], weights["right_proj.weight"],+ weights["left_gate.weight"], weights["right_gate.weight"], weights["out_gate.weight"]+ ], dim=0).to(device=device, dtype=torch.float16)++ L_p = torch.empty((B, D, N, N), device=device, dtype=torch.float16)+ R_p = torch.empty((B, D, N, N), device=device, dtype=torch.float16)+ OG_p = torch.empty((B, D, N, N), device=device, dtype=torch.float16)- # 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))-+ BLOCK_N_PROJ = 64+ # Use num_stages=2 to avoid OutOfResources error+ projection_kernel_v7[(B, N, (N + BLOCK_N_PROJ - 1) // BLOCK_N_PROJ)](+ x_norm, mask, W_concat,+ L_out=L_p, R_out=R_p, OG_out=OG_p,+ stride_xb=x_norm.stride(0), stride_xn1=x_norm.stride(1), stride_xn2=x_norm.stride(2), stride_xc=x_norm.stride(3),+ stride_mb=mask.stride(0), stride_mn1=mask.stride(1), stride_mn2=mask.stride(2),+ stride_ob=L_p.stride(0), stride_od=L_p.stride(1), stride_on1=L_p.stride(2), stride_on2=L_p.stride(3),+ B=B, N=N, C=C, D=D, BLOCK_N=BLOCK_N_PROJ, BLOCK_C=64, num_warps=8, num_stages=2+ )++ matmul_out = torch.matmul(L_p, R_p.transpose(-1, -2))++ # Pre-transpose final weight for better loading in kernel+ to_out_w_t = weights["to_out.weight"].t().contiguous().to(device=device, dtype=torch.float16)+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)+ BLOCK_N_POST = 64+ post_process_kernel_v62[(B, N, (N + BLOCK_N_POST - 1) // BLOCK_N_POST)](+ matmul_out, OG_p, weights["to_out_norm.weight"], weights["to_out_norm.bias"], to_out_w_t, out,+ matmul_out.stride(0), matmul_out.stride(1), matmul_out.stride(2), matmul_out.stride(3),+ out.stride(0), out.stride(1), out.stride(2), out.stride(3),+ B, N, D, C, BLOCK_N=BLOCK_N_POST, BLOCK_C=64, num_warps=8+ )+return out
scrolls · 301 diff lines total
Best evidence level for this revision: reported
JSON