Skip to content
KernelIndex
Search⌘K

submission 407538

Zeyu Shen · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

fused_preprocess_kernel.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-407538?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
5.91ms
#48 of 71
2026-01-27

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

mmaacc_l += tl.dot(x_n, w_l)
stages = 1num_stages=1
tile-n = 32BLOCK_SIZE_N = 32

Kernel source

fused_preprocess_kernel.py132 lines
import torch
import triton
import triton.language as tl

@triton.jit
def _fused_preprocess_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,
    stride_xb, stride_xi, stride_xj, stride_xc,
    stride_mb, stride_mi, stride_mj,
    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. Compute LayerNorm for the block [BLOCK_SIZE_N, C]
    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. Compute Projections
    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, :]
        x_n = x_n.to(tl.float16)

        # Projection weights [C, D]
        w_l = tl.load(W_L_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16)
        acc_l += tl.dot(x_n, w_l)
        
        w_lg = tl.load(W_LG_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16)
        acc_lg += tl.dot(x_n, w_lg)
        
        w_r = tl.load(W_R_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16)
        acc_r += tl.dot(x_n, w_r)
        
        w_rg = tl.load(W_RG_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16)
        acc_rg += tl.dot(x_n, w_rg)
        
        w_og = tl.load(W_OG_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16)
        acc_og += tl.dot(x_n, w_og)

    # 3. Apply 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. Store results
    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, :]
    tl.store(OG_ptr + off_og, og_final.to(tl.float16), mask=(mask_j[:, None] & (off_d[None, :] < D)))

