submission 407612
Zeyu Shen · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 148 lines, June 9 Researcher Reciprocity License v1.0.
fused_triton_kernel.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-407612?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:f434f500773e4ea4e346f60010d0a066baa992a55192af5e0f52afae3a0d1f01
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 _super_epilogue_v17(mma
acc_l += tl.dot(x_n, tl.load(w_base + 0*D + off_d[None, :], mask=c_mask[:, None]))num-warps = 8
B, N, C, D, 32, 32, num_warps=8Kernel source
fused_triton_kernel.py148 lines
import torch
import triton
import triton.language as tl
@triton.jit
def _fused_prologue_v17(
X_ptr, M_ptr, L_ptr, R_ptr, OG_ptr,
W_norm_ptr, B_norm_ptr, W_PACKED_ptr,
stride_xb, stride_xi, stride_xj, stride_xc,
B, N, C, D: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_C: tl.constexpr
):
pid_b, pid_i = tl.program_id(0), tl.program_id(1)
pid_j_start = tl.program_id(2) * BLOCK_SIZE_N
off_j = pid_j_start + tl.arange(0, BLOCK_SIZE_N)
mask_j = off_j < N
# One-pass LN statistics
sum_x = tl.zeros([BLOCK_SIZE_N], dtype=tl.float32)
sum_sq_x = tl.zeros([BLOCK_SIZE_N], dtype=tl.float32)
for c_off in range(0, C, BLOCK_SIZE_C):
off_c = c_off + tl.arange(0, BLOCK_SIZE_C)
c_mask = off_c < C
ptr = X_ptr + pid_b * stride_xb + pid_i * stride_xi + off_j[:, None] * stride_xj + off_c[None, :]
x = tl.load(ptr, mask=mask_j[:, None] & c_mask[None, :], other=0.0).to(tl.float32)
sum_x += tl.sum(x, axis=1)
sum_sq_x += tl.sum(x * x, axis=1)
mean = sum_x / C
var = (sum_sq_x / C) - (mean * mean)
rstd = 1.0 / tl.sqrt(var + 1e-5)
# Projections
off_d = tl.arange(0, D)
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)
for c_off in range(0, C, BLOCK_SIZE_C):
off_c = c_off + tl.arange(0, BLOCK_SIZE_C)
c_mask = off_c < C
x = tl.load(X_ptr + pid_b * stride_xb + pid_i * stride_xi + off_j[:, None] * stride_xj + off_c[None, :], mask=mask_j[:, None] & c_mask[None, :], other=0.0).to(tl.float32)
w_n = tl.load(W_norm_ptr + off_c, mask=c_mask)
b_n = tl.load(B_norm_ptr + off_c, mask=c_mask)
x_n = ((x - mean[:, None]) * rstd[:, None] * w_n[None, :] + b_n[None, :]).to(tl.float16)
w_base = W_PACKED_ptr + off_c[:, None] * (5 * D)
acc_l += tl.dot(x_n, tl.load(w_base + 0*D + off_d[None, :], mask=c_mask[:, None]))
acc_lg += tl.dot(x_n, tl.load(w_base + 1*D + off_d[None, :], mask=c_mask[:, None]))
acc_r += tl.dot(x_n, tl.load(w_base + 2*D + off_d[None, :], mask=c_mask[:, None]))
acc_rg += tl.dot(x_n, tl.load(w_base + 3*D + off_d[None, :], mask=c_mask[:, None]))
acc_og += tl.dot(x_n, tl.load(w_base + 4*D + off_d[None, :], mask=c_mask[:, None]))
mask_val = tl.load(M_ptr + pid_b * N * N + pid_i * N + off_j, mask=mask_j, other=0.0)[:, None]
l = acc_l * tl.sigmoid(acc_lg) * mask_val
r = acc_r * tl.sigmoid(acc_rg) * mask_val
og = tl.sigmoid(acc_og)
# Store L, R in [B, D, N, N] layout for BMM
base_idx = pid_b * D * N * N + off_d[None, :] * N * N + pid_i * N + off_j[:, None]
tl.store(L_ptr + base_idx, l.to(tl.float16), mask=mask_j[:, None])
tl.store(R_ptr + base_idx, r.to(tl.float16), mask=mask_j[:, None])
# Store OG in [B, N, N, D] layout
tl.store(OG_ptr + pid_b * N * N * D + pid_i * N * D + off_j[:, None] * D + off_d[None, :], og.to(tl.float16), mask=mask_j[:, None])
@triton.jit
def _super_epilogue_v17(
BMM_OUT_ptr, OG_ptr, OUT_ptr,
W_TN_ptr, B_TN_ptr, W_TO_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_c,
B, N, C, D: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_C: tl.constexpr
):
pid_b, pid_i = tl.program_id(0), tl.program_id(1)
pid_j_start = tl.program_id(2) * BLOCK_SIZE_N
off_j = pid_j_start + tl.arange(0, BLOCK_SIZE_N)
mask_j = off_j < N
off_d = tl.arange(0, D)
bmm_ptr = BMM_OUT_ptr + pid_b * stride_bmm_b + off_d[None, :] * stride_bmm_d + pid_i * stride_bmm_i + off_j[:, None] * stride_bmm_j
val = tl.load(bmm_ptr, mask=mask_j[:, None], other=0.0).to(tl.float32)
og_ptr = OG_ptr + pid_b * stride_og_b + pid_i * stride_og_i + off_j[:, None] * stride_og_j + off_d[None, :]
og = tl.load(og_ptr, mask=mask_j[:, None], other=0.0).to(tl.float32)
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).to(tl.float16)
for c_off in range(0, C, BLOCK_SIZE_C):
off_c = c_off + tl.arange(0, BLOCK_SIZE_C)
c_mask = off_c < C
w_to = tl.load(W_TO_ptr + off_c[None, :] * D + off_d[:, None], mask=c_mask[None, :], other=0.0).to(tl.float16)
out_chunk = tl.dot(val, w_to)
out_ptr = OUT_ptr + pid_b * stride_out_b + pid_i * stride_out_i + off_j[:, None] * stride_out_j + off_c[None, :]
tl.store(out_ptr, out_chunk.to(tl.float32), mask=mask_j[:, 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
w_packed = torch.cat([
weights["left_proj.weight"],
weights["left_gate.weight"],
weights["right_proj.weight"],
weights["right_gate.weight"],
weights["out_gate.weight"]
], dim=0).t().to(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)
_fused_prologue_v17[(B, N, (N+32-1)//32)](
x, mask, L, R, OG,
weights["norm.weight"], weights["norm.bias"], w_packed,
x.stride(0), x.stride(1), x.stride(2), x.stride(3),
B, N, C, D, 32, 32, num_warps=8
)
bmm_out = torch.bmm(L.view(B*D, N, N), R.view(B*D, N, N).transpose(-1, -2)).view(B, D, N, N)
output = torch.empty((B, N, N, C), device=device, dtype=torch.float32)
_super_epilogue_v17[(B, N, (N+32-1)//32)](
bmm_out, OG, output,
weights["to_out_norm.weight"], weights["to_out_norm.bias"], weights["to_out.weight"],
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),
output.stride(0), output.stride(1), output.stride(2), output.stride(3),
B, N, C, D, 32, 64, num_warps=4
)
return output
scrolls · 148 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 407575.
⋯ 1 unchanged linesimport tritonimport triton.language as tl+@triton.jit- def _fused_prologue_kernel(+ def _fused_prologue_v17(X_ptr, M_ptr, L_ptr, R_ptr, OG_ptr,- W_norm_ptr, B_norm_ptr,- W_PACKED_ptr,+ W_norm_ptr, B_norm_ptr, W_PACKED_ptr,stride_xb, stride_xi, stride_xj, stride_xc,- stride_mb, stride_mi, stride_mj,- stride_l_b, stride_l_d, stride_l_i, stride_l_j,- stride_r_b, stride_r_d, stride_r_i, stride_r_j,- stride_og_b, stride_og_i, stride_og_j, stride_og_d,B, N, C, D: tl.constexpr,- BLOCK_SIZE_N: tl.constexpr,- BLOCK_SIZE_C: 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_b, pid_i = tl.program_id(0), tl.program_id(1)pid_j_start = tl.program_id(2) * BLOCK_SIZE_N+ off_j = pid_j_start + tl.arange(0, BLOCK_SIZE_N)+ mask_j = off_j < N- offsets_j = pid_j_start + tl.arange(0, BLOCK_SIZE_N)- mask_j = offsets_j < N-- # 1. LayerNorm statistics (Online algorithm)- acc_sum = tl.zeros([BLOCK_SIZE_N], dtype=tl.float32)- acc_sum_sq = tl.zeros([BLOCK_SIZE_N], dtype=tl.float32)+ # One-pass LN statistics+ sum_x = tl.zeros([BLOCK_SIZE_N], dtype=tl.float32)+ sum_sq_x = tl.zeros([BLOCK_SIZE_N], dtype=tl.float32)+ for c_off in range(0, C, BLOCK_SIZE_C):+ off_c = c_off + tl.arange(0, BLOCK_SIZE_C)+ c_mask = off_c < C+ ptr = X_ptr + pid_b * stride_xb + pid_i * stride_xi + off_j[:, None] * stride_xj + off_c[None, :]+ x = tl.load(ptr, mask=mask_j[:, None] & c_mask[None, :], other=0.0).to(tl.float32)+ sum_x += tl.sum(x, axis=1)+ sum_sq_x += tl.sum(x * x, axis=1)- 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)+ mean = sum_x / C+ var = (sum_sq_x / C) - (mean * mean)rstd = 1.0 / tl.sqrt(var + 1e-5)- # 2. Projections with Packed Load, Separate Compute+ # Projections+ off_d = tl.arange(0, D)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++ for c_off in range(0, C, BLOCK_SIZE_C):+ off_c = c_off + tl.arange(0, BLOCK_SIZE_C)+ c_mask = off_c < C+ x = tl.load(X_ptr + pid_b * stride_xb + pid_i * stride_xi + off_j[:, None] * stride_xj + off_c[None, :], mask=mask_j[:, None] & c_mask[None, :], other=0.0).to(tl.float32)+ w_n = tl.load(W_norm_ptr + off_c, mask=c_mask)+ b_n = tl.load(B_norm_ptr + off_c, mask=c_mask)+ x_n = ((x - mean[:, None]) * rstd[:, None] * w_n[None, :] + b_n[None, :]).to(tl.float16)- 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, :]).to(tl.float16)+ w_base = W_PACKED_ptr + off_c[:, None] * (5 * D)+ acc_l += tl.dot(x_n, tl.load(w_base + 0*D + off_d[None, :], mask=c_mask[:, None]))+ acc_lg += tl.dot(x_n, tl.load(w_base + 1*D + off_d[None, :], mask=c_mask[:, None]))+ acc_r += tl.dot(x_n, tl.load(w_base + 2*D + off_d[None, :], mask=c_mask[:, None]))+ acc_rg += tl.dot(x_n, tl.load(w_base + 3*D + off_d[None, :], mask=c_mask[:, None]))+ acc_og += tl.dot(x_n, tl.load(w_base + 4*D + off_d[None, :], mask=c_mask[:, None]))- # Load all 5 weights in one contiguous block [C, 5*D]- # We use separate tl.dot to avoid register slicing errors- w_base = W_PACKED_ptr + cols[:, None] * (5 * D)- acc_l += tl.dot(x_n, tl.load(w_base + 0*D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16))- acc_lg += tl.dot(x_n, tl.load(w_base + 1*D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16))- acc_r += tl.dot(x_n, tl.load(w_base + 2*D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16))- acc_rg += tl.dot(x_n, tl.load(w_base + 3*D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16))- acc_og += tl.dot(x_n, tl.load(w_base + 4*D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16))+ mask_val = tl.load(M_ptr + pid_b * N * N + pid_i * N + off_j, mask=mask_j, other=0.0)[:, None]+ l = acc_l * tl.sigmoid(acc_lg) * mask_val+ r = acc_r * tl.sigmoid(acc_rg) * mask_val+ og = tl.sigmoid(acc_og)- # 3. Gating and Masking- 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)+ # Store L, R in [B, D, N, N] layout for BMM+ base_idx = pid_b * D * N * N + off_d[None, :] * N * N + pid_i * N + off_j[:, None]+ tl.store(L_ptr + base_idx, l.to(tl.float16), mask=mask_j[:, None])+ tl.store(R_ptr + base_idx, r.to(tl.float16), mask=mask_j[:, None])+ # Store OG in [B, N, N, D] layout+ tl.store(OG_ptr + pid_b * N * N * D + pid_i * N * D + off_j[:, None] * D + off_d[None, :], og.to(tl.float16), mask=mask_j[:, None])- 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)- # 4. Stores for BMM [B, D, N, N]- 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 * stride_og_b + pid_i * stride_og_i + offsets_j[:, None] * stride_og_j + 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 _super_epilogue_kernel(+ def _super_epilogue_v17(BMM_OUT_ptr, OG_ptr, OUT_ptr,W_TN_ptr, B_TN_ptr, W_TO_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_c,B, N, C, D: tl.constexpr,- BLOCK_SIZE_N: tl.constexpr,- BLOCK_SIZE_C: 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_b, pid_i = tl.program_id(0), 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_j = pid_j_start + tl.arange(0, BLOCK_SIZE_N)+ mask_j = off_j < Noff_d = tl.arange(0, D)- 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)+ bmm_ptr = BMM_OUT_ptr + pid_b * stride_bmm_b + off_d[None, :] * stride_bmm_d + pid_i * stride_bmm_i + off_j[:, None] * stride_bmm_j+ val = tl.load(bmm_ptr, mask=mask_j[:, None], other=0.0).to(tl.float32)- 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)+ og_ptr = OG_ptr + pid_b * stride_og_b + pid_i * stride_og_i + off_j[:, None] * stride_og_j + off_d[None, :]+ og = tl.load(og_ptr, mask=mask_j[:, None], other=0.0).to(tl.float32)mean = tl.sum(val, axis=1) / Dvar = (tl.sum(val * val, axis=1) / D) - (mean * mean)⋯ 4 unchanged linesval = (val - mean[:, None]) * rstd[:, None] * w_tn[None, :] + b_tn[None, :]val = (val * og).to(tl.float16)- for c_offset in range(0, C, BLOCK_SIZE_C):- off_c = c_offset + tl.arange(0, BLOCK_SIZE_C)+ for c_off in range(0, C, BLOCK_SIZE_C):+ off_c = c_off + tl.arange(0, BLOCK_SIZE_C)c_mask = off_c < Cw_to = tl.load(W_TO_ptr + off_c[None, :] * D + off_d[:, None], mask=c_mask[None, :], other=0.0).to(tl.float16)out_chunk = tl.dot(val, w_to)-- out_ptr = OUT_ptr + pid_b * stride_out_b + pid_i * stride_out_i + offsets_j[:, None] * stride_out_j + off_c[None, :]- tl.store(out_ptr, out_chunk.to(tl.float32), mask=(mask_j[:, None] & c_mask[None, :]))+ out_ptr = OUT_ptr + pid_b * stride_out_b + pid_i * stride_out_i + off_j[:, None] * stride_out_j + off_c[None, :]+ tl.store(out_ptr, out_chunk.to(tl.float32), mask=mask_j[:, None] & c_mask[None, :])+def custom_kernel(data):x, mask, weights, config = dataB, N, _, C = x.shapeD = config["hidden_dim"]device = x.device- w_fp16 = {k: v.to(torch.float16) for k, v in weights.items()}-- # Pack all 5 weights: [C, 5*D]+w_packed = torch.cat([- w_fp16["left_proj.weight"],- w_fp16["left_gate.weight"],- w_fp16["right_proj.weight"],- w_fp16["right_gate.weight"],- w_fp16["out_gate.weight"]- ], dim=0).t().contiguous()+ weights["left_proj.weight"],+ weights["left_gate.weight"],+ weights["right_proj.weight"],+ weights["right_gate.weight"],+ weights["out_gate.weight"]+ ], dim=0).t().to(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_pre = (B, N, (N + 32 - 1) // 32)- _fused_prologue_kernel[grid_pre](+ _fused_prologue_v17[(B, N, (N+32-1)//32)](x, mask, L, R, OG,- w_fp16["norm.weight"], w_fp16["norm.bias"],- w_packed,+ weights["norm.weight"], weights["norm.bias"], w_packed,x.stride(0), x.stride(1), x.stride(2), x.stride(3),- mask.stride(0), mask.stride(1), mask.stride(2),- L.stride(0), L.stride(1), L.stride(2), L.stride(3),- R.stride(0), R.stride(1), R.stride(2), R.stride(3),- OG.stride(0), OG.stride(1), OG.stride(2), OG.stride(3),- B, N, C, D, BLOCK_SIZE_N=32, BLOCK_SIZE_C=32+ B, N, C, D, 32, 32, num_warps=8)- bmm_out = torch.bmm(L.view(B * D, N, N), R.view(B * D, N, N).transpose(-1, -2)).view(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)output = torch.empty((B, N, N, C), device=device, dtype=torch.float32)- grid_epi = (B, N, (N + 32 - 1) // 32)- _super_epilogue_kernel[grid_epi](++ _super_epilogue_v17[(B, N, (N+32-1)//32)](bmm_out, OG, output,- w_fp16["to_out_norm.weight"], w_fp16["to_out_norm.bias"], w_fp16["to_out.weight"],+ weights["to_out_norm.weight"], weights["to_out_norm.bias"], weights["to_out.weight"],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),output.stride(0), output.stride(1), output.stride(2), output.stride(3),- B, N, C, D, BLOCK_SIZE_N=32, BLOCK_SIZE_C=64+ B, N, C, D, 32, 64, num_warps=4)-return output
scrolls · 247 diff lines total
Best evidence level for this revision: reported
JSON