submission 45318
leymore4172 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 122 lines, June 9 Researcher Reciprocity License v1.0.
trimul.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-45318?include=source"interfacepython
Compatibility
measured onNVIDIA A100
declared hardwareNVIDIA A100
architecturessm_80
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:f4779bf2e9bcb515ca0d222e1a750289c8ade056958fb6693546cef817fe4797
license declaredunknown
license concludedunknown
authorsleymore4172
imported2026-08-15
Kernel source
trimul.py122 lines
import torch
from torch import nn, einsum
from task import input_t, output_t
# Global cache for JIT compiled models based on dimensions
_jit_model_cache = {}
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, dtype=torch.float32)
self.right_proj = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)
self.left_gate = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)
self.right_gate = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)
self.out_gate = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)
self.to_out_norm = nn.LayerNorm(hidden_dim)
self.to_out = nn.Linear(hidden_dim, dim, bias=False, dtype=torch.float32)
def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
"""
x: [bs, seq_len, seq_len, dim]
mask: [bs, seq_len, seq_len]
Returns:
output: [bs, seq_len, seq_len, dim]
"""
# Optimize: Remove redundant shape unpacking
x = self.norm(x)
x = x.to(torch.float32)
# Optimize: Single type conversion for x
x_float32 = x # Already converted above
left = self.left_proj(x_float32)
right = self.right_proj(x_float32)
mask = mask.unsqueeze(-1)
left = left * mask
right = right * mask
left_gate = self.left_gate(x_float32).sigmoid()
right_gate = self.right_gate(x_float32).sigmoid()
out_gate = self.out_gate(x_float32).sigmoid()
left = left * left_gate
right = right * right_gate
out = einsum('... i k d, ... j k d -> ... i j d', left.to(torch.bfloat16), right.to(torch.bfloat16))
out = out.to(torch.float32)
out = self.to_out_norm(out)
out = out * out_gate
return self.to_out(out)
def custom_kernel(data: input_t) -> output_t:
"""
Reference implementation of TriMul using PyTorch.
Args:
data: Tuple of (input: torch.Tensor, mask: torch.Tensor, weights: Dict[str, torch.Tensor], config: Dict)
- input: Input tensor of shape [batch_size, seq_len, seq_len, dim]
- mask: Mask tensor of shape [batch_size, seq_len, seq_len]
- weights: Dictionary containing model weights
- config: Dictionary containing model configuration parameters
"""
input_tensor, mask, weights, config = data
# Create cache key based on dimensions and device
cache_key = (config["dim"], config["hidden_dim"], str(input_tensor.device))
# Check if we have a JIT compiled model for these dimensions
if cache_key not in _jit_model_cache:
# Create and JIT compile the model
trimul = TriMul(config["dim"], config["hidden_dim"]).to(input_tensor.device)
# Fill in the given weights of the model
trimul.norm.weight = nn.Parameter(weights['norm.weight'].to(torch.float32))
trimul.left_proj.weight = nn.Parameter(weights['left_proj.weight'].to(torch.float32))
trimul.right_proj.weight = nn.Parameter(weights['right_proj.weight'].to(torch.float32))
trimul.left_gate.weight = nn.Parameter(weights['left_gate.weight'].to(torch.float32))
trimul.right_gate.weight = nn.Parameter(weights['right_gate.weight'].to(torch.float32))
trimul.out_gate.weight = nn.Parameter(weights['out_gate.weight'].to(torch.float32))
trimul.to_out_norm.weight = nn.Parameter(weights['to_out_norm.weight'].to(torch.float32))
trimul.to_out.weight = nn.Parameter(weights['to_out.weight'].to(torch.float32))
trimul.norm.bias = nn.Parameter(weights['norm.bias'].to(torch.float32))
trimul.to_out_norm.bias = nn.Parameter(weights['to_out_norm.bias'].to(torch.float32))
# JIT compile the model using script
trimul.eval() # Set to eval mode for consistency
with torch.no_grad():
_jit_model_cache[cache_key] = torch.jit.script(trimul)
else:
# Use cached JIT compiled model and update weights
jit_model = _jit_model_cache[cache_key]
# Update weights in the JIT model
jit_model.norm.weight.data = weights['norm.weight'].to(torch.float32)
jit_model.left_proj.weight.data = weights['left_proj.weight'].to(torch.float32)
jit_model.right_proj.weight.data = weights['right_proj.weight'].to(torch.float32)
jit_model.left_gate.weight.data = weights['left_gate.weight'].to(torch.float32)
jit_model.right_gate.weight.data = weights['right_gate.weight'].to(torch.float32)
jit_model.out_gate.weight.data = weights['out_gate.weight'].to(torch.float32)
jit_model.to_out_norm.weight.data = weights['to_out_norm.weight'].to(torch.float32)
jit_model.to_out.weight.data = weights['to_out.weight'].to(torch.float32)
jit_model.norm.bias.data = weights['norm.bias'].to(torch.float32)
jit_model.to_out_norm.bias.data = weights['to_out_norm.bias'].to(torch.float32)
# Use the JIT compiled model
output = _jit_model_cache[cache_key](input_tensor, mask).to(torch.float32)
return output
scrolls · 122 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 44748.
- from utils import make_match_reference, DisableCuDNNTF32- from task import input_t, output_t-import torchfrom torch import nn, einsum- import math+ from task import input_t, output_t- # Reference code in PyTorch+ # Global cache for JIT compiled models based on dimensions+ _jit_model_cache = {}+class TriMul(nn.Module):- # Based on https://github.com/lucidrains/triangle-multiplicative-module/blob/main/triangle_multiplicative_module/triangle_multiplicative_module.pydef __init__(self,dim: int,⋯ 3 unchanged linesself.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_proj = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)+ self.right_proj = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)- 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.left_gate = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)+ self.right_gate = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)+ self.out_gate = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)self.to_out_norm = nn.LayerNorm(hidden_dim)- self.to_out = nn.Linear(hidden_dim, dim, bias=False)+ self.to_out = nn.Linear(hidden_dim, dim, bias=False, dtype=torch.float32)def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:"""⋯ 3 unchanged linesReturns:output: [bs, seq_len, seq_len, dim]"""- batch_size, seq_len, _, dim = x.shape-+ # Optimize: Remove redundant shape unpackingx = self.norm(x)+ x = x.to(torch.float32)- left = self.left_proj(x)- right = self.right_proj(x)+ # Optimize: Single type conversion for x+ x_float32 = x # Already converted above+ left = self.left_proj(x_float32)+ right = self.right_proj(x_float32)+mask = mask.unsqueeze(-1)left = left * maskright = right * mask- left_gate = self.left_gate(x).sigmoid()- right_gate = self.right_gate(x).sigmoid()- out_gate = self.out_gate(x).sigmoid()+ left_gate = self.left_gate(x_float32).sigmoid()+ right_gate = self.right_gate(x_float32).sigmoid()+ out_gate = self.out_gate(x_float32).sigmoid()left = left * left_gateright = right * right_gate- out = einsum('... i k d, ... j k d -> ... i j d', left, right)- # This einsum is the same as the following:- # out = torch.zeros(batch_size, seq_len, seq_len, dim, device=x.device)+ out = einsum('... i k d, ... j k d -> ... i j d', left.to(torch.bfloat16), right.to(torch.bfloat16))- # # Compute using nested loops- # for b in range(batch_size):- # for i in range(seq_len):- # for j in range(seq_len):- # # Compute each output element- # for k in range(seq_len):- # out[b, i, j] += left[b, i, k, :] * right[b, j, k, :]-+ out = out.to(torch.float32)out = self.to_out_norm(out)out = out * out_gatereturn self.to_out(out)⋯ 10 unchanged lines- weights: Dictionary containing model weights- config: Dictionary containing model configuration parameters"""+ input_tensor, mask, weights, config = data- # Use deterministic kernels and disable TF32 for accuracy- with DisableCuDNNTF32():- input_tensor, mask, weights, config = data- trimul = TriMul(dim=config["dim"], hidden_dim=config["hidden_dim"]).to(input_tensor.device)+ # Create cache key based on dimensions and device+ cache_key = (config["dim"], config["hidden_dim"], str(input_tensor.device))+ # Check if we have a JIT compiled model for these dimensions+ if cache_key not in _jit_model_cache:+ # Create and JIT compile the model+ trimul = TriMul(config["dim"], config["hidden_dim"]).to(input_tensor.device)+# Fill in the given weights of the model- 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'])+ trimul.norm.weight = nn.Parameter(weights['norm.weight'].to(torch.float32))+ trimul.left_proj.weight = nn.Parameter(weights['left_proj.weight'].to(torch.float32))+ trimul.right_proj.weight = nn.Parameter(weights['right_proj.weight'].to(torch.float32))+ trimul.left_gate.weight = nn.Parameter(weights['left_gate.weight'].to(torch.float32))+ trimul.right_gate.weight = nn.Parameter(weights['right_gate.weight'].to(torch.float32))+ trimul.out_gate.weight = nn.Parameter(weights['out_gate.weight'].to(torch.float32))+ trimul.to_out_norm.weight = nn.Parameter(weights['to_out_norm.weight'].to(torch.float32))+ trimul.to_out.weight = nn.Parameter(weights['to_out.weight'].to(torch.float32))+ trimul.norm.bias = nn.Parameter(weights['norm.bias'].to(torch.float32))+ trimul.to_out_norm.bias = nn.Parameter(weights['to_out_norm.bias'].to(torch.float32))- output = trimul(input_tensor, mask)-- return output-- # Input generation for the reference code- def generate_input(- seqlen: int,- bs: int,- dim: int,- hiddendim: int,- seed: int,- nomask: bool,- distribution: str,- ) -> input_t:-- # Really dumb but for now _ isn't parsing correctly.- 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 = {}-- # Generate input tensor based on distribution- if distribution == "cauchy":- # Heavier tail distribution- input_tensor = torch.distributions.Cauchy(0, 2).sample(- (batch_size, seq_len, seq_len, dim)- ).to(device='cuda', dtype=torch.float32)- else: # normal distribution- 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)+ # JIT compile the model using script+ trimul.eval() # Set to eval mode for consistency+ with torch.no_grad():+ _jit_model_cache[cache_key] = torch.jit.script(trimul)else:- mask = torch.randint(0, 2, (batch_size, seq_len, seq_len), device=input_tensor.device, generator=gen)+ # Use cached JIT compiled model and update weights+ jit_model = _jit_model_cache[cache_key]- # Initialize model weights based on distribution- 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)+ # Update weights in the JIT model+ jit_model.norm.weight.data = weights['norm.weight'].to(torch.float32)+ jit_model.left_proj.weight.data = weights['left_proj.weight'].to(torch.float32)+ jit_model.right_proj.weight.data = weights['right_proj.weight'].to(torch.float32)+ jit_model.left_gate.weight.data = weights['left_gate.weight'].to(torch.float32)+ jit_model.right_gate.weight.data = weights['right_gate.weight'].to(torch.float32)+ jit_model.out_gate.weight.data = weights['out_gate.weight'].to(torch.float32)+ jit_model.to_out_norm.weight.data = weights['to_out_norm.weight'].to(torch.float32)+ jit_model.to_out.weight.data = weights['to_out.weight'].to(torch.float32)+ jit_model.norm.bias.data = weights['norm.bias'].to(torch.float32)+ jit_model.to_out_norm.bias.data = weights['to_out_norm.bias'].to(torch.float32)- return (input_tensor, mask, weights, config)+ # Use the JIT compiled model+ output = _jit_model_cache[cache_key](input_tensor, mask).to(torch.float32)-- check_implementation = make_match_reference(custom_kernel, rtol=2e-2, atol=2e-2)+ return output
scrolls · 214 diff lines total
Best evidence level for this revision: reported
JSON