submission 39049
philip · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 162 lines, June 9 Researcher Reciprocity License v1.0.
submission_template.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-39049?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:de7d1fb56fcad995db9751be16c71dd06260fa1c3dda6ab31c97e10c51ce52a0
license declaredunknown
license concludedunknown
authorsphilip
imported2026-08-15
Kernel source
submission_template.py162 lines
#!POPCORN leaderboard trimul
# This is a submission template for popcorn leaderboard 'trimul'.
# Your task is as follows:
# > For a more complete description, see: https://tinyurl.com/gpumode-trimul
# > You will be implementing a Triangle Multiplicative Update (TriMul) module that is a core operation
# > for AlphaFold3, Chai, Protenix, and other protein structure prediction models in BioML.
# >
# > The TriMul operator operates over a 4D tensor of shape [B, N, N, C].
# >
# > Your task:
# > - Implement the "outgoing" version of the TriMul operator from the AlphaFold3 paper.
# > - You will not have to compute or store gradients for this version. You will only need to implement the forward pass.
# >
# > Input:
# > - `data`: Tuple of (input: torch.Tensor, weights: Dict[str, torch.Tensor], config: Dict)
# > - input: Input tensor of shape [bs, seq_len, seq_len, dim]
# > - mask: Mask tensor of shape [bs, seq_len, seq_len]
# > - weights: Dictionary containing model weights
# > - config: Dictionary containing model configuration parameters
# >
# > Output:
# > - Tuple containing:
# > - output: Processed tensor [bs, seq_len, seq_len, dim]
# The deadline for this leaderboard is 2025-09-30 00:00:00+00:00
# You can automatically route this file to specific GPUs by adding a line
# `#!POPCORN gpus <GPUs>` to the header of this file.
# Happy hacking!
import torch
import triton
import triton.language as tl
from torch import einsum, nn
from task import input_t, output_t
@triton.jit
def layernorm_kernel(
x_ptr,
weight_ptr,
bias_ptr,
output_ptr,
mean_ptr,
rstd_ptr,
n_rows,
n_cols,
eps,
BLOCK_SIZE: tl.constexpr,
):
row_idx = tl.program_id(0)
if row_idx >= n_rows:
return
# Load row
col_offsets = tl.arange(0, BLOCK_SIZE)
mask = col_offsets < n_cols
x_ptrs = x_ptr + row_idx * n_cols + col_offsets
x = tl.load(x_ptrs, mask=mask, other=0.0).to(tl.float32)
# Compute mean
mean = tl.sum(x, axis=0) / n_cols
tl.store(mean_ptr + row_idx, mean)
# Compute variance
x_centered = x - mean
var = tl.sum(x_centered * x_centered, axis=0) / n_cols
rstd = 1.0 / tl.sqrt(var + eps)
tl.store(rstd_ptr + row_idx, rstd)
# Normalize
x_norm = x_centered * rstd
# Apply weight and bias
weight = tl.load(weight_ptr + col_offsets, mask=mask, other=1.0)
bias = tl.load(bias_ptr + col_offsets, mask=mask, other=0.0)
output = x_norm * weight + bias
# Store output
output_ptrs = output_ptr + row_idx * n_cols + col_offsets
tl.store(output_ptrs, output, mask=mask)
def triton_layernorm(x, weight, bias, eps=1e-5):
"""
Custom Triton LayerNorm implementation
"""
# Flatten input to 2D for processing
original_shape = x.shape
x_flat = x.view(-1, original_shape[-1])
n_rows, n_cols = x_flat.shape
# Allocate output and temporary tensors
output = torch.empty_like(x_flat)
mean = torch.empty(n_rows, device=x.device, dtype=torch.float32)
rstd = torch.empty(n_rows, device=x.device, dtype=torch.float32)
# Calculate block size
BLOCK_SIZE = triton.next_power_of_2(n_cols)
# Launch kernel
grid = (n_rows,)
layernorm_kernel[grid](x_flat, weight, bias, output, mean, rstd, n_rows, n_cols, eps, BLOCK_SIZE)
return output.view(original_shape)
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
batch_size, seq_len, _, _ = input_tensor.shape
# First layer norm using Triton
# x = triton_layernorm(input_tensor, weights["norm.weight"], weights["norm.bias"])
x = torch.nn.functional.layer_norm(input_tensor, input_tensor.shape[-1:], weights["norm.weight"], weights["norm.bias"])
x = x.to(torch.float32)
# Linear projections
left = torch.nn.functional.linear(x, weights["left_proj.weight"])
right = torch.nn.functional.linear(x, weights["right_proj.weight"])
# Apply mask
mask_expanded = mask.unsqueeze(-1)
left = left * mask_expanded
right = right * mask_expanded
# Gates
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()
# Apply gates
left = left * left_gate
right = right * right_gate
# Triangle multiplication
out = einsum(
"... i k d, ... j k d -> ... i j d",
left.to(torch.bfloat16),
right.to(torch.bfloat16),
)
# Final processing
out = out.to(torch.float32)
# out = triton_layernorm(out, weights["to_out_norm.weight"], weights["to_out_norm.bias"])
# Equivalent PyTorch operation:
out = torch.nn.functional.layer_norm(out, out.shape[-1:], weights["to_out_norm.weight"], weights["to_out_norm.bias"])
out = out * out_gate
output = torch.nn.functional.linear(out, weights["to_out.weight"])
return output
scrolls · 162 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 39025.
⋯ 28 unchanged lines# Happy hacking!import torch+ import triton+ import triton.language as tlfrom torch import einsum, nnfrom task import input_t, output_t- class TriMul(nn.Module):- def __init__(- self,- dim: int,- hidden_dim: int,- ):- super().__init__()+ @triton.jit+ def layernorm_kernel(+ x_ptr,+ weight_ptr,+ bias_ptr,+ output_ptr,+ mean_ptr,+ rstd_ptr,+ n_rows,+ n_cols,+ eps,+ BLOCK_SIZE: tl.constexpr,+ ):+ row_idx = tl.program_id(0)+ if row_idx >= n_rows:+ return- self.norm = nn.LayerNorm(dim)+ # Load row+ col_offsets = tl.arange(0, BLOCK_SIZE)+ mask = col_offsets < n_cols+ x_ptrs = x_ptr + row_idx * n_cols + col_offsets+ x = tl.load(x_ptrs, mask=mask, other=0.0).to(tl.float32)- 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)+ # Compute mean+ mean = tl.sum(x, axis=0) / n_cols+ tl.store(mean_ptr + row_idx, mean)- 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)+ # Compute variance+ x_centered = x - mean+ var = tl.sum(x_centered * x_centered, axis=0) / n_cols+ rstd = 1.0 / tl.sqrt(var + eps)+ tl.store(rstd_ptr + row_idx, rstd)- self.to_out_norm = nn.LayerNorm(hidden_dim)- self.to_out = nn.Linear(hidden_dim, dim, bias=False, dtype=torch.float32)+ # Normalize+ x_norm = x_centered * rstd- def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:- """- x: [bs, seq_len, seq_len, dim]- mask: [bs, seq_len, seq_len]+ # Apply weight and bias+ weight = tl.load(weight_ptr + col_offsets, mask=mask, other=1.0)+ bias = tl.load(bias_ptr + col_offsets, mask=mask, other=0.0)+ output = x_norm * weight + bias- Returns:- output: [bs, seq_len, seq_len, dim]- """- batch_size, seq_len, _, dim = x.shape+ # Store output+ output_ptrs = output_ptr + row_idx * n_cols + col_offsets+ tl.store(output_ptrs, output, mask=mask)- x = self.norm(x)- x = x.to(torch.float32)- left = self.left_proj(x.to(torch.float32))- right = self.right_proj(x.to(torch.float32))+ def triton_layernorm(x, weight, bias, eps=1e-5):+ """+ Custom Triton LayerNorm implementation+ """+ # Flatten input to 2D for processing+ original_shape = x.shape+ x_flat = x.view(-1, original_shape[-1])+ n_rows, n_cols = x_flat.shape- mask = mask.unsqueeze(-1)- left = left * mask- right = right * mask+ # Allocate output and temporary tensors+ output = torch.empty_like(x_flat)+ mean = torch.empty(n_rows, device=x.device, dtype=torch.float32)+ rstd = torch.empty(n_rows, device=x.device, dtype=torch.float32)- left_gate = self.left_gate(x.to(torch.float32)).sigmoid()- right_gate = self.right_gate(x.to(torch.float32)).sigmoid()- out_gate = self.out_gate(x.to(torch.float32)).sigmoid()+ # Calculate block size+ BLOCK_SIZE = triton.next_power_of_2(n_cols)- left = left * left_gate- right = right * right_gate+ # Launch kernel+ grid = (n_rows,)+ layernorm_kernel[grid](x_flat, weight, bias, output, mean, rstd, n_rows, n_cols, eps, BLOCK_SIZE)- 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)+ return output.view(original_shape)- # # 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_gate- return self.to_out(out)--def custom_kernel(data: input_t) -> output_t:"""Reference implementation of TriMul using PyTorch.⋯ 6 unchanged lines- config: Dictionary containing model configuration parameters"""input_tensor, mask, weights, config = data- 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)+ batch_size, seq_len, _, _ = input_tensor.shape++ # First layer norm using Triton+ # x = triton_layernorm(input_tensor, weights["norm.weight"], weights["norm.bias"])+ x = torch.nn.functional.layer_norm(input_tensor, input_tensor.shape[-1:], weights["norm.weight"], weights["norm.bias"])+ x = x.to(torch.float32)++ # Linear projections+ left = torch.nn.functional.linear(x, weights["left_proj.weight"])+ right = torch.nn.functional.linear(x, weights["right_proj.weight"])++ # Apply mask+ mask_expanded = mask.unsqueeze(-1)+ left = left * mask_expanded+ right = right * mask_expanded++ # Gates+ 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()++ # Apply gates+ left = left * left_gate+ right = right * right_gate++ # Triangle multiplication+ out = einsum(+ "... i k d, ... j k d -> ... i j d",+ left.to(torch.bfloat16),+ right.to(torch.bfloat16),)- 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).to(torch.float32)+ # Final processing+ out = out.to(torch.float32)+ # out = triton_layernorm(out, weights["to_out_norm.weight"], weights["to_out_norm.bias"])+ # Equivalent PyTorch operation:+ out = torch.nn.functional.layer_norm(out, out.shape[-1:], weights["to_out_norm.weight"], weights["to_out_norm.bias"])+ out = out * out_gate+ output = torch.nn.functional.linear(out, weights["to_out.weight"])return output
scrolls · 209 diff lines total
Best evidence level for this revision: reported
JSON