Skip to content
KernelIndex
Search⌘K

submission 45318

leymore4172 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 122 lines, June 9 Researcher Reciprocity License v1.0.

trimul.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-45318?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
18.9ms
#47 of 69
2025-09-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f4779bf2e9bcb515ca0d222e1a750289c8ade056958fb6693546cef817fe4797
license declaredunknown
license concludedunknown
authorsleymore4172
imported2026-08-15

Kernel source

trimul.py122 lines
import torch
from torch import nn, einsum
from task import input_t, output_t

# Global cache for JIT compiled models based on dimensions
_jit_model_cache = {}

class TriMul(nn.Module):
    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, dtype=torch.float32)
        self.right_proj = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)

        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)

        self.to_out_norm = nn.LayerNorm(hidden_dim)
        self.to_out = nn.Linear(hidden_dim, dim, bias=False, dtype=torch.float32)

    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]
        """
        # Optimize: Remove redundant shape unpacking
        x = self.norm(x)
        x = x.to(torch.float32)

        # Optimize: Single type conversion for x
        x_float32 = x  # Already converted above

        left = self.left_proj(x_float32)
        right = self.right_proj(x_float32)

        mask = mask.unsqueeze(-1)
        left = left * mask
        right = right * mask

        left_gate = self.left_gate(x_float32).sigmoid()
        right_gate = self.right_gate(x_float32).sigmoid()
        out_gate = self.out_gate(x_float32).sigmoid()

        left = left * left_gate
        right = right * right_gate

        out = einsum('... i k d, ... j k d -> ... i j d', left.to(torch.bfloat16), right.to(torch.bfloat16))

        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.

    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

    # Create cache key based on dimensions and device
    cache_key = (config["dim"], config["hidden_dim"], str(input_tensor.device))

    # Check if we have a JIT compiled model for these dimensions
    if cache_key not in _jit_model_cache:
        # Create and JIT compile the model
        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))
        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))

        # JIT compile the model using script
        trimul.eval()  # Set to eval mode for consistency
        with torch.no_grad():
            _jit_model_cache[cache_key] = torch.jit.script(trimul)
    else:
        # Use cached JIT compiled model and update weights
        jit_model = _jit_model_cache[cache_key]

        # Update weights in the JIT model
        jit_model.norm.weight.data = weights['norm.weight'].to(torch.float32)
        jit_model.left_proj.weight.data = weights['left_proj.weight'].to(torch.float32)
        jit_model.right_proj.weight.data = weights['right_proj.weight'].to(torch.float32)
        jit_model.left_gate.weight.data = weights['left_gate.weight'].to(torch.float32)
        jit_model.right_gate.weight.data = weights['right_gate.weight'].to(torch.float32)
        jit_model.out_gate.weight.data = weights['out_gate.weight'].to(torch.float32)
        jit_model.to_out_norm.weight.data = weights['to_out_norm.weight'].to(torch.float32)
        jit_model.to_out.weight.data = weights['to_out.weight'].to(torch.float32)
        jit_model.norm.bias.data = weights['norm.bias'].to(torch.float32)
        jit_model.to_out_norm.bias.data = weights['to_out_norm.bias'].to(torch.float32)

    # Use the JIT compiled model
    output = _jit_model_cache[cache_key](input_tensor, mask).to(torch.float32)

    return output
scrolls · 122 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 44748.

- from utils import make_match_reference, DisableCuDNNTF32
- from task import input_t, output_t
-
import torch
from torch import nn, einsum
- import math
+ from task import input_t, output_t
- # Reference code in PyTorch
+ # Global cache for JIT compiled models based on dimensions
+ _jit_model_cache = {}
+
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,
⋯ 3 unchanged lines
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_proj = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)
+ self.right_proj = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)
- 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.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)
self.to_out_norm = nn.LayerNorm(hidden_dim)
- self.to_out = nn.Linear(hidden_dim, dim, bias=False)
+ self.to_out = nn.Linear(hidden_dim, dim, bias=False, dtype=torch.float32)
def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
"""
⋯ 3 unchanged lines
Returns:
output: [bs, seq_len, seq_len, dim]
"""
- batch_size, seq_len, _, dim = x.shape
-
+ # Optimize: Remove redundant shape unpacking
x = self.norm(x)
+ x = x.to(torch.float32)
- left = self.left_proj(x)
- right = self.right_proj(x)
+ # Optimize: Single type conversion for x
+ x_float32 = x # Already converted above
+ left = self.left_proj(x_float32)
+ right = self.right_proj(x_float32)
+
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_gate = self.left_gate(x_float32).sigmoid()
+ right_gate = self.right_gate(x_float32).sigmoid()
+ out_gate = self.out_gate(x_float32).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)
+ out = einsum('... i k d, ... j k d -> ... i j d', left.to(torch.bfloat16), right.to(torch.bfloat16))
- # # 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)
⋯ 10 unchanged lines
- weights: Dictionary containing model weights
- config: Dictionary containing model configuration parameters
"""
+ input_tensor, mask, weights, config = data
- # Use deterministic kernels and disable TF32 for accuracy
- with DisableCuDNNTF32():
- input_tensor, mask, weights, config = data
- trimul = TriMul(dim=config["dim"], hidden_dim=config["hidden_dim"]).to(input_tensor.device)
+ # Create cache key based on dimensions and device
+ cache_key = (config["dim"], config["hidden_dim"], str(input_tensor.device))
+ # Check if we have a JIT compiled model for these dimensions
+ if cache_key not in _jit_model_cache:
+ # Create and JIT compile the model
+ 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'])
- 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'])
+ trimul.norm.weight = nn.Parameter(weights['norm.weight'].to(torch.float32))
+ trimul.left_proj.weight = nn.Parameter(weights['left_proj.weight'].to(torch.float32))
+ 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)
-
- return output
-
- # Input generation for the reference code
- def generate_input(
- seqlen: int,
- bs: int,
- dim: int,
- hiddendim: int,
- seed: int,
- nomask: bool,
- distribution: str,
- ) -> input_t:
-
- # Really dumb but for now _ isn't parsing correctly.
- batch_size = bs
- seq_len = seqlen
- hidden_dim = hiddendim
- no_mask = nomask
-
- config = {
- "hidden_dim": hidden_dim,
- "dim": dim,
- }
-
- gen = torch.Generator(device='cuda')
- gen.manual_seed(seed)
-
- weights = {}
-
- # Generate input tensor based on distribution
- if distribution == "cauchy":
- # Heavier tail distribution
- input_tensor = torch.distributions.Cauchy(0, 2).sample(
- (batch_size, seq_len, seq_len, dim)
- ).to(device='cuda', dtype=torch.float32)
- else: # normal distribution
- input_tensor = torch.randn(
- (batch_size, seq_len, seq_len, dim),
- device='cuda',
- dtype=torch.float32,
- generator=gen
- ).contiguous()
-
- if no_mask:
- mask = torch.ones(batch_size, seq_len, seq_len, device=input_tensor.device)
+ # JIT compile the model using script
+ trimul.eval() # Set to eval mode for consistency
+ with torch.no_grad():
+ _jit_model_cache[cache_key] = torch.jit.script(trimul)
else:
- mask = torch.randint(0, 2, (batch_size, seq_len, seq_len), device=input_tensor.device, generator=gen)
+ # Use cached JIT compiled model and update weights
+ jit_model = _jit_model_cache[cache_key]
- # Initialize model weights based on distribution
- weights["norm.weight"] = torch.randn(dim, device="cuda", dtype=torch.float32)
- weights["norm.bias"] = torch.randn(dim, device="cuda", dtype=torch.float32)
- weights["left_proj.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
- weights["right_proj.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
- weights["left_gate.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
- weights["right_gate.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
- weights["out_gate.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
- weights["to_out_norm.weight"] = torch.randn(hidden_dim, device="cuda", dtype=torch.float32)
- weights["to_out.weight"] = torch.randn(dim, hidden_dim, device="cuda", dtype=torch.float32) / math.sqrt(dim)
- weights["to_out_norm.bias"] = torch.randn(hidden_dim, device="cuda", dtype=torch.float32)
+ # Update weights in the JIT model
+ jit_model.norm.weight.data = weights['norm.weight'].to(torch.float32)
+ jit_model.left_proj.weight.data = weights['left_proj.weight'].to(torch.float32)
+ jit_model.right_proj.weight.data = weights['right_proj.weight'].to(torch.float32)
+ jit_model.left_gate.weight.data = weights['left_gate.weight'].to(torch.float32)
+ jit_model.right_gate.weight.data = weights['right_gate.weight'].to(torch.float32)
+ jit_model.out_gate.weight.data = weights['out_gate.weight'].to(torch.float32)
+ jit_model.to_out_norm.weight.data = weights['to_out_norm.weight'].to(torch.float32)
+ jit_model.to_out.weight.data = weights['to_out.weight'].to(torch.float32)
+ jit_model.norm.bias.data = weights['norm.bias'].to(torch.float32)
+ jit_model.to_out_norm.bias.data = weights['to_out_norm.bias'].to(torch.float32)
- return (input_tensor, mask, weights, config)
+ # Use the JIT compiled model
+ output = _jit_model_cache[cache_key](input_tensor, mask).to(torch.float32)
-
- check_implementation = make_match_reference(custom_kernel, rtol=2e-2, atol=2e-2)
+ return output
scrolls · 214 diff lines total

Best evidence level for this revision: reported

JSON