submission 35618
msuiche · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 65 lines, June 9 Researcher Reciprocity License v1.0.
cuda_fast.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-35618?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:7dded55092b360fda295dad8a02c8d76c4112119c0dfb236790801e7105d78f8
license declaredunknown
license concludedunknown
authorsmsuiche
imported2026-08-15
Kernel source
cuda_fast.py65 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.functional as F
def custom_kernel(data: input_t) -> output_t:
"""
Fast implementation using PyTorch's optimized operations
with strategic operation fusion
"""
with DisableCuDNNTF32():
input_tensor, mask, weights, config = data
B, N, _, D = input_tensor.shape
H = config["hidden_dim"]
# Flatten and normalize - PyTorch's LayerNorm is highly optimized
x_flat = input_tensor.reshape(B * N * N, D)
x_norm = F.layer_norm(x_flat, [D], weights['norm.weight'], weights['norm.bias'])
# Batch all projections together for better GPU utilization
# Stack weights for a single large matmul
all_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) # Shape: [5*H, D]
# Single batched linear operation
all_proj = F.linear(x_norm, all_weights) # Shape: [B*N*N, 5*H]
all_proj = all_proj.view(B, N, N, 5, H)
# Split projections
left_proj = all_proj[..., 0, :]
right_proj = all_proj[..., 1, :]
left_gate = all_proj[..., 2, :].sigmoid()
right_gate = all_proj[..., 3, :].sigmoid()
out_gate = all_proj[..., 4, :].sigmoid()
# Apply mask and gates in a fused manner
mask_expanded = mask.unsqueeze(-1)
left = left_proj * mask_expanded * left_gate
right = right_proj * mask_expanded * right_gate
# Einsum - PyTorch's implementation is highly optimized
out = torch.einsum('bikd,bjkd->bijd', left, right)
# Output processing
out_flat = out.reshape(B * N * N, H)
out_norm = F.layer_norm(out_flat, [H],
weights['to_out_norm.weight'],
weights['to_out_norm.bias'])
out_norm = out_norm.view(B, N, N, H)
# Apply gate and final projection
out_gated = out_norm * out_gate
out_flat = out_gated.reshape(B * N * N, H)
output = F.linear(out_flat, weights['to_out.weight'])
return output.view(B, N, N, D)scrolls · 65 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 35291.
⋯ 3 unchanged linesfrom task import input_t, output_timport torch- import torch.nn as nn- # import torch.nn.functional as F # Not used- import triton- import triton.language as tl- import math+ import torch.nn.functional as F- # 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):+ def custom_kernel(data: input_t) -> output_t:"""- Triton-optimized TriMul implementation - exact copy from working H100FlashTriMul+ Fast implementation using PyTorch's optimized operations+ with strategic operation fusion"""-- def __init__(self, dim: int, hidden_dim: int):- super().__init__()- self.dim = dim- self.hidden_dim = hidden_dim+ with DisableCuDNNTF32():+ input_tensor, mask, weights, config = data- # 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)+ B, N, _, D = input_tensor.shape+ H = config["hidden_dim"]- # 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)+ # Flatten and normalize - PyTorch's LayerNorm is highly optimized+ x_flat = input_tensor.reshape(B * N * N, D)+ x_norm = F.layer_norm(x_flat, [D], weights['norm.weight'], weights['norm.bias'])- def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:- B, N, _, _ = x.shape- H = self.hidden_dim+ # Batch all projections together for better GPU utilization+ # Stack weights for a single large matmul+ all_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) # Shape: [5*H, D]- # LayerNorm and projection- x = self.norm(x)- all_features = self.fused_proj_gates(x)+ # Single batched linear operation+ all_proj = F.linear(x_norm, all_weights) # Shape: [B*N*N, 5*H]+ all_proj = all_proj.view(B, N, N, 5, H)- # Simple reshape and extraction- features = all_features.view(B, N, N, 5, H)+ # Split projections+ left_proj = all_proj[..., 0, :]+ right_proj = all_proj[..., 1, :]+ left_gate = all_proj[..., 2, :].sigmoid()+ right_gate = all_proj[..., 3, :].sigmoid()+ out_gate = all_proj[..., 4, :].sigmoid()- # 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+ # Apply mask and gates in a fused mannermask_expanded = mask.unsqueeze(-1)left = left_proj * mask_expanded * left_gateright = right_proj * mask_expanded * right_gate- # Make contiguous for einsum- left = left.contiguous()- right = right.contiguous()+ # Einsum - PyTorch's implementation is highly optimized+ out = torch.einsum('bikd,bjkd->bijd', left, right)- # 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)+ # Output processing+ out_flat = out.reshape(B * N * N, H)+ out_norm = F.layer_norm(out_flat, [H],+ weights['to_out_norm.weight'],+ weights['to_out_norm.bias'])+ out_norm = out_norm.view(B, N, N, H)- # Fused output processing- # Combine normalization, gating, and projection- out = self.to_out(self.out_norm(out) * out_gate)+ # Apply gate and final projection+ out_gated = out_norm * out_gate+ out_flat = out_gated.reshape(B * N * N, H)+ output = F.linear(out_flat, weights['to_out.weight'])- 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)No newline at end of file+ return output.view(B, N, N, D)No newline at end of file
scrolls · 401 diff lines total
Best evidence level for this revision: reported
JSON