Skip to content
KernelIndex
Search⌘K

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
NVIDIA H100
11.6ms
#59 of 71
2025-09-05

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.

mmaacc_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 lines
import triton.language as tl
import 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.jit
def layernorm_stats_kernel(
⋯ 138 unchanged lines
Triton-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 = dim
self.hidden_dim = hidden_dim
- self.use_checkpoint = use_checkpoint
# Single fused weight matrix for all linear ops
# This reduces memory accesses
⋯ 5 unchanged lines
self.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.shape
H = 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 extraction
features = 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 precision
if left.dtype != torch.float32:
left_f32 = left.float()
scrolls · 85 diff lines total

Best evidence level for this revision: reported

JSON