submission 407519
Zeyu Shen · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 73 lines, June 9 Researcher Reciprocity License v1.0.
fused_trimul.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-407519?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:7401890f4053bc3c6f83428ad4a6cdfba8685bbadf963fb3ff7950c47cd5d0b5
license declaredunknown
license concludedunknown
authorsZeyu Shen
imported2026-08-15
Kernel source
fused_trimul.py73 lines
import torch
import triton
import triton.language as tl
@triton.jit
def fused_trimul_kernel(
X_ptr, M_ptr, W_ptr, B_ptr, OUT_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
):
# 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)
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)
# Pointers for the specific batch
X_batch_ptr = X_ptr + pid_b * stride_xb
# 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 k in range(0, N, BLOCK_SIZE_K):
# Load blocks and perform contraction
# This is a 3D tiled reduction
pass
def custom_kernel(data):
input_tensor, mask, weights, config = data
dim, hidden_dim = config["dim"], config["hidden_dim"]
device = input_tensor.device
# 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)
scrolls · 73 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