Skip to content
KernelIndex
Search⌘K

submission 480973

Cookie 🍪 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

trimul_H100_gpt-5_ka_submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-480973?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
32.0ms
#63 of 71
2026-02-06

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

mmaacc_proj += tl.dot(a_tile, wproj_tile) # [M_blk, U_blk]
num-warps = 4num_warps=4, num_stages=2
stages = 2num_warps=4, num_stages=2
tile-k = 64BLOCK_K = 64
tile-m = 64BLOCK_M = 64
tile-n = 64BLOCK_N = 64

Kernel source

trimul_H100_gpt-5_ka_submission.py479 lines
# kernel.py
# Triangle Multiplicative Update (Outgoing) -- Forward-only Triton implementation.
# All compute is performed inside Triton kernels; Python only validates, allocates, and launches.
#
# Fused stages implemented across kernels (aggressively fused where feasible):
#   1) LayerNorm over D for x[b, i, j, :]
#   2) Fused linear(gate+proj) + sigmoid + masking for LEFT and RIGHT paths on rows [B, N, N]
#   3) Contraction over k: S[b, i, j, u] = sum_k LEFT[b, i, k, u] * RIGHT[b, j, k, u]
#   4) OUT gate per (b, i, j): sigmoid(x_norm[b,i,j,:] @ W_out_gate^T)
#   5) Final LayerNorm over U + affine + gate + matmul to D: out[b,i,j,:]
#
# Fusion boundaries:
# - The proj+gate+sigmoid(+mask) is fused per path to minimize memory traffic.
# - The N-reduction over k (stage 3) and the U-epilogue + projection (stage 5) use incompatible tilings:
#   * The contraction reduces along N while keeping (i,j,u), whereas the epilogue normalizes along U and projects to D.
#   * A single-kernel fusion would cause extreme register/shared-memory pressure and poor occupancy.
#   Therefore, they are launched as separate kernels for performance and maintainability.

import triton
import triton.language as tl
import torch


@triton.jit
def layer_norm_lastdim_kernel(x_ptr, y_ptr, gamma_ptr, beta_ptr,
                              M, D, eps,
                              BLOCK_D: tl.constexpr):
    # Each program normalizes one row of length D
    pid = tl.program_id(axis=0)
    if pid >= M:
        return

    row_x = x_ptr + pid * D
    row_y = y_ptr + pid * D
    offs = tl.arange(0, BLOCK_D)

    # Pass 1: mean
    sum_val = tl.zeros((), dtype=tl.float32)
    for d0 in tl.range(0, D, BLOCK_D):
        d_ids = d0 + offs
        mask = d_ids < D
        x = tl.load(row_x + d_ids, mask=mask, other=0.0).to(tl.float32)
        sum_val += tl.sum(x, axis=0)
    d_float = tl.cast(D, tl.float32)
    mean = sum_val / d_float

    # Pass 2: variance via squared deviations
    sum_sqdiff = tl.zeros((), dtype=tl.float32)
    for d0 in tl.range(0, D, BLOCK_D):
        d_ids = d0 + offs
        mask = d_ids < D
        x = tl.load(row_x + d_ids, mask=mask, other=0.0).to(tl.float32)
        diff = x - mean
        sum_sqdiff += tl.sum(diff * diff, axis=0)
    var = sum_sqdiff / d_float
    inv_std = 1.0 / tl.sqrt(var + eps)

    # Pass 3: normalize, scale, shift
    for d0 in tl.range(0, D, BLOCK_D):
        d_ids = d0 + offs
        mask = d_ids < D
        x = tl.load(row_x + d_ids, mask=mask, other=0.0).to(tl.float32)
        g = tl.load(gamma_ptr + d_ids, mask=mask, other=1.0).to(tl.float32)
        b = tl.load(beta_ptr + d_ids, mask=mask, other=0.0).to(tl.float32)
        y = (x - mean) * inv_std
        y = y * g + b
        tl.store(row_y + d_ids, y, mask=mask)


@triton.jit
def proj_gate_fused_rows_kernel(a_ptr, wproj_ptr, wgate_ptr, mask_ptr, out_ptr,
                                M, N, D, U,
                                apply_mask: tl.constexpr,
                                BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):
    """
    Fused projection and gate over rows of a 2D tensor [M, D], where M = B * N * N and each row maps to (b, i, j).
    Produces: out[row, u] = (a_row @ wproj.T) * sigmoid(a_row @ wgate.T) * (optional mask[b,i,j]).
    """
    pid_m = tl.program_id(axis=0)
    pid_n = tl.program_id(axis=1)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)

    acc_proj = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    acc_gate = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    # Loop over K = D
    for k0 in tl.range(0, D, BLOCK_K):
        k_ids = k0 + tl.arange(0, BLOCK_K)
        k_mask = k_ids < D

        # A tile: [BLOCK_M, BLOCK_K]
        a_ptrs = a_ptr + (offs_m[:, None] * D + k_ids[None, :])
        a_tile = tl.load(
            a_ptrs,
            mask=(offs_m[:, None] < M) & (k_mask[None, :]),
            other=0.0
        ).to(tl.float32)

        # Weight tiles: [BLOCK_K, BLOCK_N] (weights are [U, D])
        wproj_ptrs = wproj_ptr + (offs_n[None, :] * D + k_ids[:, None])
        wgate_ptrs = wgate_ptr + (offs_n[None, :] * D + k_ids[:, None])

        wproj_tile = tl.load(
            wproj_ptrs,
            mask=(offs_n[None, :] < U) & (k_mask[:, None]),
            other=0.0
        ).to(tl.float32)
        wgate_tile = tl.load(
            wgate_ptrs,
            mask=(offs_n[None, :] < U) & (k_mask[:, None]),
            other=0.0
        ).to(tl.float32)

        acc_proj += tl.dot(a_tile, wproj_tile)  # [M_blk, U_blk]
        acc_gate += tl.dot(a_tile, wgate_tile)

    gate = tl.sigmoid(acc_gate)
    out = acc_proj * gate

    if apply_mask:
        NN = N * N
        b_m = offs_m // NN
        rem = offs_m - b_m * NN
        i_m = rem // N
        j_m = rem - i_m * N
        m_ptrs = mask_ptr + b_m * NN + i_m * N + j_m
        mvals = tl.load(m_ptrs, mask=offs_m < M, other=0.0).to(tl.float32)  # [M_blk]
        out = out * mvals[:, None]

    out_ptrs = out_ptr + (offs_m[:, None] * U + offs_n[None, :])
    mask_out = (offs_m[:, None] < M) & (offs_n[None, :] < U)
    tl.store(out_ptrs, out, mask=mask_out)


