submission 45660
leymore4172 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 115 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-45660?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:236671cb151bcc3b21a54668f93d2189441b37eb597c2055b3efaa045cf3ae3d
license declaredunknown
license concludedunknown
authorsleymore4172
imported2026-08-15
Kernel source
submission.py115 lines
import torch
import torch.nn.functional as F
from torch import einsum
from task import input_t, output_t
@torch.compile
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)
# Fuse all linear projections with x as input
# Concatenate all weights
fused_weight = torch.cat([left_proj_weight, right_proj_weight,
left_gate_weight, right_gate_weight,
out_gate_weight], dim=0)
# Single fused linear operation
fused_output = F.linear(x, fused_weight)
# fused_output = F.linear(x.to(torch.bfloat16), fused_weight.to(torch.bfloat16)).to(torch.float32)
# Split the results
hidden_dim = left_proj_weight.shape[0]
left = fused_output[..., :hidden_dim]
right = fused_output[..., hidden_dim:2*hidden_dim]
left_gate = fused_output[..., 2*hidden_dim:3*hidden_dim].sigmoid()
right_gate = fused_output[..., 3*hidden_dim:4*hidden_dim].sigmoid()
out_gate = fused_output[..., 4*hidden_dim:5*hidden_dim].sigmoid()
mask = mask.unsqueeze(-1)
left = left * mask
right = right * mask
# 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 · 115 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 45358.
⋯ 2 unchanged linesfrom torch import einsumfrom task import input_t, output_t- @torch.jit.script+ @torch.compiledef 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,⋯ 26 unchanged lines# Layer normalizationx = 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)+ # Fuse all linear projections with x as input+ # Concatenate all weights+ fused_weight = torch.cat([left_proj_weight, right_proj_weight,+ left_gate_weight, right_gate_weight,+ out_gate_weight], dim=0)- # Apply mask+ # Single fused linear operation+ fused_output = F.linear(x, fused_weight)+ # fused_output = F.linear(x.to(torch.bfloat16), fused_weight.to(torch.bfloat16)).to(torch.float32)++ # Split the results+ hidden_dim = left_proj_weight.shape[0]+ left = fused_output[..., :hidden_dim]+ right = fused_output[..., hidden_dim:2*hidden_dim]+ left_gate = fused_output[..., 2*hidden_dim:3*hidden_dim].sigmoid()+ right_gate = fused_output[..., 3*hidden_dim:4*hidden_dim].sigmoid()+ out_gate = fused_output[..., 4*hidden_dim:5*hidden_dim].sigmoid()+mask = mask.unsqueeze(-1)left = left * maskright = 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 gatesleft = left * left_gateright = right * right_gate
scrolls · 47 diff lines total
Best evidence level for this revision: reported
JSON