submission 44207
davidberard · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 137 lines, June 9 Researcher Reciprocity License v1.0.
ref.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-44207?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:af5beada86acb530a10562489792424934cb67106d6cd664d31c048b1176af3f
license declaredunknown
license concludedunknown
authorsdavidberard
imported2026-08-15
Kernel source
ref.py137 lines
# from utils import make_match_reference, DisableCuDNNTF32
from task import input_t, output_t
import torch
from torch import nn, einsum
import math
# The flag below controls whether to allow TF32 on matmul. This flag defaults to False
# in PyTorch 1.12 and later.
torch.backends.cuda.matmul.allow_tf32 = True
# The flag below controls whether to allow TF32 on cuDNN. This flag defaults to True.
torch.backends.cudnn.allow_tf32 = True
# Reference code in PyTorch
class TriMul(nn.Module):
# Based on https://github.com/lucidrains/triangle-multiplicative-module/blob/main/triangle_multiplicative_module/triangle_multiplicative_module.py
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: [bs, seq_len, seq_len, dim]
mask: [bs, seq_len, seq_len]
Returns:
output: [bs, seq_len, seq_len, dim]
"""
batch_size, seq_len, _, dim = x.shape
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 = 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)
# # 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 = 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
hidden_dim = config["hidden_dim"]
# trimul = TriMul(dim=config["dim"], hidden_dim=config["hidden_dim"]).to(input_tensor.device)
x = input_tensor
batch_size, seq_len, _, dim = x.shape
x = torch.nn.functional.layer_norm(x, (dim,), eps=1e-5, weight=weights['norm.weight'], bias=weights['norm.bias'])
left = torch.nn.functional.linear(x, weights['left_proj.weight'])
right = torch.nn.functional.linear(x, weights['right_proj.weight'])
left = left * mask.unsqueeze(-1)
right = right * mask.unsqueeze(-1)
left_gate = torch.nn.functional.linear(x, weights['left_gate.weight']).sigmoid()
right_gate = torch.nn.functional.linear(x, weights['right_gate.weight']).sigmoid()
out_gate = torch.nn.functional.linear(x, weights['out_gate.weight']).sigmoid()
left = left * left_gate
right = right * right_gate
out = einsum('... i k d, ... j k d -> ... i j d', left, right)
out = torch.nn.functional.layer_norm(out, (hidden_dim,), eps=1e-5, weight=weights['to_out_norm.weight'], bias=weights['to_out_norm.bias'])
out = out * out_gate
return torch.nn.functional.linear(out, weights['to_out.weight'])
'''
# 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'])
output = trimul(input_tensor, mask)
return output
'''scrolls · 137 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 44194.
⋯ 4 unchanged linesfrom torch import nn, einsumimport math+ # The flag below controls whether to allow TF32 on matmul. This flag defaults to False+ # in PyTorch 1.12 and later.+ torch.backends.cuda.matmul.allow_tf32 = True++ # The flag below controls whether to allow TF32 on cuDNN. This flag defaults to True.+ torch.backends.cudnn.allow_tf32 = True+# Reference code in PyTorchclass TriMul(nn.Module):# Based on https://github.com/lucidrains/triangle-multiplicative-module/blob/main/triangle_multiplicative_module/triangle_multiplicative_module.py⋯ 72 unchanged lines"""input_tensor, mask, weights, config = data- trimul = TriMul(dim=config["dim"], hidden_dim=config["hidden_dim"]).to(input_tensor.device)+ hidden_dim = config["hidden_dim"]+ # trimul = TriMul(dim=config["dim"], hidden_dim=config["hidden_dim"]).to(input_tensor.device)+ x = input_tensor++ batch_size, seq_len, _, dim = x.shape++ x = torch.nn.functional.layer_norm(x, (dim,), eps=1e-5, weight=weights['norm.weight'], bias=weights['norm.bias'])++ left = torch.nn.functional.linear(x, weights['left_proj.weight'])+ right = torch.nn.functional.linear(x, weights['right_proj.weight'])++ left = left * mask.unsqueeze(-1)+ right = right * mask.unsqueeze(-1)++ left_gate = torch.nn.functional.linear(x, weights['left_gate.weight']).sigmoid()+ right_gate = torch.nn.functional.linear(x, weights['right_gate.weight']).sigmoid()+ out_gate = torch.nn.functional.linear(x, weights['out_gate.weight']).sigmoid()++ left = left * left_gate+ right = right * right_gate++ out = einsum('... i k d, ... j k d -> ... i j d', left, right)++ out = torch.nn.functional.layer_norm(out, (hidden_dim,), eps=1e-5, weight=weights['to_out_norm.weight'], bias=weights['to_out_norm.bias'])+ out = out * out_gate+ return torch.nn.functional.linear(out, weights['to_out.weight'])++ '''# Fill in the given weights of the modeltrimul.norm.weight = nn.Parameter(weights['norm.weight'])trimul.norm.bias = nn.Parameter(weights['norm.bias'])⋯ 8 unchanged linesoutput = trimul(input_tensor, mask)- return outputNo newline at end of file+ return output+ '''No newline at end of file
scrolls · 60 diff lines total
Best evidence level for this revision: reported
JSON