submission 45358
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-45358?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:d29ca5afacb4c065c09c9b5aa114e0aab391531a542d4572f9f467b76cd9b257
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 45338.
Best evidence level for this revision: reported
JSON