Skip to content
KernelIndex
Search⌘K

submission 44222

davidberard · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-44222?include=source"
interfacepython
Compatibility
measured onNVIDIA H100
declared hardwareNVIDIA H100
architecturessm_90
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA H100
6.59ms
#52 of 71
2025-09-25

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:06c27237cbfe1c5e6edb1c83fdfd8612837b32cce87231bc2e3700933b93d303
license declaredunknown
license concludedunknown
authorsdavidberard
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

autotuneconfigs.append(triton.Config({
mmaaccumulator1 = tl.dot(a, b1.T, accumulator1, allow_tf32=True)
num-warps = 8}, num_stages=4, num_warps=8))
persistent-kernelPersistent dual matrix multiplication: A @ B1.T and A @ B2.T using on-device TMA descriptors.
stages = 4}, num_stages=4, num_warps=8))

Kernel source

v2.py323 lines
# from utils import make_match_reference, DisableCuDNNTF32
from task import input_t, output_t

import torch
from torch import nn, einsum
import math
import os

import triton
import triton.language as tl

# 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

# Set allocator for TMA descriptors (required for on-device TMA)
def alloc_fn(size: int, alignment: int, stream=None):
    return torch.empty(size, device="cuda", dtype=torch.int8)

triton.set_allocator(alloc_fn)

os.environ['TRITON_PRINT_AUTOTUNING'] = '1'
os.environ['MLIR_ENABLE_DIAGNOSTICS'] = 'warnings,remarks'

# 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 two_mm_kernel_configs():
    configs = []
    for BLOCK_M in [64, 128]:
        for BLOCK_N in [64, 128, 256]:
            for BLOCK_K in [32, 64, 128]:
                configs.append(triton.Config({
                    'BLOCK_M': BLOCK_M,
                    'BLOCK_N': BLOCK_N,
                    'BLOCK_K': BLOCK_K,
                    'GROUP_SIZE_M': 8
                }, num_stages=4, num_warps=8))
    return configs

@triton.autotune(
    two_mm_kernel_configs(), key=["M", "N", "K"]
)
@triton.jit
def two_mm_kernel(a_ptr, b1_ptr, b2_ptr, c1_ptr, c2_ptr, mask_ptr, M, N, K, stride_am, stride_ak, stride_b1k, stride_b1n, stride_b2k, stride_b2n, stride_c1m, stride_c1n, stride_c2m, stride_c2n, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, NUM_SMS: tl.constexpr):
    # Persistent kernel using on-device TMA descriptors
    start_pid = tl.program_id(axis=0)
    num_pid_m = tl.cdiv(M, BLOCK_M)
    num_pid_n = tl.cdiv(N, BLOCK_N)
    k_tiles = tl.cdiv(K, BLOCK_K)
    num_tiles = num_pid_m * num_pid_n

    # Create on-device TMA descriptors
    a_desc = tl._experimental_make_tensor_descriptor(
        a_ptr,
        shape=[M, K],
        strides=[stride_am, stride_ak],
        block_shape=[BLOCK_M, BLOCK_K],
    )
    b1_desc = tl._experimental_make_tensor_descriptor(
        b1_ptr,
        shape=[N, K],
        strides=[stride_b1n, stride_b1k],
        block_shape=[BLOCK_N, BLOCK_K],
    )
    b2_desc = tl._experimental_make_tensor_descriptor(
        b2_ptr,
        shape=[N, K],
        strides=[stride_b2n, stride_b2k],
        block_shape=[BLOCK_N, BLOCK_K],
    )

    # tile_id_c is used in the epilogue to break the dependency between
    # the prologue and the epilogue
    tile_id_c = start_pid - NUM_SMS
    num_pid_in_group = GROUP_SIZE_M * num_pid_n

    # Persistent loop over tiles
    for tile_id in tl.range(start_pid, num_tiles, NUM_SMS, flatten=False):
        # Calculate PID for this tile using improved swizzling
        group_id = tile_id // num_pid_in_group
        first_pid_m = group_id * GROUP_SIZE_M
        group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
        pid_m = first_pid_m + (tile_id % group_size_m)
        pid_n = (tile_id % num_pid_in_group) // group_size_m

        # Calculate block offsets
        offs_am = pid_m * BLOCK_M
        offs_bn = pid_n * BLOCK_N

        # Initialize accumulators for both outputs
        accumulator1 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
        accumulator2 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

        # Main computation loop over K dimension
        for ki in range(k_tiles):
            offs_k = ki * BLOCK_K
            # Load blocks from A, B1, B2 using on-device TMA
            a = a_desc.load([offs_am, offs_k])
            b1 = b1_desc.load([offs_bn, offs_k])
            b2 = b2_desc.load([offs_bn, offs_k])

            # Perform matrix multiplications: A @ B1.T and A @ B2.T using TF32
            accumulator1 = tl.dot(a, b1.T, accumulator1, allow_tf32=True)
            accumulator2 = tl.dot(a, b2.T, accumulator2, allow_tf32=True)

        # Store results using separate tile_id_c for epilogue
        tile_id_c += NUM_SMS
        group_id = tile_id_c // num_pid_in_group
        first_pid_m = group_id * GROUP_SIZE_M
        group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
        pid_m = first_pid_m + (tile_id_c % group_size_m)
        pid_n = (tile_id_c % num_pid_in_group) // group_size_m

        # Calculate output offsets and pointers
        offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
        offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)

        # Create masks for bounds checking
        c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)

        # Calculate pointer addresses
        c1_ptrs = c1_ptr + stride_c1m * offs_cm[:, None] + stride_c1n * offs_cn[None, :]
        c2_ptrs = c2_ptr + stride_c2m * offs_cm[:, None] + stride_c2n * offs_cn[None, :]

        mask = tl.load(mask_ptr + offs_cm, mask=(offs_cm < M))

        # Broadcast mask to match accumulator dimensions [BLOCK_M, BLOCK_N]
        mask_2d = mask[:, None]  # Convert to [BLOCK_M, 1] then broadcast
        accumulator1 = tl.where(mask_2d, accumulator1, 0)
        accumulator2 = tl.where(mask_2d, accumulator2, 0)

        # Convert to appropriate output dtype and store with normal tl.store
        c1 = accumulator1.to(c1_ptr.dtype.element_ty)
        c2 = accumulator2.to(c2_ptr.dtype.element_ty)

        tl.store(c1_ptrs, c1, mask=c_mask)
        tl.store(c2_ptrs, c2, mask=c_mask)

