Skip to content
KernelIndex
Search⌘K

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
NVIDIA A100
13.1ms
#32 of 69
2025-09-25

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 lines
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
⋯ 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 model
trimul.norm.weight = nn.Parameter(weights['norm.weight'])
trimul.norm.bias = nn.Parameter(weights['norm.bias'])
⋯ 8 unchanged lines
output = trimul(input_tensor, mask)
- return output
No 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