Skip to content
KernelIndex
Search⌘K

submission 426455

Cookie 🍪 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

trimul_opus_real.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-426455?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
69.3ms
#70 of 71
2026-02-02

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a0710768fa9aaeea2b29faf25b8df6e013f299a2b1ae9f57881ce3674a244a61
license declaredunknown
license concludedunknown
authorsCookie 🍪
imported2026-08-15

Techniques

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

mmaacc += tl.dot(x0, tl.trans(x1))
tile-k = 16BLOCK_I, BLOCK_J, BLOCK_K = 16, 16, 16

Kernel source

trimul_opus_real.py272 lines
import torch
import triton
import triton.language as tl


@triton.jit
def _layernorm_kernel(
    X_ptr, Y_ptr, W_ptr, B_ptr,
    stride_x, C, eps,
    BLOCK_SIZE: tl.constexpr,
):
    row_idx = tl.program_id(0)
    row_start_ptr = X_ptr + row_idx * stride_x
    col_offsets = tl.arange(0, BLOCK_SIZE)
    mask = col_offsets < C

    x = tl.load(row_start_ptr + col_offsets, mask=mask, other=0.0).to(tl.float32)
    x_sum = tl.sum(x, axis=0)
    mean = x_sum / C
    x_centered = tl.where(mask, x - mean, 0.0)
    var_sum = tl.sum(x_centered * x_centered, axis=0)
    var = var_sum / C
    rstd = 1.0 / tl.sqrt(var + eps)
    x_norm = x_centered * rstd

    w = tl.load(W_ptr + col_offsets, mask=mask, other=1.0).to(tl.float32)
    b = tl.load(B_ptr + col_offsets, mask=mask, other=0.0).to(tl.float32)
    y = x_norm * w + b

    out_ptr = Y_ptr + row_idx * stride_x
    tl.store(out_ptr + col_offsets, y, mask=mask)


@triton.jit
def _projection_gating_kernel(
    x_ptr, mask_ptr,
    left_proj_weight_ptr, right_proj_weight_ptr,
    left_gate_weight_ptr, right_gate_weight_ptr, out_gate_weight_ptr,
    left_out_ptr, right_out_ptr, out_gate_ptr,
    B, N, C, H,
    BLOCK_H: tl.constexpr, BLOCK_C: tl.constexpr,
):
    pid = tl.program_id(0)
    b = pid // (N * N)
    remainder = pid % (N * N)
    i = remainder // N
    j = remainder % N

    x_base = b * (N * N * C) + i * (N * C) + j * C
    mask_idx = b * (N * N) + i * N + j
    mask_val = tl.load(mask_ptr + mask_idx).to(tl.float32)
    out_base = b * (N * N * H) + i * (N * H) + j * H

    for h_start in range(0, H, BLOCK_H):
        h_offs = h_start + tl.arange(0, BLOCK_H)
        h_mask = h_offs < H

        left_proj_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)
        right_proj_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)
        left_gate_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)
        right_gate_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)
        out_gate_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)

        for c_start in range(0, C, BLOCK_C):
            c_offs = c_start + tl.arange(0, BLOCK_C)
            c_mask = c_offs < C
            x_vals = tl.load(x_ptr + x_base + c_offs, mask=c_mask, other=0.0).to(tl.float32)

            weight_offsets = h_offs[:, None] * C + c_offs[None, :]
            combined_mask = h_mask[:, None] & c_mask[None, :]

            left_proj_w = tl.load(left_proj_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)
            right_proj_w = tl.load(right_proj_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)
            left_gate_w = tl.load(left_gate_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)
            right_gate_w = tl.load(right_gate_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)
            out_gate_w = tl.load(out_gate_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)

            left_proj_acc += tl.sum(left_proj_w * x_vals[None, :], axis=1)
            right_proj_acc += tl.sum(right_proj_w * x_vals[None, :], axis=1)
            left_gate_acc += tl.sum(left_gate_w * x_vals[None, :], axis=1)
            right_gate_acc += tl.sum(right_gate_w * x_vals[None, :], axis=1)
            out_gate_acc += tl.sum(out_gate_w * x_vals[None, :], axis=1)

        left_gate_sig = tl.sigmoid(left_gate_acc)
        right_gate_sig = tl.sigmoid(right_gate_acc)
        out_gate_sig = tl.sigmoid(out_gate_acc)

        left_result = left_proj_acc * mask_val * left_gate_sig
        right_result = right_proj_acc * mask_val * right_gate_sig

        out_offs = out_base + h_offs
        tl.store(left_out_ptr + out_offs, left_result, mask=h_mask)
        tl.store(right_out_ptr + out_offs, right_result, mask=h_mask)
        tl.store(out_gate_ptr + out_offs, out_gate_sig, mask=h_mask)


