Skip to content
KernelIndex
Search⌘K

submission 407575

Zeyu Shen · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

fused_packed_weights_super_epilogue.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-407575?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.67ms
#21 of 71
2026-01-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2473f844c2e969c6bf1f183ff46843da095ef439126686e662c07e84e8c65049
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-epiloguedef _super_epilogue_kernel(
mmaacc_l += tl.dot(x_n, tl.load(w_base + 0*D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16))
tile-n = 32B, N, C, D, BLOCK_SIZE_N=32, BLOCK_SIZE_C=32

Kernel source

fused_packed_weights_super_epilogue.py177 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_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
):
    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 (Online algorithm)
    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 with Packed Load, Separate Compute
    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, :]).to(tl.float16)

        # 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))

    # 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)

    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(
    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 = 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)

    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)
    
    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)

    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_offset in range(0, C, BLOCK_SIZE_C):
        off_c = c_offset + 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 + 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, :]))

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()}

    # 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()

    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](
        x, mask, L, R, OG,
        w_fp16["norm.weight"], w_fp16["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
    )

    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](
        bmm_out, OG, output,
        w_fp16["to_out_norm.weight"], w_fp16["to_out_norm.bias"], w_fp16["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
    )

    return output
scrolls · 177 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 407551.

⋯ 5 unchanged lines
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,
+ 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
⋯ 5 unchanged lines
offsets_j = pid_j_start + tl.arange(0, BLOCK_SIZE_N)
mask_j = offsets_j < N
- # 1. LayerNorm statistics
+ # 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)
⋯ 9 unchanged lines
var = (acc_sum_sq / C) - (mean * mean)
rstd = 1.0 / tl.sqrt(var + 1e-5)
- # 2. Projections
+ # 2. Projections with Packed Load, Separate Compute
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)
⋯ 10 unchanged lines
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)
+ x_n = ((x_chunk - mean[:, None]) * rstd[:, None] * w_n[None, :] + b_n[None, :]).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))
+ # 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))
+ # 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)
⋯ 1 unchanged lines
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]
+ # 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 * N * N * D + idx_nn[:, None] * D + off_d[None, :]
+ 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 _fused_epilogue_kernel(
+ def _super_epilogue_kernel(
BMM_OUT_ptr, OG_ptr, OUT_ptr,
- W_TN_ptr, B_TN_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_d,
- B, N, D: tl.constexpr,
- BLOCK_SIZE_N: tl.constexpr
+ 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 = tl.program_id(0)
pid_i = tl.program_id(1)
⋯ 3 unchanged lines
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)
⋯ 1 unchanged lines
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
+ val = (val * og).to(tl.float16)
- # 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)))
+ for c_offset in range(0, C, BLOCK_SIZE_C):
+ off_c = c_offset + 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 + 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, :]))
def custom_kernel(data):
x, mask, weights, config = data
⋯ 2 unchanged lines
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()
+
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(),
+ w_packed,
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
+ 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
)
- # 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)
+ output = torch.empty((B, N, N, C), device=device, dtype=torch.float32)
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"],
+ _super_epilogue_kernel[grid_epi](
+ bmm_out, OG, output,
+ w_fp16["to_out_norm.weight"], w_fp16["to_out_norm.bias"], w_fp16["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),
- epi_inter.stride(0), epi_inter.stride(1), epi_inter.stride(2), epi_inter.stride(3),
- B, N, D, BLOCK_SIZE_N=32
+ output.stride(0), output.stride(1), output.stride(2), output.stride(3),
+ B, N, C, D, BLOCK_SIZE_N=32, BLOCK_SIZE_C=64
)
- # Epilogue Part 2: Final Projection to C using cuBLAS.
- return (epi_inter @ w_fp16["to_out.weight"].t()).to(torch.float32)
+ return output
scrolls · 188 diff lines total

Best evidence level for this revision: reported

JSON