@triton.jit
def contract_k_kernel(L_ptr, R_ptr, S_ptr,
                      B, N, U,
                      stride_lb, stride_li, stride_lk, stride_lu,
                      stride_rb, stride_rj, stride_rk, stride_ru,
                      stride_sb, stride_si, stride_sj, stride_su,
                      BLOCK_I: tl.constexpr, BLOCK_J: tl.constexpr, BLOCK_K: tl.constexpr):
    # 3D launch: I, J, (B*U)
    pid_i = tl.program_id(axis=0)
    pid_j = tl.program_id(axis=1)
    pid_bu = tl.program_id(axis=2)

    b = pid_bu // U
    u = pid_bu % U

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

    acc = tl.zeros((BLOCK_I, BLOCK_J), dtype=tl.float32)

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

        # L[b, i, k, u] -> [BLOCK_I, BLOCK_K]
        l_ptrs = (L_ptr
                  + b * stride_lb
                  + offs_i[:, None] * stride_li
                  + offs_k[None, :] * stride_lk
                  + u * stride_lu)
        L_tile = tl.load(l_ptrs, mask=(offs_i[:, None] < N) & (k_mask[None, :]), other=0.0).to(tl.float32)

        # R[b, j, k, u] -> [BLOCK_K, BLOCK_J]
        r_ptrs = (R_ptr
                  + b * stride_rb
                  + offs_j[None, :] * stride_rj
                  + offs_k[:, None] * stride_rk
                  + u * stride_ru)
        R_tile_KJ = tl.load(r_ptrs, mask=(k_mask[:, None]) & (offs_j[None, :] < N), other=0.0).to(tl.float32)

        acc += tl.dot(L_tile, R_tile_KJ)

    s_ptrs = (S_ptr
              + b * stride_sb
              + offs_i[:, None] * stride_si
              + offs_j[None, :] * stride_sj
              + u * stride_su)
    tl.store(s_ptrs, acc, mask=(offs_i[:, None] < N) & (offs_j[None, :] < N))


@triton.jit
def out_gate_kernel(a_ptr, w_ptr, out_ptr, M, D, U,
                    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):
    # Compute: sigmoid(a @ w^T) for M rows, projecting D->U
    pid_m = tl.program_id(axis=0)
    pid_n = tl.program_id(axis=1)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    for k0 in tl.range(0, D, BLOCK_K):
        k_ids = k0 + tl.arange(0, BLOCK_K)
        k_mask = k_ids < D

        a_ptrs = a_ptr + (offs_m[:, None] * D + k_ids[None, :])
        a_tile = tl.load(a_ptrs,
                         mask=(offs_m[:, None] < M) & (k_mask[None, :]),
                         other=0.0).to(tl.float32)

        w_ptrs = w_ptr + (offs_n[None, :] * D + k_ids[:, None])
        w_tile = tl.load(w_ptrs,
                         mask=(offs_n[None, :] < U) & (k_mask[:, None]),
                         other=0.0).to(tl.float32)

        acc += tl.dot(a_tile, w_tile)

    out = tl.sigmoid(acc)
    out_ptrs = out_ptr + (offs_m[:, None] * U + offs_n[None, :])
    tl.store(out_ptrs, out, mask=(offs_m[:, None] < M) & (offs_n[None, :] < U))


