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
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