def two_mm(A, B1, B2, mask):
    """
    Persistent dual matrix multiplication: A @ B1.T and A @ B2.T using on-device TMA descriptors.

    Args:
        A: [..., K] tensor (arbitrary leading dimensions)
        B1: [N, K] matrix (will be transposed)
        B2: [N, K] matrix (will be transposed)

    Returns:
        (C1, C2): Tuple of result tensors [..., N] with same leading dims as A
    """
    # Check constraints
    assert A.shape[-1] == B1.shape[1] == B2.shape[1], "Incompatible K dimensions"
    assert A.dtype == B1.dtype == B2.dtype, "Incompatible dtypes"

    # Get dimensions
    original_shape = A.shape[:-1]  # All dimensions except the last
    K = A.shape[-1]
    N = B1.shape[0]
    dtype = A.dtype

    # Flatten A to 2D for kernel processing
    A_2d = A.view(-1, K)  # [M, K] where M is product of all leading dims
    M = A_2d.shape[0]

    # Allocate outputs as 2D then reshape
    C1_2d = torch.empty((M, N), device=A.device, dtype=dtype)
    C2_2d = torch.empty((M, N), device=A.device, dtype=dtype)

    # Get number of streaming multiprocessors
    NUM_SMS = torch.cuda.get_device_properties("cuda").multi_processor_count


    # Launch persistent kernel with limited number of blocks
    grid = lambda META: (min(NUM_SMS, triton.cdiv(M, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"])),)

    two_mm_kernel[grid](
        A_2d, B1, B2, C1_2d, C2_2d, mask,
        M, N, K,
        A_2d.stride(0), A_2d.stride(1),
        B1.stride(1), B1.stride(0),  # Note: B1 is [N, K] but we access as transposed
        B2.stride(1), B2.stride(0),  # Note: B2 is [N, K] but we access as transposed
        C1_2d.stride(0), C1_2d.stride(1),
        C2_2d.stride(0), C2_2d.stride(1),
        NUM_SMS=NUM_SMS
    )

    # Reshape outputs back to original shape + N dimension
    output_shape = original_shape + (N,)
    C1 = C1_2d.view(output_shape)
    C2 = C2_2d.view(output_shape)

    return C1, C2

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, right = two_mm(x, weights["left_proj.weight"], weights["right_proj.weight"], mask)
    # left = torch.nn.functional.linear(x, weights['left_proj.weight'].to(torch.float16))
    # right = torch.nn.functional.linear(x, weights['right_proj.weight'].to(torch.float16))

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

    '''
    left = left.to(torch.float32)
    right = right.to(torch.float32)
    x = x.to(torch.float32)
    '''

    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 · 323 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 44207.

⋯ 3 unchanged lines
import torch
from torch import nn, einsum
import math
+ import os
+ import triton
+ import triton.language as tl
+
# 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
⋯ 1 unchanged lines
# The flag below controls whether to allow TF32 on cuDNN. This flag defaults to True.
torch.backends.cudnn.allow_tf32 = True
+ # Set allocator for TMA descriptors (required for on-device TMA)
+ def alloc_fn(size: int, alignment: int, stream=None):
+ return torch.empty(size, device="cuda", dtype=torch.int8)
+
+ triton.set_allocator(alloc_fn)
+
+ os.environ['TRITON_PRINT_AUTOTUNING'] = '1'
+ os.environ['MLIR_ENABLE_DIAGNOSTICS'] = 'warnings,remarks'
+
# 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
⋯ 58 unchanged lines
out = out * out_gate
return self.to_out(out)
+ def two_mm_kernel_configs():
+ configs = []
+ for BLOCK_M in [64, 128]:
+ for BLOCK_N in [64, 128, 256]:
+ for BLOCK_K in [32, 64, 128]:
+ configs.append(triton.Config({
+ 'BLOCK_M': BLOCK_M,
+ 'BLOCK_N': BLOCK_N,
+ 'BLOCK_K': BLOCK_K,
+ 'GROUP_SIZE_M': 8
+ }, num_stages=4, num_warps=8))
+ return configs
+ @triton.autotune(
+ two_mm_kernel_configs(), key=["M", "N", "K"]
+ )
+ @triton.jit
+ def two_mm_kernel(a_ptr, b1_ptr, b2_ptr, c1_ptr, c2_ptr, mask_ptr, M, N, K, stride_am, stride_ak, stride_b1k, stride_b1n, stride_b2k, stride_b2n, stride_c1m, stride_c1n, stride_c2m, stride_c2n, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, NUM_SMS: tl.constexpr):
+ # Persistent kernel using on-device TMA descriptors
+ start_pid = tl.program_id(axis=0)
+ num_pid_m = tl.cdiv(M, BLOCK_M)
+ num_pid_n = tl.cdiv(N, BLOCK_N)
+ k_tiles = tl.cdiv(K, BLOCK_K)
+ num_tiles = num_pid_m * num_pid_n
+
+ # Create on-device TMA descriptors
+ a_desc = tl._experimental_make_tensor_descriptor(
+ a_ptr,
+ shape=[M, K],
+ strides=[stride_am, stride_ak],
+ block_shape=[BLOCK_M, BLOCK_K],
+ )
+ b1_desc = tl._experimental_make_tensor_descriptor(
+ b1_ptr,
+ shape=[N, K],
+ strides=[stride_b1n, stride_b1k],
+ block_shape=[BLOCK_N, BLOCK_K],
+ )
+ b2_desc = tl._experimental_make_tensor_descriptor(
+ b2_ptr,
+ shape=[N, K],
+ strides=[stride_b2n, stride_b2k],
+ block_shape=[BLOCK_N, BLOCK_K],
+ )
+
+ # tile_id_c is used in the epilogue to break the dependency between
+ # the prologue and the epilogue
+ tile_id_c = start_pid - NUM_SMS
+ num_pid_in_group = GROUP_SIZE_M * num_pid_n
+
+ # Persistent loop over tiles
+ for tile_id in tl.range(start_pid, num_tiles, NUM_SMS, flatten=False):
+ # Calculate PID for this tile using improved swizzling
+ group_id = tile_id // num_pid_in_group
+ first_pid_m = group_id * GROUP_SIZE_M
+ group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
+ pid_m = first_pid_m + (tile_id % group_size_m)
+ pid_n = (tile_id % num_pid_in_group) // group_size_m
+
+ # Calculate block offsets
+ offs_am = pid_m * BLOCK_M
+ offs_bn = pid_n * BLOCK_N
+
+ # Initialize accumulators for both outputs
+ accumulator1 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
+ accumulator2 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
+
+ # Main computation loop over K dimension
+ for ki in range(k_tiles):
+ offs_k = ki * BLOCK_K
+ # Load blocks from A, B1, B2 using on-device TMA
+ a = a_desc.load([offs_am, offs_k])
+ b1 = b1_desc.load([offs_bn, offs_k])
+ b2 = b2_desc.load([offs_bn, offs_k])
+
+ # Perform matrix multiplications: A @ B1.T and A @ B2.T using TF32
+ accumulator1 = tl.dot(a, b1.T, accumulator1, allow_tf32=True)
+ accumulator2 = tl.dot(a, b2.T, accumulator2, allow_tf32=True)
+
+ # Store results using separate tile_id_c for epilogue
+ tile_id_c += NUM_SMS
+ group_id = tile_id_c // num_pid_in_group
+ first_pid_m = group_id * GROUP_SIZE_M
+ group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
+ pid_m = first_pid_m + (tile_id_c % group_size_m)
+ pid_n = (tile_id_c % num_pid_in_group) // group_size_m
+
+ # Calculate output offsets and pointers
+ offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
+ offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
+
+ # Create masks for bounds checking
+ c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
+
+ # Calculate pointer addresses
+ c1_ptrs = c1_ptr + stride_c1m * offs_cm[:, None] + stride_c1n * offs_cn[None, :]
+ c2_ptrs = c2_ptr + stride_c2m * offs_cm[:, None] + stride_c2n * offs_cn[None, :]
+
+ mask = tl.load(mask_ptr + offs_cm, mask=(offs_cm < M))
+
+ # Broadcast mask to match accumulator dimensions [BLOCK_M, BLOCK_N]
+ mask_2d = mask[:, None] # Convert to [BLOCK_M, 1] then broadcast
+ accumulator1 = tl.where(mask_2d, accumulator1, 0)
+ accumulator2 = tl.where(mask_2d, accumulator2, 0)
+
+ # Convert to appropriate output dtype and store with normal tl.store
+ c1 = accumulator1.to(c1_ptr.dtype.element_ty)
+ c2 = accumulator2.to(c2_ptr.dtype.element_ty)
+
+ tl.store(c1_ptrs, c1, mask=c_mask)
+ tl.store(c2_ptrs, c2, mask=c_mask)
+
+ def two_mm(A, B1, B2, mask):
+ """
+ Persistent dual matrix multiplication: A @ B1.T and A @ B2.T using on-device TMA descriptors.
+
+ Args:
+ A: [..., K] tensor (arbitrary leading dimensions)
+ B1: [N, K] matrix (will be transposed)
+ B2: [N, K] matrix (will be transposed)
+
+ Returns:
+ (C1, C2): Tuple of result tensors [..., N] with same leading dims as A
+ """
+ # Check constraints
+ assert A.shape[-1] == B1.shape[1] == B2.shape[1], "Incompatible K dimensions"
+ assert A.dtype == B1.dtype == B2.dtype, "Incompatible dtypes"
+
+ # Get dimensions
+ original_shape = A.shape[:-1] # All dimensions except the last
+ K = A.shape[-1]
+ N = B1.shape[0]
+ dtype = A.dtype
+
+ # Flatten A to 2D for kernel processing
+ A_2d = A.view(-1, K) # [M, K] where M is product of all leading dims
+ M = A_2d.shape[0]
+
+ # Allocate outputs as 2D then reshape
+ C1_2d = torch.empty((M, N), device=A.device, dtype=dtype)
+ C2_2d = torch.empty((M, N), device=A.device, dtype=dtype)
+
+ # Get number of streaming multiprocessors
+ NUM_SMS = torch.cuda.get_device_properties("cuda").multi_processor_count
+
+
+ # Launch persistent kernel with limited number of blocks
+ grid = lambda META: (min(NUM_SMS, triton.cdiv(M, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"])),)
+
+ two_mm_kernel[grid](
+ A_2d, B1, B2, C1_2d, C2_2d, mask,
+ M, N, K,
+ A_2d.stride(0), A_2d.stride(1),
+ B1.stride(1), B1.stride(0), # Note: B1 is [N, K] but we access as transposed
+ B2.stride(1), B2.stride(0), # Note: B2 is [N, K] but we access as transposed
+ C1_2d.stride(0), C1_2d.stride(1),
+ C2_2d.stride(0), C2_2d.stride(1),
+ NUM_SMS=NUM_SMS
+ )
+
+ # Reshape outputs back to original shape + N dimension
+ output_shape = original_shape + (N,)
+ C1 = C1_2d.view(output_shape)
+ C2 = C2_2d.view(output_shape)
+
+ return C1, C2
+
def custom_kernel(data: input_t) -> output_t:
"""
Reference implementation of TriMul using PyTorch.
⋯ 16 unchanged lines
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, right = two_mm(x, weights["left_proj.weight"], weights["right_proj.weight"], mask)
+ # left = torch.nn.functional.linear(x, weights['left_proj.weight'].to(torch.float16))
+ # right = torch.nn.functional.linear(x, weights['right_proj.weight'].to(torch.float16))
- left = left * mask.unsqueeze(-1)
- right = right * mask.unsqueeze(-1)
+ # left = left * mask.unsqueeze(-1)
+ # right = right * mask.unsqueeze(-1)
+ '''
+ left = left.to(torch.float32)
+ right = right.to(torch.float32)
+ x = x.to(torch.float32)
+ '''
+
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()
scrolls · 226 diff lines total

Best evidence level for this revision: reported

JSON