@triton.jit
def final_project_kernel_affine(S_ptr, G_ptr, W_ptr, gamma_ptr, beta_ptr, Out_ptr,
                                B, N, U, D,
                                stride_sb, stride_si, stride_sj, stride_su,
                                stride_gb, stride_gi, stride_gj, stride_gu,
                                BLOCK_I: tl.constexpr, BLOCK_J: tl.constexpr, BLOCK_U: tl.constexpr, BLOCK_D: tl.constexpr):
    """
    For each (b, i, j), perform:
      Z = LayerNorm_U(S[b,i,j,:], eps=1e-5) with affine (gamma, beta)
      Z = Z * G[b,i,j,:]
      out[b,i,j,:] = Z @ W^T  where W is [D, U]
    We implement LN over U via a two-pass reduction (mean, var) to fp32 and project to D in tiles.
    """
    pid_i = tl.program_id(axis=0)
    pid_j = tl.program_id(axis=1)
    pid_bd = tl.program_id(axis=2)

    num_d_blocks = tl.cdiv(D, BLOCK_D)
    b = pid_bd // num_d_blocks
    d_block = pid_bd % num_d_blocks
    d_offs = d_block * BLOCK_D + tl.arange(0, BLOCK_D)

    i_offs = pid_i * BLOCK_I + tl.arange(0, BLOCK_I)
    j_offs = pid_j * BLOCK_J + tl.arange(0, BLOCK_J)

    # Pass 1: mean over U of S[b, i, j, :]
    sum_vals = tl.zeros((BLOCK_I, BLOCK_J), dtype=tl.float32)
    for u0 in tl.range(0, U, BLOCK_U):
        u_ids = u0 + tl.arange(0, BLOCK_U)
        u_mask = u_ids < U

        s_ptrs = (S_ptr
                  + b * stride_sb
                  + i_offs[:, None, None] * stride_si
                  + j_offs[None, :, None] * stride_sj
                  + u_ids[None, None, :] * stride_su)
        S_blk = tl.load(s_ptrs,
                        mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (u_mask[None, None, :]),
                        other=0.0).to(tl.float32)
        sum_vals += tl.sum(S_blk, axis=2)

    U_f32 = tl.cast(U, tl.float32)
    mean = sum_vals / U_f32

    # Pass 2: variance via squared deviations
    sum_sqdiff = tl.zeros((BLOCK_I, BLOCK_J), dtype=tl.float32)
    for u0 in tl.range(0, U, BLOCK_U):
        u_ids = u0 + tl.arange(0, BLOCK_U)
        u_mask = u_ids < U

        s_ptrs = (S_ptr
                  + b * stride_sb
                  + i_offs[:, None, None] * stride_si
                  + j_offs[None, :, None] * stride_sj
                  + u_ids[None, None, :] * stride_su)
        S_blk = tl.load(s_ptrs,
                        mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (u_mask[None, None, :]),
                        other=0.0).to(tl.float32)
        diff = S_blk - mean[:, :, None]
        sum_sqdiff += tl.sum(diff * diff, axis=2)

    var = sum_sqdiff / U_f32
    inv_std = 1.0 / tl.sqrt(var + 1.0e-5)

    # Accumulator for final projection to D: [BLOCK_I, BLOCK_J, BLOCK_D]
    acc_out = tl.zeros((BLOCK_I, BLOCK_J, BLOCK_D), dtype=tl.float32)

    # Pass 3: normalize + affine + gate + matmul with W to D
    for u0 in tl.range(0, U, BLOCK_U):
        u_ids = u0 + tl.arange(0, BLOCK_U)
        u_mask = u_ids < U

        s_ptrs = (S_ptr
                  + b * stride_sb
                  + i_offs[:, None, None] * stride_si
                  + j_offs[None, :, None] * stride_sj
                  + u_ids[None, None, :] * stride_su)
        g_ptrs = (G_ptr
                  + b * stride_gb
                  + i_offs[:, None, None] * stride_gi
                  + j_offs[None, :, None] * stride_gj
                  + u_ids[None, None, :] * stride_gu)

        gamma_blk = tl.load(gamma_ptr + u_ids, mask=u_mask, other=1.0).to(tl.float32)
        beta_blk = tl.load(beta_ptr + u_ids, mask=u_mask, other=0.0).to(tl.float32)

        S_blk = tl.load(s_ptrs,
                        mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (u_mask[None, None, :]),
                        other=0.0).to(tl.float32)
        G_blk = tl.load(g_ptrs,
                        mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (u_mask[None, None, :]),
                        other=0.0).to(tl.float32)

        # Normalize + affine + gate
        S_norm = (S_blk - mean[:, :, None]) * inv_std[:, :, None]
        S_affine = S_norm * gamma_blk[None, None, :] + beta_blk[None, None, :]
        Z = S_affine * G_blk  # [BLOCK_I, BLOCK_J, BLOCK_U]

        # W_ptr is [D, U]; submatrix shaped [BLOCK_U, BLOCK_D] storing W^T
        W_sub = tl.load(W_ptr + (d_offs[None, :] * U + u_ids[:, None]),
                        mask=(d_offs[None, :] < D) & (u_mask[:, None]),
                        other=0.0).to(tl.float32)

        # Flatten Z to 2D: [(BLOCK_I*BLOCK_J), BLOCK_U] then matmul
        Z_flat = tl.reshape(Z, (BLOCK_I * BLOCK_J, BLOCK_U))
        tmp = tl.dot(Z_flat, W_sub)  # [BLOCK_I*BLOCK_J, BLOCK_D]
        tmp_3d = tl.reshape(tmp, (BLOCK_I, BLOCK_J, BLOCK_D))
        acc_out += tmp_3d

    out_ptrs = (Out_ptr
                + b * (N * N * D)
                + i_offs[:, None, None] * (N * D)
                + j_offs[None, :, None] * D
                + d_offs[None, None, :])
    tl.store(out_ptrs, acc_out,
             mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (d_offs[None, None, :] < D))


def _check_inputs(x, mask, weights, config):
    assert isinstance(x, torch.Tensor) and x.is_cuda and x.dtype == torch.float32 and x.is_contiguous()
    assert isinstance(mask, torch.Tensor) and mask.is_cuda and mask.dtype == torch.float32 and mask.is_contiguous()
    assert x.ndim == 4
    B, N, N2, D = x.shape
    assert N == N2, "x must be [B, N, N, D]"
    assert mask.shape == (B, N, N)
    assert isinstance(weights, dict)
    assert isinstance(config, dict)
    U = int(config["hidden_dim"])
    assert "dim" in config and int(config["dim"]) == D

    def _ck(name, shape):
        t = weights.get(name, None)
        assert t is not None and isinstance(t, torch.Tensor) and t.is_cuda and t.dtype == torch.float32
        assert t.is_contiguous()
        assert tuple(t.shape) == tuple(shape)

    _ck("norm.weight", (D,))
    _ck("norm.bias", (D,))
    _ck("left_proj.weight", (U, D))
    _ck("right_proj.weight", (U, D))
    _ck("left_gate.weight", (U, D))
    _ck("right_gate.weight", (U, D))
    _ck("out_gate.weight", (U, D))
    _ck("to_out_norm.weight", (U,))
    _ck("to_out_norm.bias", (U,))
    _ck("to_out.weight", (D, U))


