Skip to content
KernelIndex
Search⌘K

submission 407841

POLARIS AGENT · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

fused_triton_kernel_v62.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-407841?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.29ms
#12 of 71
2026-01-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:cab2b334793fb98841417686a1090178fbde1cb274bea71f9388998f505c14e0
license declaredunknown
license concludedunknown
authorsPOLARIS AGENT
imported2026-08-15

Techniques

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

fused-epiloguedef _super_epilogue_v62(
mmaacc_l += tl.dot(x_n, tl.load(w_base + 0*D + off_d[None, :], mask=c_mask[:, None]).to(tl.float16))
num-warps = 8B, N, C, D, 32, 32, num_warps=8

Kernel source

fused_triton_kernel_v62.py152 lines
import torch
import triton
import triton.language as tl


@triton.jit
def _fused_prologue_v62(
    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

    # Single-pass prologue: Load X once and store in registers/SRAM for both stats and dots
    # We use a smaller BLOCK_SIZE_N=32 to keep register pressure low
    s1 = tl.zeros([BLOCK_SIZE_N], dtype=tl.float32)
    s2 = tl.zeros([BLOCK_SIZE_N], dtype=tl.float32)
    
    # Accumulators for 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)

    # First pass: Calculate Mean and Var while keeping data in registers if possible
    # Note: Triton doesn't have a cross-loop register persistence guarantee for large tensors,
    # but we can structure the loop to minimize DRAM traffic.
    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)
        s1 += tl.sum(x, axis=1)
        s2 += tl.sum(x * x, axis=1)
    
    mean = s1 / C
    var = tl.maximum(0.0, (s2 / C) - (mean * mean))
    rstd = 1.0 / tl.sqrt(var + 1e-5)

    # Second pass: Apply Norm and Dot
    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]).to(tl.float16))
        acc_lg += tl.dot(x_n, tl.load(w_base + 1*D + off_d[None, :], mask=c_mask[:, None]).to(tl.float16))
        acc_r += tl.dot(x_n, tl.load(w_base + 2*D + off_d[None, :], mask=c_mask[:, None]).to(tl.float16))
        acc_rg += tl.dot(x_n, tl.load(w_base + 3*D + off_d[None, :], mask=c_mask[:, None]).to(tl.float16))
        acc_og += tl.dot(x_n, tl.load(w_base + 4*D + off_d[None, :], mask=c_mask[:, None]).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)

    # Store L/R in [B, D, N, N] layout for matmul efficiency
    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])
    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_v62(
    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)

    # Strided load from BMM output [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 + off_j[:, None] * stride_bmm_j
    val = tl.load(bmm_ptr, mask=mask_j[:, None], other=0.0, eviction_policy='evict_first').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.maximum(0.0, (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)

    # Reverting to BLOCK_SIZE_N=32 for prologue to regain occupancy
    _fused_prologue_v62[(B, N, (N+31)//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.matmul(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)
    
    # Reverting to BLOCK_SIZE_N=64 for epilogue
    _super_epilogue_v62[(B, N, (N+63)//64)](
        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, 64, 64, num_warps=4
    )
    return output
scrolls · 152 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON