Skip to content
KernelIndex
Search⌘K

submission 408231

Zeyu Shen · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 138 lines, June 9 Researcher Reciprocity License v1.0.

fused_frontend_backend_v102.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-408231?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
NVIDIA H100
1.30ms
#13 of 71
2026-01-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d2b02700ea8a2ab6c72d903c5eb8e92f5239d10f1e0cc0adc16fa8db96f79592
license declaredunknown
license concludedunknown
authorsZeyu Shen
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

mmal_acc += tl.dot(xn, tl.load(W_ptr + 0*w_stride_type + w_off, mask=c_mask[:, None], other=0.0).to(tl.float16))
num-warps = 8fused_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)
stages = 2fused_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)
tile-n = 64fused_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)

Kernel source

fused_frontend_backend_v102.py138 lines
import torch
import triton
import triton.language as tl


@triton.jit
def fused_frontend_v102(
    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
    
    # Vectorized Statistics Pass (Single-pass to reduce 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)
    w_stride_type = 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]
        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))

    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]
    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
    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(
    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)
    
    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)
    
    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
        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
    
    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()
    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)
    
    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)
    return out
scrolls · 138 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 407612.

⋯ 3 unchanged lines
@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,
+ def fused_frontend_v102(
+ X_ptr, Mask_ptr,
+ W_ptr, NW_ptr, NB_ptr,
+ L_ptr, R_ptr, OG_ptr,
B, N, C, D: tl.constexpr,
- BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_C: 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, 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)
+ 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
- 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)
+ # Vectorized Statistics Pass (Single-pass to reduce 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_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)
+ 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)
+ w_stride_type = 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_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]))
+ 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))
- 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)
+ 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
- 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])
+ # 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]
+ 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
+ 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 _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
+ def fused_backend_v102(
+ 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_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)
+ 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)
- 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)
+ 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)
- 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)
+ nw = tl.load(NW_ptr + rd)[:, None]
+ nb = tl.load(NB_ptr + rd)[:, None]
+ xn = (diff * rstd * nw + nb).to(tl.float16)
+
+ 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_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, :])
+ for c_off in range(0, C, BLOCK_C):
+ rc = c_off + tl.arange(0, BLOCK_C)
+ c_mask = rc < C
+ 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):
⋯ 2 unchanged lines
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)
+ 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()
+ to_out_w = weights["to_out.weight"].to(device, torch.float16).contiguous()
- output = torch.empty((B, N, N, C), device=device, dtype=torch.float32)
+ 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)
- _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
+ 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)
+
+ 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)
+ return out
scrolls · 251 diff lines total

Best evidence level for this revision: reported

JSON