def kernel_function(*args):
    # Accept either a single tuple or unpacked args
    if len(args) == 1 and isinstance(args[0], tuple):
        x, mask, weights, config = args[0]
    else:
        x, mask, weights, config = args

    _check_inputs(x, mask, weights, config)
    B, N, _, D = x.shape
    U = int(config["hidden_dim"])
    device = x.device

    # Allocate intermediates
    x_norm = torch.empty_like(x, dtype=torch.float32, device=device)
    M = B * N * N
    x_flat = x.view(M, D)
    x_norm_flat = x_norm.view(M, D)

    # 1) LayerNorm over last dim
    eps = 1e-5
    BLOCK_LN = 128
    grid_ln = (M,)
    layer_norm_lastdim_kernel[grid_ln](
        x_flat, x_norm_flat, weights["norm.weight"], weights["norm.bias"],
        M, D, eps,
        BLOCK_D=BLOCK_LN,
        num_warps=4, num_stages=2
    )

    # 2) left and right features via fused proj+gate (+ mask) over rows [B, N, N].
    left = torch.empty((B, N, N, U), device=device, dtype=torch.float32)
    right = torch.empty((B, N, N, U), device=device, dtype=torch.float32)
    left_flat = left.view(M, U)
    right_flat = right.view(M, U)

    BLOCK_M = 64
    BLOCK_N = 64
    BLOCK_K = 64
    grid_pg = (triton.cdiv(M, BLOCK_M), triton.cdiv(U, BLOCK_N))

    # LEFT on rows (b,i,j) with mask applied per row
    proj_gate_fused_rows_kernel[grid_pg](
        x_norm_flat, weights["left_proj.weight"], weights["left_gate.weight"], mask, left_flat,
        M, N, D, U,
        apply_mask=True,
        BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
        num_warps=4, num_stages=2
    )
    # RIGHT on rows (b,i,j) with its mask applied per row
    proj_gate_fused_rows_kernel[grid_pg](
        x_norm_flat, weights["right_proj.weight"], weights["right_gate.weight"], mask, right_flat,
        M, N, D, U,
        apply_mask=True,
        BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
        num_warps=4, num_stages=2
    )

    # 3) Contraction over k to produce S: [B, N, N, U]
    S = torch.empty((B, N, N, U), device=device, dtype=torch.float32)

    # Strides for [B, N, N, U] contiguous
    stride_b = N * N * U
    stride_i = N * U
    stride_j = U
    stride_u = 1

    BLOCK_I = 32
    BLOCK_J = 32
    BLOCK_KC = 64
    grid_contract = (triton.cdiv(N, BLOCK_I), triton.cdiv(N, BLOCK_J), B * U)
    contract_k_kernel[grid_contract](
        left, right, S,
        B, N, U,
        stride_b, stride_i, stride_j, stride_u,            # L strides: b,i,k(==j),u
        stride_b, stride_i, stride_j, stride_u,            # R strides: b,j(==i),k(==j),u
        stride_b, stride_i, stride_j, stride_u,            # S strides: b,i,j,u
        BLOCK_I=BLOCK_I, BLOCK_J=BLOCK_J, BLOCK_K=BLOCK_KC,
        num_warps=4, num_stages=2
    )

    # 4) Compute out_gate: G = sigmoid(x_norm @ W_out_gate^T)
    G = torch.empty((B, N, N, U), device=device, dtype=torch.float32)
    G_flat = G.view(M, U)
    grid_gate = (triton.cdiv(M, BLOCK_M), triton.cdiv(U, BLOCK_N))
    out_gate_kernel[grid_gate](
        x_norm_flat, weights["out_gate.weight"], G_flat, M, D, U,
        BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
        num_warps=4, num_stages=2
    )

    # 5) Final projection with LayerNorm(U) + affine + gate + matmul to D
    out = torch.empty((B, N, N, D), device=device, dtype=torch.float32)

    BLOCK_I2 = 16
    BLOCK_J2 = 16
    BLOCK_U2 = 32
    BLOCK_D2 = 64
    grid_final = (triton.cdiv(N, BLOCK_I2), triton.cdiv(N, BLOCK_J2), B * triton.cdiv(D, BLOCK_D2))
    final_project_kernel_affine[grid_final](
        S, G, weights["to_out.weight"], weights["to_out_norm.weight"], weights["to_out_norm.bias"], out,
        B, N, U, D,
        stride_b, stride_i, stride_j, stride_u,  # S strides
        stride_b, stride_i, stride_j, stride_u,  # G strides
        BLOCK_I=BLOCK_I2, BLOCK_J=BLOCK_J2, BLOCK_U=BLOCK_U2, BLOCK_D=BLOCK_D2,
        num_warps=4, num_stages=2
    )

    return out

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