@triton.jit
def _triangular_mul_kernel(
    x0_ptr, x1_ptr, out_ptr,
    B, N, H,
    BLOCK_I: tl.constexpr, BLOCK_J: tl.constexpr, BLOCK_K: tl.constexpr,
):
    pid_b = tl.program_id(0)
    pid_ij = tl.program_id(1)
    pid_d = tl.program_id(2)

    num_tiles_j = tl.cdiv(N, BLOCK_J)
    pid_i = pid_ij // num_tiles_j
    pid_j = pid_ij % num_tiles_j

    offs_i = pid_i * BLOCK_I + tl.arange(0, BLOCK_I)
    offs_j = pid_j * BLOCK_J + tl.arange(0, BLOCK_J)
    mask_i = offs_i < N
    mask_j = offs_j < N

    acc = tl.zeros((BLOCK_I, BLOCK_J), dtype=tl.float32)
    stride_b = N * N * H
    stride_i = N * H
    stride_k = H
    base_b = pid_b * stride_b
    base_d = pid_d

    for k_start in range(0, N, BLOCK_K):
        offs_k = k_start + tl.arange(0, BLOCK_K)
        mask_k = offs_k < N

        x0_offsets = base_b + offs_i[:, None] * stride_i + offs_k[None, :] * stride_k + base_d
        x0 = tl.load(x0_ptr + x0_offsets, mask=mask_i[:, None] & mask_k[None, :], other=0.0).to(tl.float32)

        x1_offsets = base_b + offs_j[:, None] * stride_i + offs_k[None, :] * stride_k + base_d
        x1 = tl.load(x1_ptr + x1_offsets, mask=mask_j[:, None] & mask_k[None, :], other=0.0).to(tl.float32)

        acc += tl.dot(x0, tl.trans(x1))

    out_offsets = base_b + offs_i[:, None] * stride_i + offs_j[None, :] * stride_k + base_d
    tl.store(out_ptr + out_offsets, acc, mask=mask_i[:, None] & mask_j[None, :])


@triton.jit
def _output_norm_gate_proj_kernel(
    x0_ptr, x1_ptr, ln_weight_ptr, ln_bias_ptr, linear_weight_ptr, out_ptr,
    num_rows, H: tl.constexpr, C, eps,
    BLOCK_H: tl.constexpr, BLOCK_C: tl.constexpr,
):
    row_idx = tl.program_id(0)
    if row_idx >= num_rows:
        return

    offs_h = tl.arange(0, BLOCK_H)
    mask_h = offs_h < H

    x0 = tl.load(x0_ptr + row_idx * H + offs_h, mask=mask_h, other=0.0).to(tl.float32)
    sum_x = tl.sum(x0, axis=0)
    mean = sum_x / H
    x0_centered = x0 - mean
    sum_sq = tl.sum(x0_centered * x0_centered, axis=0)
    var = sum_sq / H
    rstd = tl.rsqrt(var + eps)
    x_norm = x0_centered * rstd

    ln_w = tl.load(ln_weight_ptr + offs_h, mask=mask_h, other=1.0).to(tl.float32)
    ln_b = tl.load(ln_bias_ptr + offs_h, mask=mask_h, other=0.0).to(tl.float32)
    x_ln = x_norm * ln_w + ln_b

    x1 = tl.load(x1_ptr + row_idx * H + offs_h, mask=mask_h, other=0.0).to(tl.float32)
    x_gated = x_ln * x1
    x_gated = tl.where(mask_h, x_gated, 0.0)

    for c_start in range(0, C, BLOCK_C):
        offs_c = c_start + tl.arange(0, BLOCK_C)
        mask_c = offs_c < C
        weight_ptrs = linear_weight_ptr + offs_c[:, None] * H + offs_h[None, :]
        weights = tl.load(weight_ptrs, mask=mask_c[:, None] & mask_h[None, :], other=0.0).to(tl.float32)
        acc = tl.sum(weights * x_gated[None, :], axis=1)
        tl.store(out_ptr + row_idx * C + offs_c, acc, mask=mask_c)


def kernel_function(input_tensor, mask, weights, config):
    B, N, _, C = input_tensor.shape
    H = config["hidden_dim"]

    # Stage 1: LayerNorm
    x_norm = torch.empty_like(input_tensor)
    n_rows = B * N * N
    BLOCK_SIZE = triton.next_power_of_2(C)
    BLOCK_SIZE = min(BLOCK_SIZE, 1024)
    _layernorm_kernel[(n_rows,)](
        input_tensor, x_norm, weights['norm.weight'], weights['norm.bias'],
        C, C, 1e-5, BLOCK_SIZE=BLOCK_SIZE
    )

    # Stage 2: Projection + Gating
    left = torch.empty((B, N, N, H), dtype=torch.float32, device='cuda')
    right = torch.empty((B, N, N, H), dtype=torch.float32, device='cuda')
    out_gate = torch.empty((B, N, N, H), dtype=torch.float32, device='cuda')

    BLOCK_H = min(32, H)
    BLOCK_C = min(32, C)
    _projection_gating_kernel[(B * N * N,)](
        x_norm, mask,
        weights['left_proj.weight'], weights['right_proj.weight'],
        weights['left_gate.weight'], weights['right_gate.weight'], weights['out_gate.weight'],
        left, right, out_gate,
        B, N, C, H, BLOCK_H, BLOCK_C
    )

    # Stage 3: Triangular multiplication
    tri_out = torch.empty((B, N, N, H), dtype=torch.float32, device='cuda')
    BLOCK_I, BLOCK_J, BLOCK_K = 16, 16, 16
    num_tiles_i = triton.cdiv(N, BLOCK_I)
    num_tiles_j = triton.cdiv(N, BLOCK_J)
    _triangular_mul_kernel[(B, num_tiles_i * num_tiles_j, H)](
        left, right, tri_out, B, N, H, BLOCK_I, BLOCK_J, BLOCK_K
    )

    # Stage 4: Output LayerNorm + Gate + Linear
    output = torch.empty((B, N, N, C), dtype=torch.float32, device='cuda')
    BLOCK_H_out = triton.next_power_of_2(H)
    BLOCK_C_out = min(128, triton.next_power_of_2(C))
    _output_norm_gate_proj_kernel[(n_rows,)](
        tri_out.reshape(-1, H), out_gate.reshape(-1, H),
        weights['to_out_norm.weight'], weights['to_out_norm.bias'], weights['to_out.weight'],
        output.reshape(-1, C), n_rows, H, C, 1e-5, BLOCK_H_out, BLOCK_C_out
    )

    return output


