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
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.
mma
acc_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 H100from utils import make_match_reference, DisableCuDNNTF32from task import input_t, output_timport torch- from torch import nn- import math+ import torch.nn as nn+ # import torch.nn.functional as F # Not usedimport 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 IDspid_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 = dimself.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.shapeH = 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 outdef 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 weightswith 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 linesweights['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_weightsmodel.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 tensorsinput_tensor = input_tensor.contiguous()mask = mask.contiguous()- # Run inference+ # Run model directly without CUDA graph to avoid timeoutwith torch.no_grad():- # No autocast - maintain FP32 precision for accuracyoutput = model(input_tensor, mask)return output- # Reference implementation+ # Reference implementation for testingclass TriMul(nn.Module):def __init__(self, dim: int, hidden_dim: int):super().__init__()⋯ 88 unchanged linesreturn (input_tensor, mask, weights, config)+ # Check implementation correctnesscheck_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