Skip to content
KernelIndex
Search⌘K

submission 45338

leymore4172 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

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

Kernel source

trimul.py107 lines
import torch
import torch.nn.functional as F
from torch import einsum
from task import input_t, output_t

@torch.jit.script
def trimul(x: torch.Tensor, mask: torch.Tensor,
          norm_weight: torch.Tensor, norm_bias: torch.Tensor,
          left_proj_weight: torch.Tensor, right_proj_weight: torch.Tensor,
          left_gate_weight: torch.Tensor, right_gate_weight: torch.Tensor,
          out_gate_weight: torch.Tensor,
          to_out_norm_weight: torch.Tensor, to_out_norm_bias: torch.Tensor,
          to_out_weight: torch.Tensor) -> torch.Tensor:
    """
    Functional implementation of TriMul.

    Args:
        x: [bs, seq_len, seq_len, dim]
        mask: [bs, seq_len, seq_len]
        norm_weight: Layer norm weight
        norm_bias: Layer norm bias
        left_proj_weight: Left projection weight
        right_proj_weight: Right projection weight
        left_gate_weight: Left gate weight
        right_gate_weight: Right gate weight
        out_gate_weight: Output gate weight
        to_out_norm_weight: Output layer norm weight
        to_out_norm_bias: Output layer norm bias
        to_out_weight: Final output projection weight

    Returns:
        output: [bs, seq_len, seq_len, dim]
    """
    batch_size, seq_len, _, dim = x.shape

    # Layer normalization
    x = F.layer_norm(x, [dim], norm_weight, norm_bias)

    # Linear projections without bias
    left = F.linear(x, left_proj_weight)
    right = F.linear(x, right_proj_weight)

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

    # Gate computations
    left_gate = F.linear(x, left_gate_weight).sigmoid()
    right_gate = F.linear(x, right_gate_weight).sigmoid()
    out_gate = F.linear(x, out_gate_weight).sigmoid()

    # Apply gates
    left = left * left_gate
    right = right * right_gate

    # Einstein summation
    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)

    # # 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)

    # Output normalization
    hidden_dim = out.shape[-1]
    out = F.layer_norm(out, [hidden_dim], to_out_norm_weight, to_out_norm_bias)
    out = out * out_gate

    # Final linear projection
    out = F.linear(out, to_out_weight)

    return out


def custom_kernel(data: input_t) -> output_t:
    """
    Reference implementation of TriMul using PyTorch functional interface.

    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

    output = trimul(
        input_tensor, mask,
        weights['norm.weight'], weights['norm.bias'],
        weights['left_proj.weight'], weights['right_proj.weight'],
        weights['left_gate.weight'], weights['right_gate.weight'],
        weights['out_gate.weight'],
        weights['to_out_norm.weight'], weights['to_out_norm.bias'],
        weights['to_out.weight']
    )

    return output
scrolls · 107 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 45318.

import torch
- from torch import nn, einsum
+ import torch.nn.functional as F
+ from torch import einsum
from task import input_t, output_t
- # Global cache for JIT compiled models based on dimensions
- _jit_model_cache = {}
+ @torch.jit.script
+ def trimul(x: torch.Tensor, mask: torch.Tensor,
+ norm_weight: torch.Tensor, norm_bias: torch.Tensor,
+ left_proj_weight: torch.Tensor, right_proj_weight: torch.Tensor,
+ left_gate_weight: torch.Tensor, right_gate_weight: torch.Tensor,
+ out_gate_weight: torch.Tensor,
+ to_out_norm_weight: torch.Tensor, to_out_norm_bias: torch.Tensor,
+ to_out_weight: torch.Tensor) -> torch.Tensor:
+ """
+ Functional implementation of TriMul.
- class TriMul(nn.Module):
- def __init__(
- self,
- dim: int,
- hidden_dim: int,
- ):
- super().__init__()
+ Args:
+ x: [bs, seq_len, seq_len, dim]
+ mask: [bs, seq_len, seq_len]
+ norm_weight: Layer norm weight
+ norm_bias: Layer norm bias
+ left_proj_weight: Left projection weight
+ right_proj_weight: Right projection weight
+ left_gate_weight: Left gate weight
+ right_gate_weight: Right gate weight
+ out_gate_weight: Output gate weight
+ to_out_norm_weight: Output layer norm weight
+ to_out_norm_bias: Output layer norm bias
+ to_out_weight: Final output projection weight
- self.norm = nn.LayerNorm(dim)
+ Returns:
+ output: [bs, seq_len, seq_len, dim]
+ """
+ batch_size, seq_len, _, dim = x.shape
- 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)
+ # Layer normalization
+ x = F.layer_norm(x, [dim], norm_weight, norm_bias)
- 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)
+ # Linear projections without bias
+ left = F.linear(x, left_proj_weight)
+ right = F.linear(x, right_proj_weight)
- self.to_out_norm = nn.LayerNorm(hidden_dim)
- self.to_out = nn.Linear(hidden_dim, dim, bias=False, dtype=torch.float32)
+ # Apply mask
+ mask = mask.unsqueeze(-1)
+ left = left * mask
+ right = right * mask
- def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
- """
- x: [bs, seq_len, seq_len, dim]
- mask: [bs, seq_len, seq_len]
+ # Gate computations
+ left_gate = F.linear(x, left_gate_weight).sigmoid()
+ right_gate = F.linear(x, right_gate_weight).sigmoid()
+ out_gate = F.linear(x, out_gate_weight).sigmoid()
- Returns:
- output: [bs, seq_len, seq_len, dim]
- """
- # Optimize: Remove redundant shape unpacking
- x = self.norm(x)
- x = x.to(torch.float32)
+ # Apply gates
+ left = left * left_gate
+ right = right * right_gate
- # Optimize: Single type conversion for x
- x_float32 = x # Already converted above
+ # Einstein summation
+ 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)
- left = self.left_proj(x_float32)
- right = self.right_proj(x_float32)
+ # # 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, :]
- mask = mask.unsqueeze(-1)
- left = left * mask
- right = right * mask
+ out = out.to(torch.float32)
- left_gate = self.left_gate(x_float32).sigmoid()
- right_gate = self.right_gate(x_float32).sigmoid()
- out_gate = self.out_gate(x_float32).sigmoid()
+ # Output normalization
+ hidden_dim = out.shape[-1]
+ out = F.layer_norm(out, [hidden_dim], to_out_norm_weight, to_out_norm_bias)
+ out = out * out_gate
- left = left * left_gate
- right = right * right_gate
+ # Final linear projection
+ out = F.linear(out, to_out_weight)
- out = einsum('... i k d, ... j k d -> ... i j d', left.to(torch.bfloat16), right.to(torch.bfloat16))
+ return out
- 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.
+ Reference implementation of TriMul using PyTorch functional interface.
Args:
data: Tuple of (input: torch.Tensor, mask: torch.Tensor, weights: Dict[str, torch.Tensor], config: Dict)
⋯ 4 unchanged lines
"""
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))
+ output = trimul(
+ input_tensor, mask,
+ weights['norm.weight'], weights['norm.bias'],
+ weights['left_proj.weight'], weights['right_proj.weight'],
+ weights['left_gate.weight'], weights['right_gate.weight'],
+ weights['out_gate.weight'],
+ weights['to_out_norm.weight'], weights['to_out_norm.bias'],
+ weights['to_out.weight']
+ )
- # 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 · 192 diff lines total

Best evidence level for this revision: reported

JSON