def test_kernel():
    torch.manual_seed(42)
    B, N, C, H = 1, 32, 128, 128

    input_tensor = torch.randn(B, N, N, C, device='cuda', dtype=torch.float32)
    mask = torch.ones(B, N, N, device='cuda', dtype=torch.float32)

    weights = {
        'norm.weight': torch.randn(C, device='cuda'), 'norm.bias': torch.randn(C, device='cuda'),
        'left_proj.weight': torch.randn(H, C, device='cuda') / (H**0.5),
        'right_proj.weight': torch.randn(H, C, device='cuda') / (H**0.5),
        'left_gate.weight': torch.randn(H, C, device='cuda') / (H**0.5),
        'right_gate.weight': torch.randn(H, C, device='cuda') / (H**0.5),
        'out_gate.weight': torch.randn(H, C, device='cuda') / (H**0.5),
        'to_out_norm.weight': torch.randn(H, device='cuda'), 'to_out_norm.bias': torch.randn(H, device='cuda'),
        'to_out.weight': torch.randn(C, H, device='cuda') / (C**0.5),
    }
    config = {"dim": C, "hidden_dim": H}

    out_triton = kernel_function(input_tensor, mask, weights, config)

    # Reference
    from torch import nn, einsum
    x = torch.nn.functional.layer_norm(input_tensor, [C], weights['norm.weight'], weights['norm.bias'])
    left = x @ weights['left_proj.weight'].T * mask.unsqueeze(-1) * torch.sigmoid(x @ weights['left_gate.weight'].T)
    right = x @ weights['right_proj.weight'].T * mask.unsqueeze(-1) * torch.sigmoid(x @ weights['right_gate.weight'].T)
    out_gate = torch.sigmoid(x @ weights['out_gate.weight'].T)
    tri = einsum('bikd,bjkd->bijd', left, right)
    tri_norm = torch.nn.functional.layer_norm(tri, [H], weights['to_out_norm.weight'], weights['to_out_norm.bias'])
    ref = (tri_norm * out_gate) @ weights['to_out.weight'].T

    if torch.allclose(out_triton, ref, rtol=2e-2, atol=2e-2):
        print("PASS")
    else:
        print(f"FAIL: max diff = {(out_triton - ref).abs().max().item()}")


if __name__ == "__main__":
    test_kernel()


def custom_kernel(input):
    return kernel_function(*input)
scrolls · 272 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 416249.