- import torch
+ # kernel.py
+ # Triangle Multiplicative Update (Outgoing) -- Forward-only Triton implementation.
+ # All compute is performed inside Triton kernels; Python only validates, allocates, and launches.
+ #
+ # Fused stages implemented across kernels (aggressively fused where feasible):
+ # 1) LayerNorm over D for x[b, i, j, :]
+ # 2) Fused linear(gate+proj) + sigmoid + masking for LEFT and RIGHT paths on rows [B, N, N]
+ # 3) Contraction over k: S[b, i, j, u] = sum_k LEFT[b, i, k, u] * RIGHT[b, j, k, u]
+ # 4) OUT gate per (b, i, j): sigmoid(x_norm[b,i,j,:] @ W_out_gate^T)
+ # 5) Final LayerNorm over U + affine + gate + matmul to D: out[b,i,j,:]
+ #
+ # Fusion boundaries:
+ # - The proj+gate+sigmoid(+mask) is fused per path to minimize memory traffic.
+ # - The N-reduction over k (stage 3) and the U-epilogue + projection (stage 5) use incompatible tilings:
+ # * The contraction reduces along N while keeping (i,j,u), whereas the epilogue normalizes along U and projects to D.
+ # * A single-kernel fusion would cause extreme register/shared-memory pressure and poor occupancy.
+ # Therefore, they are launched as separate kernels for performance and maintainability.
+
import triton
import triton.language as tl
+ import torch
@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
+ def layer_norm_lastdim_kernel(x_ptr, y_ptr, gamma_ptr, beta_ptr,
+ M, D, eps,
+ BLOCK_D: tl.constexpr):
+ # Each program normalizes one row of length D
+ pid = tl.program_id(axis=0)
+ if pid >= M:
+ return
- 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
+ row_x = x_ptr + pid * D
+ row_y = y_ptr + pid * D
+ offs = tl.arange(0, BLOCK_D)
- 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
+ # Pass 1: mean
+ sum_val = tl.zeros((), dtype=tl.float32)
+ for d0 in tl.range(0, D, BLOCK_D):
+ d_ids = d0 + offs
+ mask = d_ids < D
+ x = tl.load(row_x + d_ids, mask=mask, other=0.0).to(tl.float32)
+ sum_val += tl.sum(x, axis=0)
+ d_float = tl.cast(D, tl.float32)
+ mean = sum_val / d_float
- out_ptr = Y_ptr + row_idx * stride_x
- tl.store(out_ptr + col_offsets, y, mask=mask)
+ # Pass 2: variance via squared deviations
+ sum_sqdiff = tl.zeros((), dtype=tl.float32)
+ for d0 in tl.range(0, D, BLOCK_D):
+ d_ids = d0 + offs
+ mask = d_ids < D
+ x = tl.load(row_x + d_ids, mask=mask, other=0.0).to(tl.float32)
+ diff = x - mean
+ sum_sqdiff += tl.sum(diff * diff, axis=0)
+ var = sum_sqdiff / d_float
+ inv_std = 1.0 / tl.sqrt(var + eps)
+ # Pass 3: normalize, scale, shift
+ for d0 in tl.range(0, D, BLOCK_D):
+ d_ids = d0 + offs
+ mask = d_ids < D
+ x = tl.load(row_x + d_ids, mask=mask, other=0.0).to(tl.float32)
+ g = tl.load(gamma_ptr + d_ids, mask=mask, other=1.0).to(tl.float32)
+ b = tl.load(beta_ptr + d_ids, mask=mask, other=0.0).to(tl.float32)
+ y = (x - mean) * inv_std
+ y = y * g + b
+ tl.store(row_y + d_ids, 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
+ def proj_gate_fused_rows_kernel(a_ptr, wproj_ptr, wgate_ptr, mask_ptr, out_ptr,
+ M, N, D, U,
+ apply_mask: tl.constexpr,
+ BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):
+ """
+ Fused projection and gate over rows of a 2D tensor [M, D], where M = B * N * N and each row maps to (b, i, j).
+ Produces: out[row, u] = (a_row @ wproj.T) * sigmoid(a_row @ wgate.T) * (optional mask[b,i,j]).
+ """
+ pid_m = tl.program_id(axis=0)
+ pid_n = tl.program_id(axis=1)
- 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
+ offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
+ offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
- for h_start in range(0, H, BLOCK_H):
- h_offs = h_start + tl.arange(0, BLOCK_H)
- h_mask = h_offs < H
+ acc_proj = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
+ acc_gate = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
- 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)
+ # Loop over K = D
+ for k0 in tl.range(0, D, BLOCK_K):
+ k_ids = k0 + tl.arange(0, BLOCK_K)
+ k_mask = k_ids < D
- 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)
+ # A tile: [BLOCK_M, BLOCK_K]
+ a_ptrs = a_ptr + (offs_m[:, None] * D + k_ids[None, :])
+ a_tile = tl.load(
+ a_ptrs,
+ mask=(offs_m[:, None] < M) & (k_mask[None, :]),
+ other=0.0
+ ).to(tl.float32)
- weight_offsets = h_offs[:, None] * C + c_offs[None, :]
- combined_mask = h_mask[:, None] & c_mask[None, :]
+ # Weight tiles: [BLOCK_K, BLOCK_N] (weights are [U, D])
+ wproj_ptrs = wproj_ptr + (offs_n[None, :] * D + k_ids[:, None])
+ wgate_ptrs = wgate_ptr + (offs_n[None, :] * D + k_ids[:, 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)
+ wproj_tile = tl.load(
+ wproj_ptrs,
+ mask=(offs_n[None, :] < U) & (k_mask[:, None]),
+ other=0.0
+ ).to(tl.float32)
+ wgate_tile = tl.load(
+ wgate_ptrs,
+ mask=(offs_n[None, :] < U) & (k_mask[:, None]),
+ 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)
+ acc_proj += tl.dot(a_tile, wproj_tile) # [M_blk, U_blk]
+ acc_gate += tl.dot(a_tile, wgate_tile)
- left_gate_sig = tl.sigmoid(left_gate_acc)
- right_gate_sig = tl.sigmoid(right_gate_acc)
- out_gate_sig = tl.sigmoid(out_gate_acc)
+ gate = tl.sigmoid(acc_gate)
+ out = acc_proj * gate
- left_result = left_proj_acc * mask_val * left_gate_sig
- right_result = right_proj_acc * mask_val * right_gate_sig
+ if apply_mask:
+ NN = N * N
+ b_m = offs_m // NN
+ rem = offs_m - b_m * NN
+ i_m = rem // N
+ j_m = rem - i_m * N
+ m_ptrs = mask_ptr + b_m * NN + i_m * N + j_m
+ mvals = tl.load(m_ptrs, mask=offs_m < M, other=0.0).to(tl.float32) # [M_blk]
+ out = out * mvals[:, None]
- 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)
+ out_ptrs = out_ptr + (offs_m[:, None] * U + offs_n[None, :])
+ mask_out = (offs_m[:, None] < M) & (offs_n[None, :] < U)
+ tl.store(out_ptrs, out, mask=mask_out)
@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)
+ def contract_k_kernel(L_ptr, R_ptr, S_ptr,
+ B, N, U,
+ stride_lb, stride_li, stride_lk, stride_lu,
+ stride_rb, stride_rj, stride_rk, stride_ru,
+ stride_sb, stride_si, stride_sj, stride_su,
+ BLOCK_I: tl.constexpr, BLOCK_J: tl.constexpr, BLOCK_K: tl.constexpr):
+ # 3D launch: I, J, (B*U)
+ pid_i = tl.program_id(axis=0)
+ pid_j = tl.program_id(axis=1)
+ pid_bu = tl.program_id(axis=2)
- num_tiles_j = tl.cdiv(N, BLOCK_J)
- pid_i = pid_ij // num_tiles_j
- pid_j = pid_ij % num_tiles_j
+ b = pid_bu // U
+ u = pid_bu % U
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
+ for k0 in tl.range(0, N, BLOCK_K):
+ offs_k = k0 + tl.arange(0, BLOCK_K)
+ k_mask = 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)
+ # L[b, i, k, u] -> [BLOCK_I, BLOCK_K]
+ l_ptrs = (L_ptr
+ + b * stride_lb
+ + offs_i[:, None] * stride_li
+ + offs_k[None, :] * stride_lk
+ + u * stride_lu)
+ L_tile = tl.load(l_ptrs, mask=(offs_i[:, None] < N) & (k_mask[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)
+ # R[b, j, k, u] -> [BLOCK_K, BLOCK_J]
+ r_ptrs = (R_ptr
+ + b * stride_rb
+ + offs_j[None, :] * stride_rj
+ + offs_k[:, None] * stride_rk
+ + u * stride_ru)
+ R_tile_KJ = tl.load(r_ptrs, mask=(k_mask[:, None]) & (offs_j[None, :] < N), other=0.0).to(tl.float32)
- acc += tl.dot(x0, tl.trans(x1))
+ acc += tl.dot(L_tile, R_tile_KJ)
- 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, :])
+ s_ptrs = (S_ptr
+ + b * stride_sb
+ + offs_i[:, None] * stride_si
+ + offs_j[None, :] * stride_sj
+ + u * stride_su)
+ tl.store(s_ptrs, acc, mask=(offs_i[:, None] < N) & (offs_j[None, :] < N))
@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
+ def out_gate_kernel(a_ptr, w_ptr, out_ptr, M, D, U,
+ BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):
+ # Compute: sigmoid(a @ w^T) for M rows, projecting D->U
+ pid_m = tl.program_id(axis=0)
+ pid_n = tl.program_id(axis=1)
- offs_h = tl.arange(0, BLOCK_H)
- mask_h = offs_h < H
+ offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
+ offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
- 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
+ acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
- 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
+ for k0 in tl.range(0, D, BLOCK_K):
+ k_ids = k0 + tl.arange(0, BLOCK_K)
+ k_mask = k_ids < D
- 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)
+ a_ptrs = a_ptr + (offs_m[:, None] * D + k_ids[None, :])
+ a_tile = tl.load(a_ptrs,
+ mask=(offs_m[:, None] < M) & (k_mask[None, :]),
+ other=0.0).to(tl.float32)
- 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)
+ w_ptrs = w_ptr + (offs_n[None, :] * D + k_ids[:, None])
+ w_tile = tl.load(w_ptrs,
+ mask=(offs_n[None, :] < U) & (k_mask[:, None]),
+ other=0.0).to(tl.float32)
+ acc += tl.dot(a_tile, w_tile)
- def kernel_function(input_tensor, mask, weights, config):
- B, N, _, C = input_tensor.shape
- H = config["hidden_dim"]
+ out = tl.sigmoid(acc)
+ out_ptrs = out_ptr + (offs_m[:, None] * U + offs_n[None, :])
+ tl.store(out_ptrs, out, mask=(offs_m[:, None] < M) & (offs_n[None, :] < U))
- # 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')
+ @triton.jit
+ def final_project_kernel_affine(S_ptr, G_ptr, W_ptr, gamma_ptr, beta_ptr, Out_ptr,
+ B, N, U, D,
+ stride_sb, stride_si, stride_sj, stride_su,
+ stride_gb, stride_gi, stride_gj, stride_gu,
+ BLOCK_I: tl.constexpr, BLOCK_J: tl.constexpr, BLOCK_U: tl.constexpr, BLOCK_D: tl.constexpr):
+ """
+ For each (b, i, j), perform:
+ Z = LayerNorm_U(S[b,i,j,:], eps=1e-5) with affine (gamma, beta)
+ Z = Z * G[b,i,j,:]
+ out[b,i,j,:] = Z @ W^T where W is [D, U]
+ We implement LN over U via a two-pass reduction (mean, var) to fp32 and project to D in tiles.
+ """
+ pid_i = tl.program_id(axis=0)
+ pid_j = tl.program_id(axis=1)
+ pid_bd = tl.program_id(axis=2)
- 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
- )
+ num_d_blocks = tl.cdiv(D, BLOCK_D)
+ b = pid_bd // num_d_blocks
+ d_block = pid_bd % num_d_blocks
+ d_offs = d_block * BLOCK_D + tl.arange(0, 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
- )
+ i_offs = pid_i * BLOCK_I + tl.arange(0, BLOCK_I)
+ j_offs = pid_j * BLOCK_J + tl.arange(0, BLOCK_J)
- # 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
- )
+ # Pass 1: mean over U of S[b, i, j, :]
+ sum_vals = tl.zeros((BLOCK_I, BLOCK_J), dtype=tl.float32)
+ for u0 in tl.range(0, U, BLOCK_U):
+ u_ids = u0 + tl.arange(0, BLOCK_U)
+ u_mask = u_ids < U
- return output
+ s_ptrs = (S_ptr
+ + b * stride_sb
+ + i_offs[:, None, None] * stride_si
+ + j_offs[None, :, None] * stride_sj
+ + u_ids[None, None, :] * stride_su)
+ S_blk = tl.load(s_ptrs,
+ mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (u_mask[None, None, :]),
+ other=0.0).to(tl.float32)
+ sum_vals += tl.sum(S_blk, axis=2)
+ U_f32 = tl.cast(U, tl.float32)
+ mean = sum_vals / U_f32
- def test_kernel():
- torch.manual_seed(42)
- B, N, C, H = 1, 32, 128, 128
+ # Pass 2: variance via squared deviations
+ sum_sqdiff = tl.zeros((BLOCK_I, BLOCK_J), dtype=tl.float32)
+ for u0 in tl.range(0, U, BLOCK_U):
+ u_ids = u0 + tl.arange(0, BLOCK_U)
+ u_mask = u_ids < U
- input_tensor = torch.randn(B, N, N, C, device='cuda', dtype=torch.float32)
- mask = torch.ones(B, N, N, device='cuda', dtype=torch.float32)
+ s_ptrs = (S_ptr
+ + b * stride_sb
+ + i_offs[:, None, None] * stride_si
+ + j_offs[None, :, None] * stride_sj
+ + u_ids[None, None, :] * stride_su)
+ S_blk = tl.load(s_ptrs,
+ mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (u_mask[None, None, :]),
+ other=0.0).to(tl.float32)
+ diff = S_blk - mean[:, :, None]
+ sum_sqdiff += tl.sum(diff * diff, axis=2)
- 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}
+ var = sum_sqdiff / U_f32
+ inv_std = 1.0 / tl.sqrt(var + 1.0e-5)
- out_triton = kernel_function(input_tensor, mask, weights, config)
+ # Accumulator for final projection to D: [BLOCK_I, BLOCK_J, BLOCK_D]
+ acc_out = tl.zeros((BLOCK_I, BLOCK_J, BLOCK_D), dtype=tl.float32)
- # 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
+ # Pass 3: normalize + affine + gate + matmul with W to D
+ for u0 in tl.range(0, U, BLOCK_U):
+ u_ids = u0 + tl.arange(0, BLOCK_U)
+ u_mask = u_ids < U
- if torch.allclose(out_triton, ref, rtol=2e-2, atol=2e-2):
- print("PASS")
+ s_ptrs = (S_ptr
+ + b * stride_sb
+ + i_offs[:, None, None] * stride_si
+ + j_offs[None, :, None] * stride_sj
+ + u_ids[None, None, :] * stride_su)
+ g_ptrs = (G_ptr
+ + b * stride_gb
+ + i_offs[:, None, None] * stride_gi
+ + j_offs[None, :, None] * stride_gj
+ + u_ids[None, None, :] * stride_gu)
+
+ gamma_blk = tl.load(gamma_ptr + u_ids, mask=u_mask, other=1.0).to(tl.float32)
+ beta_blk = tl.load(beta_ptr + u_ids, mask=u_mask, other=0.0).to(tl.float32)
+
+ S_blk = tl.load(s_ptrs,
+ mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (u_mask[None, None, :]),
+ other=0.0).to(tl.float32)
+ G_blk = tl.load(g_ptrs,
+ mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (u_mask[None, None, :]),
+ other=0.0).to(tl.float32)
+
+ # Normalize + affine + gate
+ S_norm = (S_blk - mean[:, :, None]) * inv_std[:, :, None]
+ S_affine = S_norm * gamma_blk[None, None, :] + beta_blk[None, None, :]
+ Z = S_affine * G_blk # [BLOCK_I, BLOCK_J, BLOCK_U]
+
+ # W_ptr is [D, U]; submatrix shaped [BLOCK_U, BLOCK_D] storing W^T
+ W_sub = tl.load(W_ptr + (d_offs[None, :] * U + u_ids[:, None]),
+ mask=(d_offs[None, :] < D) & (u_mask[:, None]),
+ other=0.0).to(tl.float32)
+
+ # Flatten Z to 2D: [(BLOCK_I*BLOCK_J), BLOCK_U] then matmul
+ Z_flat = tl.reshape(Z, (BLOCK_I * BLOCK_J, BLOCK_U))
+ tmp = tl.dot(Z_flat, W_sub) # [BLOCK_I*BLOCK_J, BLOCK_D]
+ tmp_3d = tl.reshape(tmp, (BLOCK_I, BLOCK_J, BLOCK_D))
+ acc_out += tmp_3d
+
+ out_ptrs = (Out_ptr
+ + b * (N * N * D)
+ + i_offs[:, None, None] * (N * D)
+ + j_offs[None, :, None] * D
+ + d_offs[None, None, :])
+ tl.store(out_ptrs, acc_out,
+ mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (d_offs[None, None, :] < D))
+
+
+ def _check_inputs(x, mask, weights, config):
+ assert isinstance(x, torch.Tensor) and x.is_cuda and x.dtype == torch.float32 and x.is_contiguous()
+ assert isinstance(mask, torch.Tensor) and mask.is_cuda and mask.dtype == torch.float32 and mask.is_contiguous()
+ assert x.ndim == 4
+ B, N, N2, D = x.shape
+ assert N == N2, "x must be [B, N, N, D]"
+ assert mask.shape == (B, N, N)
+ assert isinstance(weights, dict)
+ assert isinstance(config, dict)
+ U = int(config["hidden_dim"])
+ assert "dim" in config and int(config["dim"]) == D
+
+ def _ck(name, shape):
+ t = weights.get(name, None)
+ assert t is not None and isinstance(t, torch.Tensor) and t.is_cuda and t.dtype == torch.float32
+ assert t.is_contiguous()
+ assert tuple(t.shape) == tuple(shape)
+
+ _ck("norm.weight", (D,))
+ _ck("norm.bias", (D,))
+ _ck("left_proj.weight", (U, D))
+ _ck("right_proj.weight", (U, D))
+ _ck("left_gate.weight", (U, D))
+ _ck("right_gate.weight", (U, D))
+ _ck("out_gate.weight", (U, D))
+ _ck("to_out_norm.weight", (U,))
+ _ck("to_out_norm.bias", (U,))
+ _ck("to_out.weight", (D, U))
+
+
+ def kernel_function(*args):
+ # Accept either a single tuple or unpacked args
+ if len(args) == 1 and isinstance(args[0], tuple):
+ x, mask, weights, config = args[0]
else:
- print(f"FAIL: max diff = {(out_triton - ref).abs().max().item()}")
+ x, mask, weights, config = args
+ _check_inputs(x, mask, weights, config)
+ B, N, _, D = x.shape
+ U = int(config["hidden_dim"])
+ device = x.device
- if __name__ == "__main__":
- test_kernel()
+ # Allocate intermediates
+ x_norm = torch.empty_like(x, dtype=torch.float32, device=device)
+ M = B * N * N
+ x_flat = x.view(M, D)
+ x_norm_flat = x_norm.view(M, D)
+ # 1) LayerNorm over last dim
+ eps = 1e-5
+ BLOCK_LN = 128
+ grid_ln = (M,)
+ layer_norm_lastdim_kernel[grid_ln](
+ x_flat, x_norm_flat, weights["norm.weight"], weights["norm.bias"],
+ M, D, eps,
+ BLOCK_D=BLOCK_LN,
+ num_warps=4, num_stages=2
+ )
+ # 2) left and right features via fused proj+gate (+ mask) over rows [B, N, N].
+ left = torch.empty((B, N, N, U), device=device, dtype=torch.float32)
+ right = torch.empty((B, N, N, U), device=device, dtype=torch.float32)
+ left_flat = left.view(M, U)
+ right_flat = right.view(M, U)
+
+ BLOCK_M = 64
+ BLOCK_N = 64
+ BLOCK_K = 64
+ grid_pg = (triton.cdiv(M, BLOCK_M), triton.cdiv(U, BLOCK_N))
+
+ # LEFT on rows (b,i,j) with mask applied per row
+ proj_gate_fused_rows_kernel[grid_pg](
+ x_norm_flat, weights["left_proj.weight"], weights["left_gate.weight"], mask, left_flat,
+ M, N, D, U,
+ apply_mask=True,
+ BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
+ num_warps=4, num_stages=2
+ )
+ # RIGHT on rows (b,i,j) with its mask applied per row
+ proj_gate_fused_rows_kernel[grid_pg](
+ x_norm_flat, weights["right_proj.weight"], weights["right_gate.weight"], mask, right_flat,
+ M, N, D, U,
+ apply_mask=True,
+ BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
+ num_warps=4, num_stages=2
+ )
+
+ # 3) Contraction over k to produce S: [B, N, N, U]
+ S = torch.empty((B, N, N, U), device=device, dtype=torch.float32)
+
+ # Strides for [B, N, N, U] contiguous
+ stride_b = N * N * U
+ stride_i = N * U
+ stride_j = U
+ stride_u = 1
+
+ BLOCK_I = 32
+ BLOCK_J = 32
+ BLOCK_KC = 64
+ grid_contract = (triton.cdiv(N, BLOCK_I), triton.cdiv(N, BLOCK_J), B * U)
+ contract_k_kernel[grid_contract](
+ left, right, S,
+ B, N, U,
+ stride_b, stride_i, stride_j, stride_u, # L strides: b,i,k(==j),u
+ stride_b, stride_i, stride_j, stride_u, # R strides: b,j(==i),k(==j),u
+ stride_b, stride_i, stride_j, stride_u, # S strides: b,i,j,u
+ BLOCK_I=BLOCK_I, BLOCK_J=BLOCK_J, BLOCK_K=BLOCK_KC,
+ num_warps=4, num_stages=2
+ )
+
+ # 4) Compute out_gate: G = sigmoid(x_norm @ W_out_gate^T)
+ G = torch.empty((B, N, N, U), device=device, dtype=torch.float32)
+ G_flat = G.view(M, U)
+ grid_gate = (triton.cdiv(M, BLOCK_M), triton.cdiv(U, BLOCK_N))
+ out_gate_kernel[grid_gate](
+ x_norm_flat, weights["out_gate.weight"], G_flat, M, D, U,
+ BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
+ num_warps=4, num_stages=2
+ )
+
+ # 5) Final projection with LayerNorm(U) + affine + gate + matmul to D
+ out = torch.empty((B, N, N, D), device=device, dtype=torch.float32)
+
+ BLOCK_I2 = 16
+ BLOCK_J2 = 16
+ BLOCK_U2 = 32
+ BLOCK_D2 = 64
+ grid_final = (triton.cdiv(N, BLOCK_I2), triton.cdiv(N, BLOCK_J2), B * triton.cdiv(D, BLOCK_D2))
+ final_project_kernel_affine[grid_final](
+ S, G, weights["to_out.weight"], weights["to_out_norm.weight"], weights["to_out_norm.bias"], out,
+ B, N, U, D,
+ stride_b, stride_i, stride_j, stride_u, # S strides
+ stride_b, stride_i, stride_j, stride_u, # G strides
+ BLOCK_I=BLOCK_I2, BLOCK_J=BLOCK_J2, BLOCK_U=BLOCK_U2, BLOCK_D=BLOCK_D2,
+ num_warps=4, num_stages=2
+ )
+
+ return out
+
def custom_kernel(input):
return kernel_function(*input)
scrolls · 684 diff lines total

Best evidence level for this revision: reported

JSON