def custom_kernel(data):
    x, mask, weights, config = data
    B, N, _, C = x.shape
    D = config["hidden_dim"]
    device = x.device

    # Use Half for intermediate activations to speed up BMM and save memory
    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)

    BLOCK_SIZE_N = 32
    BLOCK_SIZE_C = 128
    
    grid = (B, N, (N + BLOCK_SIZE_N - 1) // BLOCK_SIZE_N)
    _fused_preprocess_kernel[grid](
        x, mask, L, R, OG,
        weights["norm.weight"], weights["norm.bias"],
        weights["left_proj.weight"].t().contiguous(), weights["right_proj.weight"].t().contiguous(),
        weights["left_gate.weight"].t().contiguous(), weights["right_gate.weight"].t().contiguous(),
        weights["out_gate.weight"].t().contiguous(),
        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=BLOCK_SIZE_N,
        BLOCK_SIZE_C=BLOCK_SIZE_C,
        num_stages=1
    )

    # BMM expects [B*D, N, N] and returns [B*D, N, N] in Half
    bmm_out = torch.bmm(L.view(B * D, N, N), R.view(B * D, N, N).transpose(-1, -2))
    bmm_out = bmm_out.view(B, D, N, N).permute(0, 2, 3, 1) # [B, N, N, D]

    # Fix: Ensure bmm_out matches weight dtype (Float32) for LayerNorm
    target_dtype = weights["to_out_norm.weight"].dtype
    bmm_out = bmm_out.to(target_dtype)

    out = torch.nn.functional.layer_norm(bmm_out, (D,), weights["to_out_norm.weight"], weights["to_out_norm.bias"])
    out = out * OG.to(target_dtype)
    return (out @ weights["to_out.weight"].t()).to(torch.float32)
scrolls · 132 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 407519.

⋯ 2 unchanged lines
import triton.language as tl
@triton.jit
- def fused_trimul_kernel(
- X_ptr, M_ptr, W_ptr, B_ptr, OUT_ptr,
+ def _fused_preprocess_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,
stride_xb, stride_xi, stride_xj, stride_xc,
stride_mb, stride_mi, stride_mj,
- stride_ob, stride_oi, stride_oj, stride_oc,
- B, N, C, H,
- BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
- BLOCK_SIZE_D: tl.constexpr
+ B, N, C, D: tl.constexpr,
+ BLOCK_SIZE_N: tl.constexpr,
+ BLOCK_SIZE_C: tl.constexpr
):
- # This kernel handles the core einsum: out[i, j, d] = sum_k (left[i, k, d] * right[j, k, d])
- # For simplicity in this first iteration, we assume projections are pre-computed or handled.
- # However, to beat the baseline, we must fuse.
- # Let's implement a simplified fused version focusing on the O(N^3) part.
-
- pid_m = tl.program_id(0)
- pid_n = tl.program_id(1)
- pid_b = tl.program_id(2)
+ pid_b = tl.program_id(0)
+ pid_i = tl.program_id(1)
+ pid_j_start = tl.program_id(2) * BLOCK_SIZE_N
- rm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
- rn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
- rk = tl.arange(0, BLOCK_SIZE_K)
- rd = tl.arange(0, BLOCK_SIZE_D)
+ offsets_j = pid_j_start + tl.arange(0, BLOCK_SIZE_N)
+ mask_j = offsets_j < N
- # Pointers for the specific batch
- X_batch_ptr = X_ptr + pid_b * stride_xb
+ # 1. Compute LayerNorm for the block [BLOCK_SIZE_N, C]
+ acc_sum = tl.zeros([BLOCK_SIZE_N], dtype=tl.float32)
+ acc_sum_sq = tl.zeros([BLOCK_SIZE_N], dtype=tl.float32)
- # In a real optimized version, we'd load X, apply LayerNorm and Projections here.
- # For this submission, we'll focus on the structure of the einsum contraction.
- # out[i, j, d] = sum_k left[i, k, d] * right[j, k, d]
-
- acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_D), 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)
- for k in range(0, N, BLOCK_SIZE_K):
- # Load blocks and perform contraction
- # This is a 3D tiled reduction
- pass
+ mean = acc_sum / C
+ var = (acc_sum_sq / C) - (mean * mean)
+ rstd = 1.0 / tl.sqrt(var + 1e-5)
+ # 2. Compute Projections
+ 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, :]
+ x_n = x_n.to(tl.float16)
+
+ # Projection weights [C, D]
+ w_l = tl.load(W_L_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16)
+ acc_l += tl.dot(x_n, w_l)
+
+ w_lg = tl.load(W_LG_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16)
+ acc_lg += tl.dot(x_n, w_lg)
+
+ w_r = tl.load(W_R_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16)
+ acc_r += tl.dot(x_n, w_r)
+
+ w_rg = tl.load(W_RG_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16)
+ acc_rg += tl.dot(x_n, w_rg)
+
+ w_og = tl.load(W_OG_ptr + cols[:, None] * D + off_d[None, :], mask=c_mask[:, None], other=0.0).to(tl.float16)
+ acc_og += tl.dot(x_n, w_og)
+
+ # 3. Apply 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. Store results
+ 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, :]
+ tl.store(OG_ptr + off_og, og_final.to(tl.float16), mask=(mask_j[:, None] & (off_d[None, :] < D)))
+
def custom_kernel(data):
- input_tensor, mask, weights, config = data
- dim, hidden_dim = config["dim"], config["hidden_dim"]
- device = input_tensor.device
+ x, mask, weights, config = data
+ B, N, _, C = x.shape
+ D = config["hidden_dim"]
+ device = x.device
+
+ # Use Half for intermediate activations to speed up BMM and save memory
+ 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)
+
+ BLOCK_SIZE_N = 32
+ BLOCK_SIZE_C = 128
- # 1. LayerNorm
- x = torch.nn.functional.layer_norm(input_tensor, (dim,), weights["norm.weight"], weights["norm.bias"])
-
- # 2. Projections
- # left = (x @ W_l) * mask * sigmoid(x @ W_lg)
- # right = (x @ W_r) * mask * sigmoid(x @ W_rg)
- left = torch.matmul(x, weights["left_proj.weight"].t()) * mask.unsqueeze(-1) * torch.sigmoid(torch.matmul(x, weights["left_gate.weight"].t()))
- right = torch.matmul(x, weights["right_proj.weight"].t()) * mask.unsqueeze(-1) * torch.sigmoid(torch.matmul(x, weights["right_gate.weight"].t()))
-
- # 3. Core Einsum (The O(N^3) part)
- # out = einsum('... i k d, ... j k d -> ... i j d', left, right)
- # Optimization: Reshape to use batch matmul
- # left: [B, N, N, H] -> [B, H, N, N]
- # right: [B, N, N, H] -> [B, H, N, N]
- # result: [B, H, N, N] -> [B, N, N, H]
- l_r = left.permute(0, 3, 1, 2) # [B, H, Ni, Nk]
- r_r = right.permute(0, 3, 2, 1) # [B, H, Nk, Nj]
- out = torch.matmul(l_r, r_r).permute(0, 2, 3, 1) # [B, Ni, Nj, H]
-
- # 4. Epilogue
- out = torch.nn.functional.layer_norm(out, (hidden_dim,), weights["to_out_norm.weight"], weights["to_out_norm.bias"])
- out = out * torch.sigmoid(torch.matmul(x, weights["out_gate.weight"].t()))
- out = torch.matmul(out, weights["to_out.weight"].t())
-
- return out.to(torch.float32)
+ grid = (B, N, (N + BLOCK_SIZE_N - 1) // BLOCK_SIZE_N)
+ _fused_preprocess_kernel[grid](
+ x, mask, L, R, OG,
+ weights["norm.weight"], weights["norm.bias"],
+ weights["left_proj.weight"].t().contiguous(), weights["right_proj.weight"].t().contiguous(),
+ weights["left_gate.weight"].t().contiguous(), weights["right_gate.weight"].t().contiguous(),
+ weights["out_gate.weight"].t().contiguous(),
+ 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=BLOCK_SIZE_N,
+ BLOCK_SIZE_C=BLOCK_SIZE_C,
+ num_stages=1
+ )
+
+ # BMM expects [B*D, N, N] and returns [B*D, N, N] in Half
+ bmm_out = torch.bmm(L.view(B * D, N, N), R.view(B * D, N, N).transpose(-1, -2))
+ bmm_out = bmm_out.view(B, D, N, N).permute(0, 2, 3, 1) # [B, N, N, D]
+
+ # Fix: Ensure bmm_out matches weight dtype (Float32) for LayerNorm
+ target_dtype = weights["to_out_norm.weight"].dtype
+ bmm_out = bmm_out.to(target_dtype)
+
+ out = torch.nn.functional.layer_norm(bmm_out, (D,), weights["to_out_norm.weight"], weights["to_out_norm.bias"])
+ out = out * OG.to(target_dtype)
+ return (out @ weights["to_out.weight"].t()).to(torch.float32)
scrolls · 187 diff lines total

Best evidence level for this revision: reported

JSON