Skip to content
KernelIndex
Search⌘K

submission 35286

msuiche · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 372 lines, June 9 Researcher Reciprocity License v1.0.

triton_optimized.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-35286?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.9ms
#60 of 71
2025-09-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6063b7eb3eb10a0d5b0331b21e432288e14c862844650ba9be55ae72e4185e9f
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.py372 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

# 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(
    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, use_checkpoint: bool = False):
        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
        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 _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)
        
        # 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
        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
        
        # Use JIT-compiled fusion for gate application
        left, right, out_gate = fused_gate_application(projections, gates_logits, mask)
        
        # 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 · 372 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 35087.

#!POPCORN leaderboard trimul
+ #!POPCORN gpu H100
from utils import make_match_reference, DisableCuDNNTF32
from task import input_t, output_t
import torch
- from torch import nn
- import math
+ import torch.nn as nn
+ # import torch.nn.functional as F # Not used
import triton
- import triton.language as tl
+ 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 layer_norm_kernel(
- x_ptr, out_ptr, weight_ptr, bias_ptr,
- N, eps,
- BLOCK_SIZE: tl.constexpr
+ 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
):
- """Fused layer normalization kernel"""
- row = tl.program_id(0)
+ pid_b = tl.program_id(0) # batch index
+ pid_n = tl.program_id(1) # N index (combined n1*n2)
- # Compute mean
- mean = 0.0
- for idx in range(0, N, BLOCK_SIZE):
- cols = idx + tl.arange(0, BLOCK_SIZE)
- mask = cols < N
- x = tl.load(x_ptr + row * N + cols, mask=mask, other=0.0)
- mean += tl.sum(x, axis=0)
- mean = mean / N
+ offs_k = tl.arange(0, BLOCK_K)
- # Compute variance
- var = 0.0
- for idx in range(0, N, BLOCK_SIZE):
- cols = idx + tl.arange(0, BLOCK_SIZE)
- mask = cols < N
- x = tl.load(x_ptr + row * N + cols, mask=mask, other=0.0)
- var += tl.sum((x - mean) * (x - mean), axis=0)
- var = var / N
+ # Compute mean and variance
+ sum_x = tl.zeros((1,), dtype=tl.float32)
+ sum_x2 = tl.zeros((1,), dtype=tl.float32)
- # Normalize and apply weight/bias
- rstd = 1.0 / tl.sqrt(var + eps)
- for idx in range(0, N, BLOCK_SIZE):
- cols = idx + tl.arange(0, BLOCK_SIZE)
- mask = cols < N
- x = tl.load(x_ptr + row * N + cols, mask=mask)
- w = tl.load(weight_ptr + cols, mask=mask)
- b = tl.load(bias_ptr + cols, mask=mask)
- out = (x - mean) * rstd * w + b
- tl.store(out_ptr + row * N + cols, out, mask=mask)
+ 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_forward_kernel(
- # Input tensors
+ def trimul_fused_kernel(
x_ptr, mask_ptr,
- # Weight pointers
- norm_w_ptr, norm_b_ptr,
- fused_proj_ptr, # All projections/gates in one weight matrix
+ ln_w_ptr, ln_b_ptr,
+ proj_gates_w_ptr,
out_norm_w_ptr, out_norm_b_ptr,
- final_proj_ptr,
- # Output
- output_ptr,
- # Dimensions
- batch_size, seq_len, dim, hidden_dim,
- # Strides
- stride_xb, stride_xi, stride_xj, stride_xd,
- stride_mb, stride_mi, stride_mj,
- stride_ob, stride_oi, stride_oj, stride_od,
- # Block configuration
- BLOCK_B: tl.constexpr,
- BLOCK_I: tl.constexpr,
- BLOCK_J: tl.constexpr,
- BLOCK_K: tl.constexpr,
- BLOCK_D: tl.constexpr,
+ 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
):
- """
- Fully fused TriMul kernel with branched masking
- Computes the entire TriMul operation in a single kernel
- """
- # Program IDs
pid_b = tl.program_id(0)
- pid_ij = tl.program_id(1)
- pid_d = tl.program_id(2)
+ pid_m = tl.program_id(1) # N index for output
+ pid_n = tl.program_id(2) # H index for output
- # Compute i, j indices from flattened pid_ij
- pid_i = pid_ij // (seq_len // BLOCK_J)
- pid_j = pid_ij % (seq_len // BLOCK_J)
+ 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)
- # Block offsets
- offs_b = pid_b * BLOCK_B + tl.arange(0, BLOCK_B)
- offs_i = pid_i * BLOCK_I + tl.arange(0, BLOCK_I)
- offs_j = pid_j * BLOCK_J + tl.arange(0, BLOCK_J)
- offs_d = pid_d * BLOCK_D + tl.arange(0, BLOCK_D)
+ mask_m = offs_m < N
+ mask_n = offs_n < H
- # Masks for bounds checking
- mask_b = offs_b < batch_size
- mask_i = offs_i < seq_len
- mask_j = offs_j < seq_len
- mask_d = offs_d < hidden_dim
+ # 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 accumulator for einsum
- acc = tl.zeros([BLOCK_I, BLOCK_J, BLOCK_D], dtype=tl.float32)
+ # 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)
- # Loop over K dimension (contraction dimension)
- for k_start in range(0, seq_len, BLOCK_K):
- offs_k = k_start + tl.arange(0, BLOCK_K)
- mask_k = offs_k < seq_len
+ # 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 mask values
- mask_ik_ptr = mask_ptr + offs_b[:, None, None] * stride_mb + \
- offs_i[None, :, None] * stride_mi + \
- offs_k[None, None, :] * stride_mj
- mask_jk_ptr = mask_ptr + offs_b[:, None, None] * stride_mb + \
- offs_j[None, :, None] * stride_mi + \
- offs_k[None, None, :] * stride_mj
-
- mask_ik = tl.load(mask_ik_ptr,
- mask=mask_b[:, None, None] & mask_i[None, :, None] & mask_k[None, None, :],
- other=0.0)
- mask_jk = tl.load(mask_jk_ptr,
- mask=mask_b[:, None, None] & mask_j[None, :, None] & mask_k[None, None, :],
- other=0.0)
+ # 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)
- # Branching: skip computation if mask is zero
- # This is the key optimization for sparse masks
- if tl.sum(mask_ik) > 0 and tl.sum(mask_jk) > 0:
- # Load left[b, i, k, d] with LayerNorm applied
- left_ptr = x_ptr + offs_b[:, None, None, None] * stride_xb + \
- offs_i[None, :, None, None] * stride_xi + \
- offs_k[None, None, :, None] * stride_xj + \
- offs_d[None, None, None, :] * stride_xd
-
- left = tl.load(left_ptr,
- mask=mask_b[:, None, None, None] & mask_i[None, :, None, None] &
- mask_k[None, None, :, None] & mask_d[None, None, None, :],
- other=0.0)
-
- # Load right[b, j, k, d]
- right_ptr = x_ptr + offs_b[:, None, None, None] * stride_xb + \
- offs_j[None, :, None, None] * stride_xi + \
- offs_k[None, None, :, None] * stride_xj + \
- offs_d[None, None, None, :] * stride_xd
-
- right = tl.load(right_ptr,
- mask=mask_b[:, None, None, None] & mask_j[None, :, None, None] &
- mask_k[None, None, :, None] & mask_d[None, None, None, :],
- other=0.0)
-
- # Apply masks
- left = left * mask_ik[:, :, :, None]
- right = right * mask_jk[:, :, :, None]
-
- # Accumulate einsum: sum over batch and k dimensions
- for b in range(BLOCK_B):
- if offs_b[b] < batch_size:
- acc += tl.sum(left[b, :, :, :, None] * right[b, None, :, :, :], axis=2)
+ # 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))
- # Store output
- output_offs = offs_b[:, None, None, None] * stride_ob + \
- offs_i[None, :, None, None] * stride_oi + \
- offs_j[None, None, :, None] * stride_oj + \
- offs_d[None, None, None, :] * stride_od
+ # 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)
- output_mask = mask_b[:, None, None, None] & mask_i[None, :, None, None] & \
- mask_j[None, None, :, None] & mask_d[None, None, None, :]
+ # 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)
- # Average over batch dimension before storing
- acc_avg = acc / BLOCK_B
+ # 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
- for b in range(BLOCK_B):
- if offs_b[b] < batch_size:
- tl.store(output_ptr + output_offs[b], acc_avg, mask=output_mask[b])
+ 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)
- @triton.jit
- def trimul_ultra_fused_kernel(
- # Inputs
- x_ptr, mask_ptr,
- # All weights concatenated
- weights_ptr,
- # Output
- output_ptr,
- # Dimensions
- B, N, D, H,
- # Strides for x [B, N, N, D]
- sx_b, sx_i, sx_j, sx_d,
- # Strides for mask [B, N, N]
- sm_b, sm_i, sm_j,
- # Strides for output [B, N, N, H]
- so_b, so_i, so_j, so_h,
- # Fusion config
- TILE_I: tl.constexpr,
- TILE_J: tl.constexpr,
- TILE_K: tl.constexpr,
- TILE_H: tl.constexpr,
- ):
- """
- Ultra-optimized fused kernel that performs:
- 1. LayerNorm
- 2. Projections and gates
- 3. Masked einsum
- 4. Output projection
- All in a single kernel pass
- """
- pid = tl.program_id(0)
- grid_i = (N + TILE_I - 1) // TILE_I
- grid_j = (N + TILE_J - 1) // TILE_J
-
- # Decode 2D grid position
- pid_i = pid // grid_j
- pid_j = pid % grid_j
-
- # Tile boundaries
- i_start = pid_i * TILE_I
- j_start = pid_j * TILE_J
-
- # Initialize accumulator
- acc = tl.zeros([TILE_I, TILE_J, TILE_H], dtype=tl.float32)
-
- # Main loop over K dimension
- for k in range(0, N, TILE_K):
- # Load tiles with boundary checks
- for ti in range(TILE_I):
- for tj in range(TILE_J):
- for tk in range(TILE_K):
- i = i_start + ti
- j = j_start + tj
- kk = k + tk
-
- if i < N and j < N and kk < N:
- # Load and apply mask
- mask_val = tl.load(mask_ptr + sm_i * i + sm_j * kk)
-
- if mask_val > 0: # Branch on mask
- # Load input and apply transformations
- for h in range(TILE_H):
- if h < H:
- # Fused computation
- val_i = tl.load(x_ptr + sx_i * i + sx_j * kk + sx_d * (h % D))
- val_j = tl.load(x_ptr + sx_i * j + sx_j * kk + sx_d * (h % D))
-
- # Apply mask and accumulate
- acc[ti, tj, h] += val_i * val_j * mask_val
-
- # Store results
- for ti in range(TILE_I):
- for tj in range(TILE_J):
- i = i_start + ti
- j = j_start + tj
-
- if i < N and j < N:
- for h in range(TILE_H):
- if h < H:
- out_idx = so_i * i + so_j * j + so_h * h
- tl.store(output_ptr + out_idx, acc[ti, tj, h])
-
-
class TritonTriMul(nn.Module):
"""
- Triton-accelerated TriMul implementation
- Achieves ~4x speedup over PyTorch implementation
+ Triton-optimized TriMul implementation - exact copy from working H100FlashTriMul
"""
- def __init__(self, dim: int, hidden_dim: int):
+ def __init__(self, dim: int, hidden_dim: int, use_checkpoint: bool = False):
super().__init__()
self.dim = dim
self.hidden_dim = hidden_dim
+ self.use_checkpoint = use_checkpoint
- # CRITICAL OPTIMIZATION: Use Linear instead of Parameter for better BLAS
- # MI300X has optimized rocBLAS that Linear layers leverage better
- self.mega_proj = nn.Linear(dim, hidden_dim * 5, bias=False)
+ # 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 (can't fuse different dimensions easily)
- self.norm = nn.LayerNorm(dim)
- self.out_norm = nn.LayerNorm(hidden_dim)
- self.final_proj = nn.Linear(hidden_dim, dim, 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 _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, _, D = x.shape
+ B, N, _, _ = x.shape
H = self.hidden_dim
- # Normalize
- x = self.norm(x)
+ # Compute features directly (no checkpointing for inference)
+ all_features = self._compute_features(x)
- # OPTIMIZATION 1: Use Linear layer for better MI300X GEMM performance
- all_features = self.mega_proj(x)
+ # Memory-efficient reshape without copying data
+ # First reshape to optimal layout for extraction
+ all_features = all_features.view(B, N * N, 5, H)
- # OPTIMIZATION 2: Efficient reshape (view is zero-copy)
- all_features = all_features.view(B, N, N, 5, H)
+ # Extract components with memory-efficient views
+ features = all_features.view(B, N, N, 5, H)
- # Extract components (views, not copies)
- left_proj = all_features[..., 0, :]
- right_proj = all_features[..., 1, :]
+ # Fused extraction and activation using JIT-compiled function
+ projections = features[..., :2, :] # left and right projections
+ gates_logits = features[..., 2:, :] # all gate logits
- # OPTIMIZATION 3: Batch sigmoid for better GPU utilization
- gates = torch.sigmoid(all_features[..., 2:, :])
- left_gate = gates[..., 0, :]
- right_gate = gates[..., 1, :]
- out_gate = gates[..., 2, :]
+ # Use JIT-compiled fusion for gate application
+ left, right, out_gate = fused_gate_application(projections, gates_logits, mask)
- # OPTIMIZATION 4: Type casting and efficient masking
- mask = mask.unsqueeze(-1).to(left_proj.dtype)
+ # 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)
- # OPTIMIZATION 5: Fully fused multiplication using single operation
- # This ensures the compiler can fuse into a single kernel
- left = left_proj * mask * left_gate # Single fused elementwise kernel
- right = right_proj * mask * right_gate # Single fused elementwise kernel
+ # Fused output processing
+ # Combine normalization, gating, and projection
+ out = self.to_out(self.out_norm(out) * out_gate)
- # OPTIMIZATION 6: Ensure contiguous for optimal einsum
- left = left.contiguous()
- right = right.contiguous()
-
- # Core computation - einsum is still most efficient for this pattern
- output = torch.einsum('bikd,bjkd->bijd', left, right)
-
- # OPTIMIZATION 7: Fused output processing
- output = self.out_norm(output).mul(out_gate)
-
- return self.final_proj(output)
+ return out
def custom_kernel(data: input_t) -> output_t:
"""
- Custom kernel using Triton acceleration
+ Custom kernel implementation for TriMul using Triton optimization
"""
with DisableCuDNNTF32():
input_tensor, mask, weights, config = data
- model = TritonTriMul(
- dim=config["dim"],
- hidden_dim=config["hidden_dim"]
- ).to(input_tensor.device)
+ dim = config["dim"]
+ hidden_dim = config["hidden_dim"]
- # Efficient weight loading for Linear layer
+ # 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():
- # Combine all projection/gate weights for mega_proj
- combined_weights = torch.cat([
+ # Stack all projection and gate weights
+ proj_weights = torch.cat([
weights['left_proj.weight'],
weights['right_proj.weight'],
weights['left_gate.weight'],
⋯ 1 unchanged lines
weights['out_gate.weight']
], dim=0)
- # Direct data assignment is faster than Parameter wrapping
- model.mega_proj.weight.data = combined_weights
+ 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.final_proj.weight.data = weights['to_out.weight']
+ model.to_out.weight.data = weights['to_out.weight']
- # MI300X-specific optimizations
- torch.backends.cudnn.benchmark = True # Auto-tune for best kernels
-
- # Ensure optimal tensor layout
+ # Ensure contiguous tensors
input_tensor = input_tensor.contiguous()
mask = mask.contiguous()
- # Run inference
+ # Run model directly without CUDA graph to avoid timeout
with torch.no_grad():
- # No autocast - maintain FP32 precision for accuracy
output = model(input_tensor, mask)
return output
- # Reference implementation
+ # Reference implementation for testing
class TriMul(nn.Module):
def __init__(self, dim: int, hidden_dim: int):
super().__init__()
⋯ 88 unchanged lines
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
scrolls · 570 diff lines total

Best evidence level for this revision: reported

JSON