- from typing import Dict, Tuple, TypeVar
-
import torch
import triton
import triton.language as tl
⋯ 1 unchanged lines
@triton.jit
def _layernorm_kernel(
- x_ptr,
- out_ptr,
- gamma_ptr,
- beta_ptr,
- N,
- D,
- stride_n,
- stride_d,
- eps: tl.constexpr,
- BLOCK_D: tl.constexpr,
+ X_ptr, Y_ptr, W_ptr, B_ptr,
+ stride_x, C, eps,
+ BLOCK_SIZE: tl.constexpr,
):
- pid = tl.program_id(0)
- offs_d = tl.arange(0, BLOCK_D)
- mask = offs_d < D
+ row_idx = tl.program_id(0)
+ row_start_ptr = X_ptr + row_idx * stride_x
+ col_offsets = tl.arange(0, BLOCK_SIZE)
+ mask = col_offsets < C
- x_ptrs = x_ptr + pid * stride_n + offs_d * stride_d
- x = tl.load(x_ptrs, mask=mask, other=0.0)
-
- mean = tl.sum(x, axis=0) / D
- x_centered = x - mean
- var = tl.sum(x_centered * x_centered, axis=0) / D
+ x = tl.load(row_start_ptr + col_offsets, mask=mask, other=0.0).to(tl.float32)
+ x_sum = tl.sum(x, axis=0)
+ mean = x_sum / C
+ x_centered = tl.where(mask, x - mean, 0.0)
+ var_sum = tl.sum(x_centered * x_centered, axis=0)
+ var = var_sum / C
rstd = 1.0 / tl.sqrt(var + eps)
+ x_norm = x_centered * rstd
- gamma = tl.load(gamma_ptr + offs_d, mask=mask, other=1.0)
- beta = tl.load(beta_ptr + offs_d, mask=mask, other=0.0)
+ w = tl.load(W_ptr + col_offsets, mask=mask, other=1.0).to(tl.float32)
+ b = tl.load(B_ptr + col_offsets, mask=mask, other=0.0).to(tl.float32)
+ y = x_norm * w + b
- out = x_centered * rstd * gamma + beta
- out_ptrs = out_ptr + pid * stride_n + offs_d * stride_d
- tl.store(out_ptrs, out, mask=mask)
+ out_ptr = Y_ptr + row_idx * stride_x
+ tl.store(out_ptr + col_offsets, y, mask=mask)
@triton.jit
- def _linear_sigmoid_kernel(
- x_ptr,
- w_ptr,
- out_ptr,
- M,
- K,
- N,
- stride_xm,
- stride_xk,
- stride_wn,
- stride_wk,
- stride_om,
- stride_on,
- apply_sigmoid: tl.constexpr,
- BLOCK_M: tl.constexpr,
- BLOCK_N: tl.constexpr,
- BLOCK_K: tl.constexpr,
+ def _projection_gating_kernel(
+ x_ptr, mask_ptr,
+ left_proj_weight_ptr, right_proj_weight_ptr,
+ left_gate_weight_ptr, right_gate_weight_ptr, out_gate_weight_ptr,
+ left_out_ptr, right_out_ptr, out_gate_ptr,
+ B, N, C, H,
+ BLOCK_H: tl.constexpr, BLOCK_C: tl.constexpr,
):
- pid_m = tl.program_id(0)
- pid_n = tl.program_id(1)
+ pid = tl.program_id(0)
+ b = pid // (N * N)
+ remainder = pid % (N * N)
+ i = remainder // N
+ j = remainder % N
- offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
- offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
- offs_k = tl.arange(0, BLOCK_K)
+ x_base = b * (N * N * C) + i * (N * C) + j * C
+ mask_idx = b * (N * N) + i * N + j
+ mask_val = tl.load(mask_ptr + mask_idx).to(tl.float32)
+ out_base = b * (N * N * H) + i * (N * H) + j * H
- acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
+ for h_start in range(0, H, BLOCK_H):
+ h_offs = h_start + tl.arange(0, BLOCK_H)
+ h_mask = h_offs < H
- for k_start in range(0, K, BLOCK_K):
- k_offs = k_start + offs_k
- k_mask = k_offs < K
+ left_proj_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)
+ right_proj_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)
+ left_gate_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)
+ right_gate_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)
+ out_gate_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)
- x_ptrs = x_ptr + offs_m[:, None] * stride_xm + k_offs[None, :] * stride_xk
- w_ptrs = w_ptr + offs_n[:, None] * stride_wn + k_offs[None, :] * stride_wk
+ for c_start in range(0, C, BLOCK_C):
+ c_offs = c_start + tl.arange(0, BLOCK_C)
+ c_mask = c_offs < C
+ x_vals = tl.load(x_ptr + x_base + c_offs, mask=c_mask, other=0.0).to(tl.float32)
- x_mask = (offs_m[:, None] < M) & k_mask[None, :]
- w_mask = (offs_n[:, None] < N) & k_mask[None, :]
+ weight_offsets = h_offs[:, None] * C + c_offs[None, :]
+ combined_mask = h_mask[:, None] & c_mask[None, :]
- x_vals = tl.load(x_ptrs, mask=x_mask, other=0.0)
- w_vals = tl.load(w_ptrs, mask=w_mask, other=0.0)
+ left_proj_w = tl.load(left_proj_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)
+ right_proj_w = tl.load(right_proj_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)
+ left_gate_w = tl.load(left_gate_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)
+ right_gate_w = tl.load(right_gate_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)
+ out_gate_w = tl.load(out_gate_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)
- acc += tl.dot(x_vals, tl.trans(w_vals))
+ left_proj_acc += tl.sum(left_proj_w * x_vals[None, :], axis=1)
+ right_proj_acc += tl.sum(right_proj_w * x_vals[None, :], axis=1)
+ left_gate_acc += tl.sum(left_gate_w * x_vals[None, :], axis=1)
+ right_gate_acc += tl.sum(right_gate_w * x_vals[None, :], axis=1)
+ out_gate_acc += tl.sum(out_gate_w * x_vals[None, :], axis=1)
- if apply_sigmoid:
- acc = tl.sigmoid(acc)
+ left_gate_sig = tl.sigmoid(left_gate_acc)
+ right_gate_sig = tl.sigmoid(right_gate_acc)
+ out_gate_sig = tl.sigmoid(out_gate_acc)
- out_ptrs = out_ptr + offs_m[:, None] * stride_om + offs_n[None, :] * stride_on
- out_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
- tl.store(out_ptrs, acc, mask=out_mask)
+ left_result = left_proj_acc * mask_val * left_gate_sig
+ right_result = right_proj_acc * mask_val * right_gate_sig
+ out_offs = out_base + h_offs
+ tl.store(left_out_ptr + out_offs, left_result, mask=h_mask)
+ tl.store(right_out_ptr + out_offs, right_result, mask=h_mask)
+ tl.store(out_gate_ptr + out_offs, out_gate_sig, mask=h_mask)
+
@triton.jit
- def _masked_gating_einsum_kernel(
- left_ptr,
- right_ptr,
- left_gate_ptr,
- right_gate_ptr,
- mask_ptr,
- out_ptr,
- B,
- seq_len,
- hidden_dim,
- stride_lb,
- stride_li,
- stride_lk,
- stride_ld,
- stride_rb,
- stride_ri,
- stride_rk,
- stride_rd,
- stride_lgb,
- stride_lgi,
- stride_lgk,
- stride_lgd,
- stride_rgb,
- stride_rgi,
- stride_rgk,
- stride_rgd,
- stride_mb,
- stride_mi,
- stride_mk,
- stride_ob,
- stride_oi,
- stride_oj,
- stride_od,
- BLOCK_D: tl.constexpr,
+ def _triangular_mul_kernel(
+ x0_ptr, x1_ptr, out_ptr,
+ B, N, H,
+ BLOCK_I: tl.constexpr, BLOCK_J: tl.constexpr, BLOCK_K: tl.constexpr,
):
- pid_bi = tl.program_id(0)
- pid_j = tl.program_id(1)
+ pid_b = tl.program_id(0)
+ pid_ij = tl.program_id(1)
pid_d = tl.program_id(2)
- pid_b = pid_bi // seq_len
- pid_i = pid_bi % seq_len
+ num_tiles_j = tl.cdiv(N, BLOCK_J)
+ pid_i = pid_ij // num_tiles_j
+ pid_j = pid_ij % num_tiles_j
- d_offs = pid_d * BLOCK_D + tl.arange(0, BLOCK_D)
- d_mask = d_offs < hidden_dim
+ offs_i = pid_i * BLOCK_I + tl.arange(0, BLOCK_I)
+ offs_j = pid_j * BLOCK_J + tl.arange(0, BLOCK_J)
+ mask_i = offs_i < N
+ mask_j = offs_j < N
- acc = tl.zeros((BLOCK_D,), dtype=tl.float32)
+ acc = tl.zeros((BLOCK_I, BLOCK_J), dtype=tl.float32)
+ stride_b = N * N * H
+ stride_i = N * H
+ stride_k = H
+ base_b = pid_b * stride_b
+ base_d = pid_d
- for k in range(seq_len):
- mask_l = tl.load(
- mask_ptr + pid_b * stride_mb + pid_i * stride_mi + k * stride_mk
- )
- mask_r = tl.load(
- mask_ptr + pid_b * stride_mb + pid_j * stride_mi + k * stride_mk
- )
+ for k_start in range(0, N, BLOCK_K):
+ offs_k = k_start + tl.arange(0, BLOCK_K)
+ mask_k = offs_k < N
- left_off = (
- pid_b * stride_lb + pid_i * stride_li + k * stride_lk + d_offs * stride_ld
- )
- left_val = tl.load(left_ptr + left_off, mask=d_mask, other=0.0)
- lg_off = (
- pid_b * stride_lgb
- + pid_i * stride_lgi
- + k * stride_lgk
- + d_offs * stride_lgd
- )
- lg_val = tl.load(left_gate_ptr + lg_off, mask=d_mask, other=0.0)
- left_gated = left_val * mask_l * lg_val
+ x0_offsets = base_b + offs_i[:, None] * stride_i + offs_k[None, :] * stride_k + base_d
+ x0 = tl.load(x0_ptr + x0_offsets, mask=mask_i[:, None] & mask_k[None, :], other=0.0).to(tl.float32)
- right_off = (
- pid_b * stride_rb + pid_j * stride_ri + k * stride_rk + d_offs * stride_rd
- )
- right_val = tl.load(right_ptr + right_off, mask=d_mask, other=0.0)
- rg_off = (
- pid_b * stride_rgb
- + pid_j * stride_rgi
- + k * stride_rgk
- + d_offs * stride_rgd
- )
- rg_val = tl.load(right_gate_ptr + rg_off, mask=d_mask, other=0.0)
- right_gated = right_val * mask_r * rg_val
+ x1_offsets = base_b + offs_j[:, None] * stride_i + offs_k[None, :] * stride_k + base_d
+ x1 = tl.load(x1_ptr + x1_offsets, mask=mask_j[:, None] & mask_k[None, :], other=0.0).to(tl.float32)
- acc += left_gated * right_gated
+ acc += tl.dot(x0, tl.trans(x1))
- out_off = (
- pid_b * stride_ob + pid_i * stride_oi + pid_j * stride_oj + d_offs * stride_od
- )
- tl.store(out_ptr + out_off, acc, mask=d_mask)
+ out_offsets = base_b + offs_i[:, None] * stride_i + offs_j[None, :] * stride_k + base_d
+ tl.store(out_ptr + out_offsets, acc, mask=mask_i[:, None] & mask_j[None, :])
@triton.jit
- def _layernorm_gate_linear_kernel(
- x_ptr,
- gate_ptr,
- gamma_ptr,
- beta_ptr,
- w_ptr,
- out_ptr,
- N,
- hidden_dim,
- out_dim,
- stride_xn,
- stride_xd,
- stride_gn,
- stride_gd,
- stride_wout,
- stride_whid,
- stride_on,
- stride_od,
- eps: tl.constexpr,
- BLOCK_HD: tl.constexpr,
- BLOCK_OUT: tl.constexpr,
+ def _output_norm_gate_proj_kernel(
+ x0_ptr, x1_ptr, ln_weight_ptr, ln_bias_ptr, linear_weight_ptr, out_ptr,
+ num_rows, H: tl.constexpr, C, eps,
+ BLOCK_H: tl.constexpr, BLOCK_C: tl.constexpr,
):
- pid_n = tl.program_id(0)
- pid_out = tl.program_id(1)
+ row_idx = tl.program_id(0)
+ if row_idx >= num_rows:
+ return
- offs_hd = tl.arange(0, BLOCK_HD)
- hd_mask = offs_hd < hidden_dim
+ offs_h = tl.arange(0, BLOCK_H)
+ mask_h = offs_h < H
- x_ptrs = x_ptr + pid_n * stride_xn + offs_hd * stride_xd
- x = tl.load(x_ptrs, mask=hd_mask, other=0.0)
+ x0 = tl.load(x0_ptr + row_idx * H + offs_h, mask=mask_h, other=0.0).to(tl.float32)
+ sum_x = tl.sum(x0, axis=0)
+ mean = sum_x / H
+ x0_centered = x0 - mean
+ sum_sq = tl.sum(x0_centered * x0_centered, axis=0)
+ var = sum_sq / H
+ rstd = tl.rsqrt(var + eps)
+ x_norm = x0_centered * rstd
- mean = tl.sum(x, axis=0) / hidden_dim
- x_centered = x - mean
- var = tl.sum(x_centered * x_centered, axis=0) / hidden_dim
- rstd = 1.0 / tl.sqrt(var + eps)
+ ln_w = tl.load(ln_weight_ptr + offs_h, mask=mask_h, other=1.0).to(tl.float32)
+ ln_b = tl.load(ln_bias_ptr + offs_h, mask=mask_h, other=0.0).to(tl.float32)
+ x_ln = x_norm * ln_w + ln_b
- gamma = tl.load(gamma_ptr + offs_hd, mask=hd_mask, other=1.0)
- beta = tl.load(beta_ptr + offs_hd, mask=hd_mask, other=0.0)
- x_normed = x_centered * rstd * gamma + beta
+ x1 = tl.load(x1_ptr + row_idx * H + offs_h, mask=mask_h, other=0.0).to(tl.float32)
+ x_gated = x_ln * x1
+ x_gated = tl.where(mask_h, x_gated, 0.0)
- gate_ptrs = gate_ptr + pid_n * stride_gn + offs_hd * stride_gd
- gate = tl.load(gate_ptrs, mask=hd_mask, other=0.0)
- x_gated = x_normed * gate
+ for c_start in range(0, C, BLOCK_C):
+ offs_c = c_start + tl.arange(0, BLOCK_C)
+ mask_c = offs_c < C
+ weight_ptrs = linear_weight_ptr + offs_c[:, None] * H + offs_h[None, :]
+ weights = tl.load(weight_ptrs, mask=mask_c[:, None] & mask_h[None, :], other=0.0).to(tl.float32)
+ acc = tl.sum(weights * x_gated[None, :], axis=1)
+ tl.store(out_ptr + row_idx * C + offs_c, acc, mask=mask_c)
- offs_out = pid_out * BLOCK_OUT + tl.arange(0, BLOCK_OUT)
- out_mask = offs_out < out_dim
- acc = tl.zeros((BLOCK_OUT,), dtype=tl.float32)
- for hd_start in range(0, hidden_dim, BLOCK_HD):
- hd_offs = hd_start + tl.arange(0, BLOCK_HD)
- hd_m = hd_offs < hidden_dim
+ def kernel_function(input_tensor, mask, weights, config):
+ B, N, _, C = input_tensor.shape
+ H = config["hidden_dim"]
- if hd_start == 0:
- x_chunk = x_gated
- else:
- x_ptrs2 = x_ptr + pid_n * stride_xn + hd_offs * stride_xd
- x2 = tl.load(x_ptrs2, mask=hd_m, other=0.0)
- mean2 = tl.sum(x2, axis=0) / hidden_dim
- var2 = tl.sum((x2 - mean2) * (x2 - mean2), axis=0) / hidden_dim
- rstd2 = 1.0 / tl.sqrt(var2 + eps)
- gamma2 = tl.load(gamma_ptr + hd_offs, mask=hd_m, other=1.0)
- beta2 = tl.load(beta_ptr + hd_offs, mask=hd_m, other=0.0)
- x_normed2 = (x2 - mean2) * rstd2 * gamma2 + beta2
- gate2 = tl.load(
- gate_ptr + pid_n * stride_gn + hd_offs * stride_gd, mask=hd_m, other=0.0
- )
- x_chunk = x_normed2 * gate2
+ # Stage 1: LayerNorm
+ x_norm = torch.empty_like(input_tensor)
+ n_rows = B * N * N
+ BLOCK_SIZE = triton.next_power_of_2(C)
+ BLOCK_SIZE = min(BLOCK_SIZE, 1024)
+ _layernorm_kernel[(n_rows,)](
+ input_tensor, x_norm, weights['norm.weight'], weights['norm.bias'],
+ C, C, 1e-5, BLOCK_SIZE=BLOCK_SIZE
+ )
- for out_idx in range(BLOCK_OUT):
- if pid_out * BLOCK_OUT + out_idx < out_dim:
- w_ptrs = (
- w_ptr
- + (pid_out * BLOCK_OUT + out_idx) * stride_wout
- + hd_offs * stride_whid
- )
- w_vals = tl.load(w_ptrs, mask=hd_m, other=0.0)
- acc = tl.where(
- tl.arange(0, BLOCK_OUT) == out_idx,
- acc + tl.sum(x_chunk * w_vals),
- acc,
- )
- break
+ # Stage 2: Projection + Gating
+ left = torch.empty((B, N, N, H), dtype=torch.float32, device='cuda')
+ right = torch.empty((B, N, N, H), dtype=torch.float32, device='cuda')
+ out_gate = torch.empty((B, N, N, H), dtype=torch.float32, device='cuda')
- out_ptrs = out_ptr + pid_n * stride_on + offs_out * stride_od
- tl.store(out_ptrs, acc, mask=out_mask)
-
-
- def kernel_function(x, mask, weights, config):
- B, seq_len, _, dim = x.shape
- hidden_dim = config["hidden_dim"]
- device = x.device
- dtype = x.dtype
-
- # Flatten for processing
- N = B * seq_len * seq_len
- x_flat = x.reshape(N, dim).contiguous()
-
- # LayerNorm
- x_normed = torch.empty_like(x_flat)
- BLOCK_D = triton.next_power_of_2(dim)
- _layernorm_kernel[(N,)](
- x_flat,
- x_normed,
- weights["norm.weight"],
- weights["norm.bias"],
- N,
- dim,
- dim,
- 1,
- eps=1e-5,
- BLOCK_D=BLOCK_D,
+ BLOCK_H = min(32, H)
+ BLOCK_C = min(32, C)
+ _projection_gating_kernel[(B * N * N,)](
+ x_norm, mask,
+ weights['left_proj.weight'], weights['right_proj.weight'],
+ weights['left_gate.weight'], weights['right_gate.weight'], weights['out_gate.weight'],
+ left, right, out_gate,
+ B, N, C, H, BLOCK_H, BLOCK_C
)
- # Linear projections using torch.mm (for simplicity, we use matmul here)
- x_normed_2d = x_normed.view(N, dim)
- left = x_normed_2d @ weights["left_proj.weight"].T
- right = x_normed_2d @ weights["right_proj.weight"].T
- left_gate = torch.sigmoid(x_normed_2d @ weights["left_gate.weight"].T)
- right_gate = torch.sigmoid(x_normed_2d @ weights["right_gate.weight"].T)
- out_gate = torch.sigmoid(x_normed_2d @ weights["out_gate.weight"].T)
-
- # Reshape back
- left = left.view(B, seq_len, seq_len, hidden_dim)
- right = right.view(B, seq_len, seq_len, hidden_dim)
- left_gate = left_gate.view(B, seq_len, seq_len, hidden_dim)
- right_gate = right_gate.view(B, seq_len, seq_len, hidden_dim)
- out_gate = out_gate.view(B, seq_len, seq_len, hidden_dim)
-
- # Masked gating einsum
- out = torch.empty((B, seq_len, seq_len, hidden_dim), device=device, dtype=dtype)
- BLOCK_D = min(64, triton.next_power_of_2(hidden_dim))
- grid = (B * seq_len, seq_len, triton.cdiv(hidden_dim, BLOCK_D))
- _masked_gating_einsum_kernel[grid](
- left,
- right,
- left_gate,
- right_gate,
- mask.float(),
- out,
- B,
- seq_len,
- hidden_dim,
- *left.stride(),
- *right.stride(),
- *left_gate.stride(),
- *right_gate.stride(),
- *mask.stride(),
- *out.stride(),
- BLOCK_D=BLOCK_D,
+ # Stage 3: Triangular multiplication
+ tri_out = torch.empty((B, N, N, H), dtype=torch.float32, device='cuda')
+ BLOCK_I, BLOCK_J, BLOCK_K = 16, 16, 16
+ num_tiles_i = triton.cdiv(N, BLOCK_I)
+ num_tiles_j = triton.cdiv(N, BLOCK_J)
+ _triangular_mul_kernel[(B, num_tiles_i * num_tiles_j, H)](
+ left, right, tri_out, B, N, H, BLOCK_I, BLOCK_J, BLOCK_K
)
- # Output: LayerNorm, gate, linear
- out_flat = out.reshape(N, hidden_dim)
- out_normed = torch.empty_like(out_flat)
- BLOCK_HD = triton.next_power_of_2(hidden_dim)
- _layernorm_kernel[(N,)](
- out_flat,
- out_normed,
- weights["to_out_norm.weight"],
- weights["to_out_norm.bias"],
- N,
- hidden_dim,
- hidden_dim,
- 1,
- eps=1e-5,
- BLOCK_D=BLOCK_HD,
+ # Stage 4: Output LayerNorm + Gate + Linear
+ output = torch.empty((B, N, N, C), dtype=torch.float32, device='cuda')
+ BLOCK_H_out = triton.next_power_of_2(H)
+ BLOCK_C_out = min(128, triton.next_power_of_2(C))
+ _output_norm_gate_proj_kernel[(n_rows,)](
+ tri_out.reshape(-1, H), out_gate.reshape(-1, H),
+ weights['to_out_norm.weight'], weights['to_out_norm.bias'], weights['to_out.weight'],
+ output.reshape(-1, C), n_rows, H, C, 1e-5, BLOCK_H_out, BLOCK_C_out
)
- out_gated = out_normed * out_gate.reshape(N, hidden_dim)
- result = out_gated @ weights["to_out.weight"].T
- return result.view(B, seq_len, seq_len, dim)
+ return output
- def test_kernel():
- from math import sqrt
+ def test_kernel():
torch.manual_seed(42)
- B, seq_len, dim, hidden_dim = 2, 8, 32, 64
- x = torch.randn(B, seq_len, seq_len, dim, device="cuda")
- mask = torch.randint(0, 2, (B, seq_len, seq_len), device="cuda").float()
+ B, N, C, H = 1, 32, 128, 128
+
+ input_tensor = torch.randn(B, N, N, C, device='cuda', dtype=torch.float32)
+ mask = torch.ones(B, N, N, device='cuda', dtype=torch.float32)
+
weights = {
- "norm.weight": torch.randn(dim, device="cuda"),
- "norm.bias": torch.randn(dim, device="cuda"),
- "left_proj.weight": torch.randn(hidden_dim, dim, device="cuda")
- / sqrt(hidden_dim),
- "right_proj.weight": torch.randn(hidden_dim, dim, device="cuda")
- / sqrt(hidden_dim),
- "left_gate.weight": torch.randn(hidden_dim, dim, device="cuda")
- / sqrt(hidden_dim),
- "right_gate.weight": torch.randn(hidden_dim, dim, device="cuda")
- / sqrt(hidden_dim),
- "out_gate.weight": torch.randn(hidden_dim, dim, device="cuda")
- / sqrt(hidden_dim),
- "to_out_norm.weight": torch.randn(hidden_dim, device="cuda"),
- "to_out_norm.bias": torch.randn(hidden_dim, device="cuda"),
- "to_out.weight": torch.randn(dim, hidden_dim, device="cuda") / sqrt(dim),
+ 'norm.weight': torch.randn(C, device='cuda'), 'norm.bias': torch.randn(C, device='cuda'),
+ 'left_proj.weight': torch.randn(H, C, device='cuda') / (H**0.5),
+ 'right_proj.weight': torch.randn(H, C, device='cuda') / (H**0.5),
+ 'left_gate.weight': torch.randn(H, C, device='cuda') / (H**0.5),
+ 'right_gate.weight': torch.randn(H, C, device='cuda') / (H**0.5),
+ 'out_gate.weight': torch.randn(H, C, device='cuda') / (H**0.5),
+ 'to_out_norm.weight': torch.randn(H, device='cuda'), 'to_out_norm.bias': torch.randn(H, device='cuda'),
+ 'to_out.weight': torch.randn(C, H, device='cuda') / (C**0.5),
}
- config = {"dim": dim, "hidden_dim": hidden_dim}
- out_triton = kernel_function(x, mask, weights, config)
+ config = {"dim": C, "hidden_dim": H}
- from torch import nn
+ out_triton = kernel_function(input_tensor, mask, weights, config)
- model = nn.Module()
- x_n = nn.functional.layer_norm(
- x, [dim], weights["norm.weight"], weights["norm.bias"]
- )
- left = x_n @ weights["left_proj.weight"].T
- right = x_n @ weights["right_proj.weight"].T
- lg = torch.sigmoid(x_n @ weights["left_gate.weight"].T)
- rg = torch.sigmoid(x_n @ weights["right_gate.weight"].T)
- og = torch.sigmoid(x_n @ weights["out_gate.weight"].T)
- m = mask.unsqueeze(-1)
- left = left * m * lg
- right = right * m * rg
- out = torch.einsum("bikd,bjkd->bijd", left, right)
- out = nn.functional.layer_norm(
- out, [hidden_dim], weights["to_out_norm.weight"], weights["to_out_norm.bias"]
- )
- out = out * og
- out_ref = out @ weights["to_out.weight"].T
+ # Reference
+ from torch import nn, einsum
+ x = torch.nn.functional.layer_norm(input_tensor, [C], weights['norm.weight'], weights['norm.bias'])
+ left = x @ weights['left_proj.weight'].T * mask.unsqueeze(-1) * torch.sigmoid(x @ weights['left_gate.weight'].T)
+ right = x @ weights['right_proj.weight'].T * mask.unsqueeze(-1) * torch.sigmoid(x @ weights['right_gate.weight'].T)
+ out_gate = torch.sigmoid(x @ weights['out_gate.weight'].T)
+ tri = einsum('bikd,bjkd->bijd', left, right)
+ tri_norm = torch.nn.functional.layer_norm(tri, [H], weights['to_out_norm.weight'], weights['to_out_norm.bias'])
+ ref = (tri_norm * out_gate) @ weights['to_out.weight'].T
- if torch.allclose(out_triton, out_ref, rtol=2e-2, atol=2e-2):
+ if torch.allclose(out_triton, ref, rtol=2e-2, atol=2e-2):
print("PASS")
else:
- print(f"FAIL: max diff = {(out_triton - out_ref).abs().max()}")
+ print(f"FAIL: max diff = {(out_triton - ref).abs().max().item()}")
- # if __name__ == "__main__":
- # test_kernel()
+ if __name__ == "__main__":
+ test_kernel()
- input_t = TypeVar(
- "input_t", bound=Tuple[torch.Tensor, torch.Tensor, Dict[str, torch.Tensor], Dict]
- )
-
-
- def custom_kernel(input: input_t) -> torch.Tensor:
- x, mask, weights, config = input
- return kernel_function(x, mask, weights, config)
+ def custom_kernel(input):
+ return kernel_function(*input)
scrolls · 619 diff lines total

Best evidence level for this revision: reported

JSON