submission 35291
msuiche · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 358 lines, June 9 Researcher Reciprocity License v1.0.
triton_optimized.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-35291?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:e7ce7e5916d80d50510e6302aa6848dda36c4b4bd580a22497fa735011619953
license declaredunknown
license concludedunknown
authorsmsuiche
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
acc_left_proj += tl.dot(norm_x, tl.trans(w_left))Kernel source
triton_optimized.py358 lines
#!POPCORN leaderboard trimul
#!POPCORN gpu H100
from utils import make_match_reference, DisableCuDNNTF32
from task import input_t, output_t
import torch
import torch.nn as nn
# import torch.nn.functional as F # Not used
import triton
import triton.language as tl
import math
# Triton kernel for computing layernorm statistics
@triton.jit
def layernorm_stats_kernel(
x_ptr, ln_stats_ptr, N, D,
stride_x_b, stride_x_n1, stride_x_n2, stride_x_d,
stride_ln_b, stride_ln_n,
BLOCK_K: tl.constexpr
):
pid_b = tl.program_id(0) # batch index
pid_n = tl.program_id(1) # N index (combined n1*n2)
offs_k = tl.arange(0, BLOCK_K)
# Compute mean and variance
sum_x = tl.zeros((1,), dtype=tl.float32)
sum_x2 = tl.zeros((1,), dtype=tl.float32)
num_k_blocks = tl.cdiv(D, BLOCK_K)
for kb in range(num_k_blocks):
k_idx = kb * BLOCK_K + offs_k
valid_k = k_idx < D
# Load input block
x_ptrs = x_ptr + pid_b * stride_x_b + pid_n * stride_x_n1 + k_idx * stride_x_d
x_block = tl.load(x_ptrs, mask=valid_k, other=0.0)
sum_x += tl.sum(x_block, axis=0)
sum_x2 += tl.sum(x_block * x_block, axis=0)
mean = sum_x / D
var = tl.maximum(sum_x2 / D - mean * mean, 0.0)
inv_std = tl.rsqrt(var + 1e-5)
# Store stats
ln_stats_ptr_b = ln_stats_ptr + pid_b * stride_ln_b + pid_n * stride_ln_n
tl.store(ln_stats_ptr_b, mean)
tl.store(ln_stats_ptr_b + 1, inv_std)
# Main fused Triton kernel
@triton.jit
def trimul_fused_kernel(
x_ptr, mask_ptr,
ln_w_ptr, ln_b_ptr,
proj_gates_w_ptr,
out_norm_w_ptr, out_norm_b_ptr,
to_out_w_ptr,
ln_stats_ptr,
left_out_ptr, right_out_ptr,
B, N, D, H,
stride_x_b, stride_x_n1, stride_x_n2, stride_x_d,
stride_mask_b, stride_mask_n1, stride_mask_n2,
stride_ln_b, stride_ln_n,
stride_w_h, stride_w_d,
stride_out_b, stride_out_n1, stride_out_n2, stride_out_h,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr
):
pid_b = tl.program_id(0)
pid_m = tl.program_id(1) # N index for output
pid_n = tl.program_id(2) # H index for output
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
offs_k = tl.arange(0, BLOCK_K)
mask_m = offs_m < N
mask_n = offs_n < H
# Load precomputed stats
ln_stats_ptr_bm = ln_stats_ptr + pid_b * stride_ln_b + offs_m * stride_ln_n
mean = tl.load(ln_stats_ptr_bm, mask=mask_m, other=0.0)
inv_std = tl.load(ln_stats_ptr_bm + 1, mask=mask_m, other=1.0)
# Initialize accumulators
acc_left_proj = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
acc_right_proj = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
acc_left_gate = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
acc_right_gate = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
acc_out_gate = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# Process input in blocks
num_k_blocks = tl.cdiv(D, BLOCK_K)
for kb in range(num_k_blocks):
k_idx = kb * BLOCK_K + offs_k
valid_k = k_idx < D
# Load and normalize input
n1 = offs_m // N
n2 = offs_m % N
x_ptrs = x_ptr + pid_b * stride_x_b + n1[:, None] * stride_x_n1 + n2[:, None] * stride_x_n2 + k_idx[None, :] * stride_x_d
x_block = tl.load(x_ptrs, mask=(mask_m[:, None] & valid_k[None, :]), other=0.0)
# LayerNorm
ln_w = tl.load(ln_w_ptr + k_idx, mask=valid_k, other=1.0)
ln_b = tl.load(ln_b_ptr + k_idx, mask=valid_k, other=0.0)
norm_x = ((x_block - mean[:, None]) * inv_std[:, None] * ln_w[None, :]) + ln_b[None, :]
# Load weights and accumulate
# Left projection
w_left_ptr = proj_gates_w_ptr + offs_n[:, None] * stride_w_h + k_idx[None, :] * stride_w_d
w_left = tl.load(w_left_ptr, mask=(mask_n[:, None] & valid_k[None, :]), other=0.0)
acc_left_proj += tl.dot(norm_x, tl.trans(w_left))
# Right projection
w_right_ptr = proj_gates_w_ptr + (H + offs_n)[:, None] * stride_w_h + k_idx[None, :] * stride_w_d
w_right = tl.load(w_right_ptr, mask=(mask_n[:, None] & valid_k[None, :]), other=0.0)
acc_right_proj += tl.dot(norm_x, tl.trans(w_right))
# Gates
w_left_gate_ptr = proj_gates_w_ptr + (2*H + offs_n)[:, None] * stride_w_h + k_idx[None, :] * stride_w_d
w_left_gate = tl.load(w_left_gate_ptr, mask=(mask_n[:, None] & valid_k[None, :]), other=0.0)
acc_left_gate += tl.dot(norm_x, tl.trans(w_left_gate))
w_right_gate_ptr = proj_gates_w_ptr + (3*H + offs_n)[:, None] * stride_w_h + k_idx[None, :] * stride_w_d
w_right_gate = tl.load(w_right_gate_ptr, mask=(mask_n[:, None] & valid_k[None, :]), other=0.0)
acc_right_gate += tl.dot(norm_x, tl.trans(w_right_gate))
w_out_gate_ptr = proj_gates_w_ptr + (4*H + offs_n)[:, None] * stride_w_h + k_idx[None, :] * stride_w_d
w_out_gate = tl.load(w_out_gate_ptr, mask=(mask_n[:, None] & valid_k[None, :]), other=0.0)
acc_out_gate += tl.dot(norm_x, tl.trans(w_out_gate))
# Apply mask and gates
n1 = offs_m // N
n2 = offs_m % N
mask_ptrs = mask_ptr + pid_b * stride_mask_b + n1 * stride_mask_n1 + n2 * stride_mask_n2
mask_val = tl.load(mask_ptrs, mask=mask_m, other=0.0)
# Apply gates directly without clamping for accuracy
left_gated = acc_left_proj * mask_val[:, None] * tl.sigmoid(acc_left_gate)
right_gated = acc_right_proj * mask_val[:, None] * tl.sigmoid(acc_right_gate)
# Store intermediate results
left_out_ptrs = left_out_ptr + pid_b * stride_out_b + n1[:, None] * stride_out_n1 + n2[:, None] * stride_out_n2 + offs_n[None, :] * stride_out_h
right_out_ptrs = right_out_ptr + pid_b * stride_out_b + n1[:, None] * stride_out_n1 + n2[:, None] * stride_out_n2 + offs_n[None, :] * stride_out_h
store_mask = mask_m[:, None] & mask_n[None, :]
tl.store(left_out_ptrs, left_gated, mask=store_mask)
tl.store(right_out_ptrs, right_gated, mask=store_mask)
class TritonTriMul(nn.Module):
"""
Triton-optimized TriMul implementation - exact copy from working H100FlashTriMul
"""
def __init__(self, dim: int, hidden_dim: int):
super().__init__()
self.dim = dim
self.hidden_dim = hidden_dim
# Single fused weight matrix for all linear ops
# This reduces memory accesses
total_params = hidden_dim * 5
self.fused_proj_gates = nn.Linear(dim, total_params, bias=False)
# Separate norms (hard to fuse efficiently)
self.norm = nn.LayerNorm(dim) # Use default eps to match reference
self.out_norm = nn.LayerNorm(hidden_dim)
self.to_out = nn.Linear(hidden_dim, dim, bias=False)
def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
B, N, _, _ = x.shape
H = self.hidden_dim
# LayerNorm and projection
x = self.norm(x)
all_features = self.fused_proj_gates(x)
# Simple reshape and extraction
features = all_features.view(B, N, N, 5, H)
# Extract projections and compute gates
left_proj = features[..., 0, :]
right_proj = features[..., 1, :]
left_gate = features[..., 2, :].sigmoid()
right_gate = features[..., 3, :].sigmoid()
out_gate = features[..., 4, :].sigmoid()
# Apply mask and gates
mask_expanded = mask.unsqueeze(-1)
left = left_proj * mask_expanded * left_gate
right = right_proj * mask_expanded * right_gate
# Make contiguous for einsum
left = left.contiguous()
right = right.contiguous()
# Einsum with FP32 accumulation for better precision
if left.dtype != torch.float32:
left_f32 = left.float()
right_f32 = right.float()
out = torch.einsum('bikd,bjkd->bijd', left_f32, right_f32)
out = out.to(left.dtype)
else:
out = torch.einsum('bikd,bjkd->bijd', left, right)
# Fused output processing
# Combine normalization, gating, and projection
out = self.to_out(self.out_norm(out) * out_gate)
return out
def custom_kernel(data: input_t) -> output_t:
"""
Custom kernel implementation for TriMul using Triton optimization
"""
with DisableCuDNNTF32():
input_tensor, mask, weights, config = data
dim = config["dim"]
hidden_dim = config["hidden_dim"]
# Create model
model = TritonTriMul(dim=dim, hidden_dim=hidden_dim)
model = model.to(input_tensor.device)
# Skip compilation to avoid timeout
# Compilation adds overhead for first run which can cause timeout
pass
# Load weights
with torch.no_grad():
# Stack all projection and gate weights
proj_weights = torch.cat([
weights['left_proj.weight'],
weights['right_proj.weight'],
weights['left_gate.weight'],
weights['right_gate.weight'],
weights['out_gate.weight']
], dim=0)
model.fused_proj_gates.weight.data = proj_weights
model.norm.weight.data = weights['norm.weight']
model.norm.bias.data = weights['norm.bias']
model.out_norm.weight.data = weights['to_out_norm.weight']
model.out_norm.bias.data = weights['to_out_norm.bias']
model.to_out.weight.data = weights['to_out.weight']
# Ensure contiguous tensors
input_tensor = input_tensor.contiguous()
mask = mask.contiguous()
# Run model directly without CUDA graph to avoid timeout
with torch.no_grad():
output = model(input_tensor, mask)
return output
# Reference implementation for testing
class TriMul(nn.Module):
def __init__(self, dim: int, hidden_dim: int):
super().__init__()
self.norm = nn.LayerNorm(dim)
self.left_proj = nn.Linear(dim, hidden_dim, bias=False)
self.right_proj = nn.Linear(dim, hidden_dim, bias=False)
self.left_gate = nn.Linear(dim, hidden_dim, bias=False)
self.right_gate = nn.Linear(dim, hidden_dim, bias=False)
self.out_gate = nn.Linear(dim, hidden_dim, bias=False)
self.to_out_norm = nn.LayerNorm(hidden_dim)
self.to_out = nn.Linear(hidden_dim, dim, bias=False)
def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
x = self.norm(x)
left = self.left_proj(x)
right = self.right_proj(x)
mask = mask.unsqueeze(-1)
left = left * mask
right = right * mask
left_gate = self.left_gate(x).sigmoid()
right_gate = self.right_gate(x).sigmoid()
out_gate = self.out_gate(x).sigmoid()
left = left * left_gate
right = right * right_gate
out = torch.einsum('... i k d, ... j k d -> ... i j d', left, right)
out = self.to_out_norm(out)
out = out * out_gate
return self.to_out(out)
def ref_kernel(data: input_t) -> output_t:
with DisableCuDNNTF32():
input_tensor, mask, weights, config = data
trimul = TriMul(dim=config["dim"], hidden_dim=config["hidden_dim"]).to(input_tensor.device)
trimul.norm.weight = nn.Parameter(weights['norm.weight'])
trimul.norm.bias = nn.Parameter(weights['norm.bias'])
trimul.left_proj.weight = nn.Parameter(weights['left_proj.weight'])
trimul.right_proj.weight = nn.Parameter(weights['right_proj.weight'])
trimul.left_gate.weight = nn.Parameter(weights['left_gate.weight'])
trimul.right_gate.weight = nn.Parameter(weights['right_gate.weight'])
trimul.out_gate.weight = nn.Parameter(weights['out_gate.weight'])
trimul.to_out_norm.weight = nn.Parameter(weights['to_out_norm.weight'])
trimul.to_out_norm.bias = nn.Parameter(weights['to_out_norm.bias'])
trimul.to_out.weight = nn.Parameter(weights['to_out.weight'])
output = trimul(input_tensor, mask)
return output
def generate_input(seqlen: int, bs: int, dim: int, hiddendim: int,
seed: int, nomask: bool, distribution: str) -> input_t:
batch_size = bs
seq_len = seqlen
hidden_dim = hiddendim
no_mask = nomask
config = {"hidden_dim": hidden_dim, "dim": dim}
gen = torch.Generator(device='cuda')
gen.manual_seed(seed)
weights = {}
if distribution == "cauchy":
input_tensor = torch.distributions.Cauchy(0, 2).sample(
(batch_size, seq_len, seq_len, dim)
).to(device='cuda', dtype=torch.float32)
else:
input_tensor = torch.randn(
(batch_size, seq_len, seq_len, dim),
device='cuda', dtype=torch.float32, generator=gen
).contiguous()
if no_mask:
mask = torch.ones(batch_size, seq_len, seq_len, device=input_tensor.device)
else:
mask = torch.randint(0, 2, (batch_size, seq_len, seq_len),
device=input_tensor.device, generator=gen)
weights["norm.weight"] = torch.randn(dim, device="cuda", dtype=torch.float32)
weights["norm.bias"] = torch.randn(dim, device="cuda", dtype=torch.float32)
weights["left_proj.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
weights["right_proj.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
weights["left_gate.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
weights["right_gate.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
weights["out_gate.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
weights["to_out_norm.weight"] = torch.randn(hidden_dim, device="cuda", dtype=torch.float32)
weights["to_out.weight"] = torch.randn(dim, hidden_dim, device="cuda", dtype=torch.float32) / math.sqrt(dim)
weights["to_out_norm.bias"] = torch.randn(hidden_dim, device="cuda", dtype=torch.float32)
return (input_tensor, mask, weights, config)
# Check implementation correctness
check_implementation = make_match_reference(custom_kernel, rtol=2e-2, atol=2e-2)scrolls · 358 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 35286.
⋯ 9 unchanged linesimport triton.language as tlimport math- # JIT-compiled fusion function for gate application- @torch.jit.script- def fused_gate_application(projections: torch.Tensor,- gates_logits: torch.Tensor,- mask: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:- """Fused gate application with sigmoid and masking"""- gates = torch.sigmoid(gates_logits)- mask_expanded = mask.unsqueeze(-1)-- left = projections[..., 0, :] * mask_expanded * gates[..., 0, :]- right = projections[..., 1, :] * mask_expanded * gates[..., 1, :]- out_gate = gates[..., 2, :]-- return left.contiguous(), right.contiguous(), out_gate-# Triton kernel for computing layernorm statistics@triton.jitdef layernorm_stats_kernel(⋯ 138 unchanged linesTriton-optimized TriMul implementation - exact copy from working H100FlashTriMul"""- def __init__(self, dim: int, hidden_dim: int, use_checkpoint: bool = False):+ def __init__(self, dim: int, hidden_dim: int):super().__init__()self.dim = dimself.hidden_dim = hidden_dim- self.use_checkpoint = use_checkpoint# Single fused weight matrix for all linear ops# This reduces memory accesses⋯ 5 unchanged linesself.out_norm = nn.LayerNorm(hidden_dim)self.to_out = nn.Linear(hidden_dim, dim, bias=False)- def _compute_features(self, x: torch.Tensor):- """Compute features - can be checkpointed for memory efficiency"""- return self.fused_proj_gates(self.norm(x))-def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:B, N, _, _ = x.shapeH = self.hidden_dim- # Compute features directly (no checkpointing for inference)- all_features = self._compute_features(x)+ # LayerNorm and projection+ x = self.norm(x)+ all_features = self.fused_proj_gates(x)- # Memory-efficient reshape without copying data- # First reshape to optimal layout for extraction- all_features = all_features.view(B, N * N, 5, H)-- # Extract components with memory-efficient views+ # Simple reshape and extractionfeatures = all_features.view(B, N, N, 5, H)- # Fused extraction and activation using JIT-compiled function- projections = features[..., :2, :] # left and right projections- gates_logits = features[..., 2:, :] # all gate logits+ # Extract projections and compute gates+ left_proj = features[..., 0, :]+ right_proj = features[..., 1, :]+ left_gate = features[..., 2, :].sigmoid()+ right_gate = features[..., 3, :].sigmoid()+ out_gate = features[..., 4, :].sigmoid()- # Use JIT-compiled fusion for gate application- left, right, out_gate = fused_gate_application(projections, gates_logits, mask)+ # Apply mask and gates+ mask_expanded = mask.unsqueeze(-1)+ left = left_proj * mask_expanded * left_gate+ right = right_proj * mask_expanded * right_gate+ # Make contiguous for einsum+ left = left.contiguous()+ right = right.contiguous()+# Einsum with FP32 accumulation for better precisionif left.dtype != torch.float32:left_f32 = left.float()
scrolls · 85 diff lines total
Best evidence level for this revision: reported
JSON