submission 45338
leymore4172 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 107 lines, June 9 Researcher Reciprocity License v1.0.
trimul.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-45338?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:3ce393c5225c7b4f96472ae0f9d6c24442f78d126f293cbad913226791d6bf11
license declaredunknown
license concludedunknown
authorsleymore4172
imported2026-08-15
Kernel source
trimul.py107 lines
import torch
import torch.nn.functional as F
from torch import einsum
from task import input_t, output_t
@torch.jit.script
def trimul(x: torch.Tensor, mask: torch.Tensor,
norm_weight: torch.Tensor, norm_bias: torch.Tensor,
left_proj_weight: torch.Tensor, right_proj_weight: torch.Tensor,
left_gate_weight: torch.Tensor, right_gate_weight: torch.Tensor,
out_gate_weight: torch.Tensor,
to_out_norm_weight: torch.Tensor, to_out_norm_bias: torch.Tensor,
to_out_weight: torch.Tensor) -> torch.Tensor:
"""
Functional implementation of TriMul.
Args:
x: [bs, seq_len, seq_len, dim]
mask: [bs, seq_len, seq_len]
norm_weight: Layer norm weight
norm_bias: Layer norm bias
left_proj_weight: Left projection weight
right_proj_weight: Right projection weight
left_gate_weight: Left gate weight
right_gate_weight: Right gate weight
out_gate_weight: Output gate weight
to_out_norm_weight: Output layer norm weight
to_out_norm_bias: Output layer norm bias
to_out_weight: Final output projection weight
Returns:
output: [bs, seq_len, seq_len, dim]
"""
batch_size, seq_len, _, dim = x.shape
# Layer normalization
x = F.layer_norm(x, [dim], norm_weight, norm_bias)
# Linear projections without bias
left = F.linear(x, left_proj_weight)
right = F.linear(x, right_proj_weight)
# Apply mask
mask = mask.unsqueeze(-1)
left = left * mask
right = right * mask
# Gate computations
left_gate = F.linear(x, left_gate_weight).sigmoid()
right_gate = F.linear(x, right_gate_weight).sigmoid()
out_gate = F.linear(x, out_gate_weight).sigmoid()
# Apply gates
left = left * left_gate
right = right * right_gate
# Einstein summation
out = einsum('... i k d, ... j k d -> ... i j d', left.to(torch.bfloat16), right.to(torch.bfloat16))
# This einsum is the same as the following:
# out = torch.zeros(batch_size, seq_len, seq_len, dim, device=x.device)
# # 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)
# Output normalization
hidden_dim = out.shape[-1]
out = F.layer_norm(out, [hidden_dim], to_out_norm_weight, to_out_norm_bias)
out = out * out_gate
# Final linear projection
out = F.linear(out, to_out_weight)
return out
def custom_kernel(data: input_t) -> output_t:
"""
Reference implementation of TriMul using PyTorch functional interface.
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
output = trimul(
input_tensor, mask,
weights['norm.weight'], weights['norm.bias'],
weights['left_proj.weight'], weights['right_proj.weight'],
weights['left_gate.weight'], weights['right_gate.weight'],
weights['out_gate.weight'],
weights['to_out_norm.weight'], weights['to_out_norm.bias'],
weights['to_out.weight']
)
return output
scrolls · 107 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 45318.
import torch- from torch import nn, einsum+ import torch.nn.functional as F+ from torch import einsumfrom task import input_t, output_t- # Global cache for JIT compiled models based on dimensions- _jit_model_cache = {}+ @torch.jit.script+ def trimul(x: torch.Tensor, mask: torch.Tensor,+ norm_weight: torch.Tensor, norm_bias: torch.Tensor,+ left_proj_weight: torch.Tensor, right_proj_weight: torch.Tensor,+ left_gate_weight: torch.Tensor, right_gate_weight: torch.Tensor,+ out_gate_weight: torch.Tensor,+ to_out_norm_weight: torch.Tensor, to_out_norm_bias: torch.Tensor,+ to_out_weight: torch.Tensor) -> torch.Tensor:+ """+ Functional implementation of TriMul.- class TriMul(nn.Module):- def __init__(- self,- dim: int,- hidden_dim: int,- ):- super().__init__()+ Args:+ x: [bs, seq_len, seq_len, dim]+ mask: [bs, seq_len, seq_len]+ norm_weight: Layer norm weight+ norm_bias: Layer norm bias+ left_proj_weight: Left projection weight+ right_proj_weight: Right projection weight+ left_gate_weight: Left gate weight+ right_gate_weight: Right gate weight+ out_gate_weight: Output gate weight+ to_out_norm_weight: Output layer norm weight+ to_out_norm_bias: Output layer norm bias+ to_out_weight: Final output projection weight- self.norm = nn.LayerNorm(dim)+ Returns:+ output: [bs, seq_len, seq_len, dim]+ """+ batch_size, seq_len, _, dim = x.shape- 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)+ # Layer normalization+ x = F.layer_norm(x, [dim], norm_weight, norm_bias)- 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)+ # Linear projections without bias+ left = F.linear(x, left_proj_weight)+ right = F.linear(x, right_proj_weight)- self.to_out_norm = nn.LayerNorm(hidden_dim)- self.to_out = nn.Linear(hidden_dim, dim, bias=False, dtype=torch.float32)+ # Apply mask+ mask = mask.unsqueeze(-1)+ left = left * mask+ right = right * mask- def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:- """- x: [bs, seq_len, seq_len, dim]- mask: [bs, seq_len, seq_len]+ # Gate computations+ left_gate = F.linear(x, left_gate_weight).sigmoid()+ right_gate = F.linear(x, right_gate_weight).sigmoid()+ out_gate = F.linear(x, out_gate_weight).sigmoid()- Returns:- output: [bs, seq_len, seq_len, dim]- """- # Optimize: Remove redundant shape unpacking- x = self.norm(x)- x = x.to(torch.float32)+ # Apply gates+ left = left * left_gate+ right = right * right_gate- # Optimize: Single type conversion for x- x_float32 = x # Already converted above+ # Einstein summation+ out = einsum('... i k d, ... j k d -> ... i j d', left.to(torch.bfloat16), right.to(torch.bfloat16))+ # This einsum is the same as the following:+ # out = torch.zeros(batch_size, seq_len, seq_len, dim, device=x.device)- left = self.left_proj(x_float32)- right = self.right_proj(x_float32)+ # # 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, :]- mask = mask.unsqueeze(-1)- left = left * mask- right = right * mask+ out = out.to(torch.float32)- left_gate = self.left_gate(x_float32).sigmoid()- right_gate = self.right_gate(x_float32).sigmoid()- out_gate = self.out_gate(x_float32).sigmoid()+ # Output normalization+ hidden_dim = out.shape[-1]+ out = F.layer_norm(out, [hidden_dim], to_out_norm_weight, to_out_norm_bias)+ out = out * out_gate- left = left * left_gate- right = right * right_gate+ # Final linear projection+ out = F.linear(out, to_out_weight)- out = einsum('... i k d, ... j k d -> ... i j d', left.to(torch.bfloat16), right.to(torch.bfloat16))+ return out- 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.+ Reference implementation of TriMul using PyTorch functional interface.Args:data: Tuple of (input: torch.Tensor, mask: torch.Tensor, weights: Dict[str, torch.Tensor], config: Dict)⋯ 4 unchanged lines"""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))+ output = trimul(+ input_tensor, mask,+ weights['norm.weight'], weights['norm.bias'],+ weights['left_proj.weight'], weights['right_proj.weight'],+ weights['left_gate.weight'], weights['right_gate.weight'],+ weights['out_gate.weight'],+ weights['to_out_norm.weight'], weights['to_out_norm.bias'],+ weights['to_out.weight']+ )- # 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 · 192 diff lines total
Best evidence level for this revision: reported
JSON