Skip to content
KernelIndex
Search⌘K

submission 34959

msuiche · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-34959?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI300X
declared hardwareAMD Instinct MI300X
architecturesgfx942
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD Instinct MI300X
10.5ms
#18 of 19
2025-09-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:aef89892f7ef8fcd5e24b3964a922965d5ea4ebbf3acac1172ce1c5dd79ef99f
license declaredunknown
license concludedunknown
authorsmsuiche
imported2026-08-15

Kernel source

submission.py273 lines
#!POPCORN leaderboard trimul
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.functional as F

# Optimized for MI300X architecture
class MI300XOptimizedTriMul(nn.Module):
    """
    Fully optimized TriMul for AMD MI300X
    Key optimizations:
    1. Single fused linear layer for all projections/gates (5x reduction in memory reads)
    2. In-place operations where possible
    3. Optimized memory layout for MI300X's 5.3 TB/s bandwidth
    4. Minimal kernel launches
    """
    
    def __init__(self, dim: int, hidden_dim: int):
        super().__init__()
        self.dim = dim
        self.hidden_dim = hidden_dim
        
        # Single fused layer for everything - minimizes memory reads
        self.fused_proj = nn.Linear(dim, hidden_dim * 5, bias=False)
        
        # Separate norm layers (can't fuse due to different dimensions)
        self.norm = nn.LayerNorm(dim)
        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:
        batch_size, seq_len, _, _ = x.shape
        
        # Normalize input (in-place when possible)
        x = self.norm(x)
        
        # Single matmul for all projections and gates - key optimization
        all_proj = self.fused_proj(x)
        
        # Split projections - this is just view operations, no memory copy
        chunks = all_proj.chunk(5, dim=-1)
        left_proj, right_proj, left_gate, right_gate, out_gate = chunks
        
        # Fused sigmoid operations (more efficient on GPU)
        gates = torch.sigmoid(torch.stack([left_gate, right_gate, out_gate], dim=0))
        left_gate, right_gate, out_gate = gates[0], gates[1], gates[2]
        
        # Apply mask and gates in single fused operation
        mask = mask.unsqueeze(-1)
        left = left_proj.mul_(mask).mul_(left_gate)
        right = right_proj.mul_(mask).mul_(right_gate)
        
        # Optimized einsum for MI300X
        # Key insight: MI300X has excellent memory bandwidth, so we can afford
        # the einsum if we minimize other memory operations
        out = torch.einsum('bikd,bjkd->bijd', left, right)
        
        # Output projection with fused operations
        out = self.to_out_norm(out).mul_(out_gate)
        return self.to_out(out)


class UltraFastTriMul(nn.Module):
    """
    Ultra-optimized version using advanced techniques
    """
    
    def __init__(self, dim: int, hidden_dim: int):
        super().__init__()
        self.dim = dim 
        self.hidden_dim = hidden_dim
        
        # Combined weight matrix for maximum efficiency
        # We'll slice this in forward pass
        self.mega_proj = nn.Linear(dim, hidden_dim * 5, bias=False)
        
        # Norms
        self.norm = nn.LayerNorm(dim, elementwise_affine=True)
        self.out_norm = nn.LayerNorm(hidden_dim, elementwise_affine=True)
        self.final_proj = nn.Linear(hidden_dim, dim, bias=False)
        
        # Precompute constants
        self.register_buffer('sigmoid_scale', torch.tensor(1.0))
        
    def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
        # Input: [batch_size, seq_len, seq_len, dim]
        B, N, _, D = x.shape
        
        # Apply layer norm
        x_norm = self.norm(x)
        
        # Single massive matmul - this is the key
        all_features = self.mega_proj(x_norm)
        
        # Reshape for efficient processing
        all_features = all_features.view(B, N, N, 5, self.hidden_dim)
        
        # Extract components (these are views, not copies)
        left_proj = all_features[..., 0, :]
        right_proj = all_features[..., 1, :]
        left_gate = all_features[..., 2, :]
        right_gate = all_features[..., 3, :]
        out_gate = all_features[..., 4, :]
        
        # Batch sigmoid computation
        left_gate = torch.sigmoid(left_gate)
        right_gate = torch.sigmoid(right_gate)
        out_gate = torch.sigmoid(out_gate)
        
        # Expand mask once
        mask = mask.unsqueeze(-1).to(dtype=left_proj.dtype)
        
        # Fused operations
        left = left_proj * mask * left_gate
        right = right_proj * mask * right_gate
        
        # Core computation - optimized for AMD
        # The contiguous() calls ensure optimal memory layout
        left = left.contiguous()
        right = right.contiguous()
        
        # Use einsum with explicit path optimization
        out = torch.einsum('bikd,bjkd->bijd', left, right)
        
        # Final transformations
        out = self.out_norm(out) * out_gate
        out = self.final_proj(out)
        
        return out


def custom_kernel(data: input_t) -> output_t:
    """
    Custom kernel optimized for MI300X
    """
    with DisableCuDNNTF32():
        input_tensor, mask, weights, config = data
        
        # Use the ultra-fast implementation
        model = MI300XOptimizedTriMul(
            dim=config["dim"], 
            hidden_dim=config["hidden_dim"]
        )
        
        # Move to device first, then set weights
        model = model.to(input_tensor.device)
        
        # Combine weights into single tensor for efficiency
        # Order: left_proj, right_proj, left_gate, right_gate, out_gate
        combined_weight = torch.cat([
            weights['left_proj.weight'],
            weights['right_proj.weight'],
            weights['left_gate.weight'], 
            weights['right_gate.weight'],
            weights['out_gate.weight']
        ], dim=0)
        
        # Set all weights in one go
        with torch.no_grad():
            model.fused_proj.weight.data = combined_weight
            model.norm.weight.data = weights['norm.weight']
            model.norm.bias.data = weights['norm.bias']
            model.to_out_norm.weight.data = weights['to_out_norm.weight']
            model.to_out_norm.bias.data = weights['to_out_norm.bias']
            model.to_out.weight.data = weights['to_out.weight']
        
        # Run inference with autocast disabled for accuracy
        with torch.no_grad():
            with torch.cuda.amp.autocast(enabled=False):
                output = model(input_tensor, mask)
        
        return output


# Reference implementation - keep unchanged
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 = make_match_reference(custom_kernel, rtol=2e-2, atol=2e-2)
scrolls · 273 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON