Skip to content
KernelIndex
Search⌘K

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
NVIDIA H100
10.2ms
#58 of 71
2025-09-06

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 lines
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
+ 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 manner
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 - 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