Skip to content
KernelIndex
Search⌘K

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
NVIDIA A100
17.9ms
#44 of 69
2025-09-15

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 tl
from torch import einsum, nn
from 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