Skip to content
KernelIndex
Search⌘K

submission 413027

Emmett Bicker · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

best_result_H100.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-413027?include=source"
interfacepython
Compatibility
measured onNVIDIA A100
declared hardwareNVIDIA A100
architecturessm_80
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA A100
2.38ms
#7 of 69
2026-01-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6a7fc9059d21a365654422b9a17979ea5d90162f03f515fcf9fc3ea266abd545
license declaredunknown
license concludedunknown
authorsEmmett Bicker
imported2026-08-15

Techniques

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

autotune@triton.autotune(
mmalp += tl.dot(x16, tl.load(w_lp + w_off, mask=wm, other=0.0).to(tl.float16))
num-warps = 4triton.Config({'BM': 64, 'BD': 64, 'BH': 64}, num_warps=4, num_stages=3),
persistent-kernel- torch.baddbmm with a persistent output buffer to avoid per-call allocations.
stages = 3triton.Config({'BM': 64, 'BD': 64, 'BH': 64}, num_warps=4, num_stages=3),
tile-m = 256BM=256, BD=64,

Kernel source

best_result_H100.py311 lines
import torch
from torch import nn
import triton
import triton.language as tl
import math

# Fused Triton head: LayerNorm(x) + 5 pointwise projections + sigmoid gates + optional mask on flattened [B*N*N, dim],
# directly pack L/R into [B*hidden, N*N] layout for fast tensor-core torch.bmm, store gates [B*N*N, hidden].
# Triton tail: unpack [B*hidden, N*N] back to [B*N*N, hidden], LayerNorm + out_gate mul + final proj to [B*N*N, dim].
def _get_w16_T(weights, name, ref):
    key = name + "_T_fp16"
    w = weights.get(key, None)
    if w is None or w.device != ref.device:
        w0 = weights[name]
        if w0.dtype != torch.float16 or w0.device != ref.device:
            w0 = w0.to(device=ref.device, dtype=torch.float16)
        w = w0.t().contiguous()
        weights[key] = w
    return w

def _get_f16(weights, name, ref):
    # Cache LN vectors as fp16 to cut bandwidth inside kernels.
    key = name + "_fp16"
    v = weights.get(key, None)
    if v is None or v.device != ref.device:
        v0 = weights[name]
        if v0.dtype != torch.float16 or v0.device != ref.device:
            v0 = v0.to(device=ref.device, dtype=torch.float16)
        v = v0.contiguous()
        weights[key] = v
    return v

@triton.jit
def _ln_stats_kernel(
    x_ptr, mean_ptr, rstd_ptr,
    M: tl.constexpr, D: tl.constexpr,
    s_xm: tl.constexpr, s_xd: tl.constexpr,
    BM: tl.constexpr, BD: tl.constexpr,
):
    # Compute LayerNorm statistics for each row of x2d [M, D] once.
    pid = tl.program_id(0)
    offs_m = pid * BM + tl.arange(0, BM)
    m_m = offs_m < M
    s1 = tl.zeros((BM,), tl.float32)
    s2 = tl.zeros((BM,), tl.float32)
    for kd in range(0, D, BD):
        offs_d = kd + tl.arange(0, BD)
        m_d = offs_d < D
        x = tl.load(
            x_ptr + offs_m[:, None] * s_xm + offs_d[None, :] * s_xd,
            mask=m_m[:, None] & m_d[None, :],
            other=0.0,
        ).to(tl.float32)
        s1 += tl.sum(x, axis=1)
        s2 += tl.sum(x * x, axis=1)
    mean = s1 / D
    var = s2 / D - mean * mean
    rstd = tl.math.rsqrt(var + 1e-5)
    tl.store(mean_ptr + offs_m, mean, mask=m_m)
    tl.store(rstd_ptr + offs_m, rstd, mask=m_m)

@triton.autotune(
    configs=[
        triton.Config({'BM': 64, 'BD': 64, 'BH': 64}, num_warps=4, num_stages=3),
        triton.Config({'BM': 128, 'BD': 32, 'BH': 32}, num_warps=4, num_stages=3),
        triton.Config({'BM': 64, 'BD': 32, 'BH': 64}, num_warps=4, num_stages=3),
    ],
    key=['M', 'D', 'H'],
)
@triton.jit
def _head_fused_kernel(
    x_ptr, mask_ptr, mean_ptr, rstd_ptr,
    w_lp, w_rp, w_lg, w_rg, w_og,
    ln_w, ln_b,
    l_out_ptr, r_out_ptr, g_out_ptr,
    M: tl.constexpr, D: tl.constexpr, H: tl.constexpr, NN: tl.constexpr,
    s_xm: tl.constexpr, s_xd: tl.constexpr,
    s_wk: tl.constexpr, s_wh: tl.constexpr,
    HAS_MASK: tl.constexpr,
    BM: tl.constexpr, BD: tl.constexpr, BH: tl.constexpr,
):
    # Key change: do NOT recompute LN stats per pid_h tile (that was a massive redundancy).
    pid_h = tl.program_id(0)
    pid_m = tl.program_id(1)

    offs_m = pid_m * BM + tl.arange(0, BM)
    offs_h = pid_h * BH + tl.arange(0, BH)
    m_m = offs_m < M
    m_h = offs_h < H

    mean = tl.load(mean_ptr + offs_m, mask=m_m, other=0.0).to(tl.float32)
    rstd = tl.load(rstd_ptr + offs_m, mask=m_m, other=0.0).to(tl.float32)

    lp = tl.zeros((BM, BH), tl.float32)
    rp = tl.zeros((BM, BH), tl.float32)
    lg = tl.zeros((BM, BH), tl.float32)
    rg = tl.zeros((BM, BH), tl.float32)
    og = tl.zeros((BM, BH), tl.float32)

    for kd in range(0, D, BD):
        offs_d = kd + tl.arange(0, BD)
        m_d = offs_d < D

        x = tl.load(
            x_ptr + offs_m[:, None] * s_xm + offs_d[None, :] * s_xd,
            mask=m_m[:, None] & m_d[None, :],
            other=0.0,
        ).to(tl.float32)

        # LN affine in fp16 (stats kept in fp32)
        w = tl.load(ln_w + offs_d, mask=m_d, other=0.0).to(tl.float16)
        b = tl.load(ln_b + offs_d, mask=m_d, other=0.0).to(tl.float16)
        x16 = ((x - mean[:, None]) * rstd[:, None]).to(tl.float16)
        x16 = x16 * w[None, :] + b[None, :]

        w_off = offs_d[:, None] * s_wk + offs_h[None, :] * s_wh
        wm = m_d[:, None] & m_h[None, :]
        lp += tl.dot(x16, tl.load(w_lp + w_off, mask=wm, other=0.0).to(tl.float16))
        rp += tl.dot(x16, tl.load(w_rp + w_off, mask=wm, other=0.0).to(tl.float16))
        lg += tl.dot(x16, tl.load(w_lg + w_off, mask=wm, other=0.0).to(tl.float16))
        rg += tl.dot(x16, tl.load(w_rg + w_off, mask=wm, other=0.0).to(tl.float16))
        og += tl.dot(x16, tl.load(w_og + w_off, mask=wm, other=0.0).to(tl.float16))

    l = lp * tl.sigmoid(lg)
    r = rp * tl.sigmoid(rg)
    g = tl.sigmoid(og)

    if HAS_MASK:
        mm = tl.load(mask_ptr + offs_m, mask=m_m, other=0.0).to(tl.float32)
        l *= mm[:, None]
        r *= mm[:, None]

    st = m_m[:, None] & m_h[None, :]
    tl.store(g_out_ptr + offs_m[:, None] * H + offs_h[None, :], g.to(tl.float16), mask=st)

    # Pack L/R contiguously as [B*H, N*N]. Let cuBLAS handle transpose via op(B)=T.
    b_idx = offs_m // NN
    rem = offs_m - b_idx * NN
    addr = (b_idx[:, None] * H + offs_h[None, :]) * NN + rem[:, None]
    tl.store(l_out_ptr + addr, l.to(tl.float16), mask=st)
    tl.store(r_out_ptr + addr, r.to(tl.float16), mask=st)

# NOTE: repack kernel removed from the hotpath. Repacking [B*H, N*N] -> [M, H]
# is a full extra bandwidth pass and was a major regression at N=768/1024.
# Keep a tiny stub to avoid editing more code; it is no longer invoked.
@triton.jit
def _repack_bh_nn_to_m_h_kernel(
    inp_ptr, out_ptr,
    M: tl.constexpr, H: tl.constexpr, NN: tl.constexpr,
    BM: tl.constexpr, BH: tl.constexpr,
):
    return


@triton.jit
def _tail_fused_kernel(
    bmm_ptr, g_ptr,
    w_out, ln_w, ln_b,
    out_ptr,
    M: tl.constexpr, H: tl.constexpr, D: tl.constexpr, NN: tl.constexpr,
    s_wh: tl.constexpr, s_wd: tl.constexpr,
    BM: tl.constexpr, BD: tl.constexpr, BH: tl.constexpr,
):
    # Read directly from [B*H, N*N] (flattened) with address math.
    # Avoids materializing/repacking [M, H].
    pid = tl.program_id(0)
    offs_m = pid * BM + tl.arange(0, BM)
    m_m = offs_m < M

    offs_h = tl.arange(0, BH)
    m_h = offs_h < H

    b_idx = offs_m // NN
    rem = offs_m - b_idx * NN
    addr = (b_idx[:, None] * H + offs_h[None, :]) * NN + rem[:, None]

    v = tl.load(
        bmm_ptr + addr,
        mask=m_m[:, None] & m_h[None, :],
        other=0.0,
    ).to(tl.float32)

    g = tl.load(
        g_ptr + offs_m[:, None] * H + offs_h[None, :],
        mask=m_m[:, None] & m_h[None, :],
        other=0.0,
    ).to(tl.float16)

    mean = tl.sum(v, axis=1) / H
    var = tl.sum(v * v, axis=1) / H - mean * mean
    rstd = tl.math.rsqrt(var + 1e-5)

    # fp16 LN affine + fp16 gate for tensorcore dot
    w = tl.load(ln_w + offs_h, mask=m_h, other=0.0).to(tl.float16)
    b = tl.load(ln_b + offs_h, mask=m_h, other=0.0).to(tl.float16)

    v16 = ((v - mean[:, None]) * rstd[:, None]).to(tl.float16)
    v16 = (v16 * w[None, :] + b[None, :]) * g

    for kd in range(0, D, BD):
        offs_d = kd + tl.arange(0, BD)
        m_d = offs_d < D
        w_tile = tl.load(
            w_out + offs_h[:, None] * s_wh + offs_d[None, :] * s_wd,
            mask=m_h[:, None] & m_d[None, :],
            other=0.0,
        ).to(tl.float16)
        o = tl.dot(v16, w_tile)
        tl.store(
            out_ptr + offs_m[:, None] * D + offs_d[None, :],
            o.to(tl.float32),
            mask=m_m[:, None] & m_d[None, :],
        )


# NOTE: TriMul nn.Module removed (not used by the evaluator); keeping only custom_kernel reduces code size/compile time.


def custom_kernel(data):
    """
    Performance-oriented TriMul(outgoing) forward:
      - Triton head: LayerNorm(x) + 5 projections + sigmoid gates (+ optional mask),
        and directly pack L/R into [B*H, N*N] for tensor-core bmm; store out_gate.
        This avoids materializing the massive [M,5H] 'proj' tensor (which can exceed 1GB).
      - torch.baddbmm with a persistent output buffer to avoid per-call allocations.
      - Triton tail: LayerNorm(out) + out_gate + final projection to dim (fp32 output).
    """
    x, mask, weights, config = data
    D, H = config["dim"], config["hidden_dim"]
    B, N, _, _ = x.shape
    NN = N * N
    M = B * NN

    # flatten x to [M, D]
    x2d = x.reshape(M, D)
    mask_flat = mask.reshape(M) if mask is not None else None

    # cache fp16 transposed weights for tl.dot (shape [D,H] / [H,D])
    w_lp = _get_w16_T(weights, "left_proj.weight", x)
    w_rp = _get_w16_T(weights, "right_proj.weight", x)
    w_lg = _get_w16_T(weights, "left_gate.weight", x)
    w_rg = _get_w16_T(weights, "right_gate.weight", x)
    w_og = _get_w16_T(weights, "out_gate.weight", x)
    w_to = _get_w16_T(weights, "to_out.weight", x)  # [H, D]

    # Reuse large buffers (critical for N=768/1024). Allocations here dominate otherwise.
    scratch = weights.setdefault("_triumul_scratch", {})
    skey = (B, N, D, H, x.device)
    buf = scratch.get(skey, None)
    if buf is None:
        buf = {
            "l_bmm": torch.empty((B * H, NN), device=x.device, dtype=torch.float16),
            "r_bmm": torch.empty((B * H, NN), device=x.device, dtype=torch.float16),
            "g_out": torch.empty((M, H), device=x.device, dtype=torch.float16),
            "out_bmm": torch.empty((B * H, N, N), device=x.device, dtype=torch.float16),
            "out2d": torch.empty((M, D), device=x.device, dtype=torch.float32),
            # LN stats scratch (computed once per row, reused across hidden tiles)
            "mean": torch.empty((M,), device=x.device, dtype=torch.float32),
            "rstd": torch.empty((M,), device=x.device, dtype=torch.float32),
        }
        scratch[skey] = buf

    l_bmm = buf["l_bmm"]
    r_bmm = buf["r_bmm"]
    g_out = buf["g_out"]

    # 1) LN stats once per row (avoid recomputing per hidden-tile pid_h)
    mean = buf["mean"]
    rstd = buf["rstd"]
    _ln_stats_kernel[(triton.cdiv(M, 256),)](
        x2d, mean, rstd,
        M=M, D=D,
        s_xm=x2d.stride(0), s_xd=x2d.stride(1),
        BM=256, BD=64,
        num_warps=4,
    )

    # 2) head: projections + gates (+ optional mask) + pack L/R for BMM
    grid_head = lambda META: (triton.cdiv(H, META["BH"]), triton.cdiv(M, META["BM"]))
    _head_fused_kernel[grid_head](
        x2d, mask_flat if mask_flat is not None else x2d, mean, rstd,
        w_lp, w_rp, w_lg, w_rg, w_og,
        _get_f16(weights, "norm.weight", x), _get_f16(weights, "norm.bias", x),
        l_bmm, r_bmm, g_out,
        M=M, D=D, H=H, NN=NN,
        s_xm=x2d.stride(0), s_xd=x2d.stride(1),
        s_wk=w_lp.stride(0), s_wh=w_lp.stride(1),
        HAS_MASK=(mask_flat is not None),
    )

    # 3) tensor-core contraction; let cuBLAS handle transpose via op(B)=T (no custom repacking)
    out_bmm = buf["out_bmm"]
    A = l_bmm.view(B * H, N, N)
    Bt = r_bmm.view(B * H, N, N).transpose(1, 2)
    torch.baddbmm(out_bmm, A, Bt, beta=0.0, alpha=1.0, out=out_bmm)

    # 4) tail: read directly from [B*H, NN], LN + gate + final projection => fp32 [M, D]
    out2d = buf["out2d"]
    BD_TAIL = 128 if D == 128 else 64
    grid_tail = (triton.cdiv(M, 64),)
    _tail_fused_kernel[grid_tail](
        out_bmm.view(B * H, NN), g_out,
        w_to, _get_f16(weights, "to_out_norm.weight", x), _get_f16(weights, "to_out_norm.bias", x),
        out2d,
        M=M, H=H, D=D, NN=NN,
        s_wh=w_to.stride(0), s_wd=w_to.stride(1),
        BM=64, BD=BD_TAIL, BH=128,
        num_warps=4,
    )
    return out2d.view(B, N, N, D)
scrolls · 311 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 412370.

⋯ 17 unchanged lines
weights[key] = w
return w
+ def _get_f16(weights, name, ref):
+ # Cache LN vectors as fp16 to cut bandwidth inside kernels.
+ key = name + "_fp16"
+ v = weights.get(key, None)
+ if v is None or v.device != ref.device:
+ v0 = weights[name]
+ if v0.dtype != torch.float16 or v0.device != ref.device:
+ v0 = v0.to(device=ref.device, dtype=torch.float16)
+ v = v0.contiguous()
+ weights[key] = v
+ return v
+
+ @triton.jit
+ def _ln_stats_kernel(
+ x_ptr, mean_ptr, rstd_ptr,
+ M: tl.constexpr, D: tl.constexpr,
+ s_xm: tl.constexpr, s_xd: tl.constexpr,
+ BM: tl.constexpr, BD: tl.constexpr,
+ ):
+ # Compute LayerNorm statistics for each row of x2d [M, D] once.
+ pid = tl.program_id(0)
+ offs_m = pid * BM + tl.arange(0, BM)
+ m_m = offs_m < M
+ s1 = tl.zeros((BM,), tl.float32)
+ s2 = tl.zeros((BM,), tl.float32)
+ for kd in range(0, D, BD):
+ offs_d = kd + tl.arange(0, BD)
+ m_d = offs_d < D
+ x = tl.load(
+ x_ptr + offs_m[:, None] * s_xm + offs_d[None, :] * s_xd,
+ mask=m_m[:, None] & m_d[None, :],
+ other=0.0,
+ ).to(tl.float32)
+ s1 += tl.sum(x, axis=1)
+ s2 += tl.sum(x * x, axis=1)
+ mean = s1 / D
+ var = s2 / D - mean * mean
+ rstd = tl.math.rsqrt(var + 1e-5)
+ tl.store(mean_ptr + offs_m, mean, mask=m_m)
+ tl.store(rstd_ptr + offs_m, rstd, mask=m_m)
+
@triton.autotune(
configs=[
triton.Config({'BM': 64, 'BD': 64, 'BH': 64}, num_warps=4, num_stages=3),
⋯ 4 unchanged lines
)
@triton.jit
def _head_fused_kernel(
- x_ptr, mask_ptr,
+ x_ptr, mask_ptr, mean_ptr, rstd_ptr,
w_lp, w_rp, w_lg, w_rg, w_og,
ln_w, ln_b,
l_out_ptr, r_out_ptr, g_out_ptr,
⋯ 3 unchanged lines
HAS_MASK: tl.constexpr,
BM: tl.constexpr, BD: tl.constexpr, BH: tl.constexpr,
):
+ # Key change: do NOT recompute LN stats per pid_h tile (that was a massive redundancy).
pid_h = tl.program_id(0)
pid_m = tl.program_id(1)
⋯ 2 unchanged lines
m_m = offs_m < M
m_h = offs_h < H
- # LN stats over D (per row)
- s1 = tl.zeros((BM,), tl.float32)
- s2 = tl.zeros((BM,), tl.float32)
- for kd in range(0, D, BD):
- offs_d = kd + tl.arange(0, BD)
- m_d = offs_d < D
- x = tl.load(x_ptr + offs_m[:, None] * s_xm + offs_d[None, :] * s_xd,
- mask=m_m[:, None] & m_d[None, :], other=0.0).to(tl.float32)
- s1 += tl.sum(x, axis=1)
- s2 += tl.sum(x * x, axis=1)
- mean = s1 / D
- var = s2 / D - mean * mean
- rstd = 1.0 / tl.sqrt(var + 1e-5)
+ mean = tl.load(mean_ptr + offs_m, mask=m_m, other=0.0).to(tl.float32)
+ rstd = tl.load(rstd_ptr + offs_m, mask=m_m, other=0.0).to(tl.float32)
- # 5 projections on LN(x)
lp = tl.zeros((BM, BH), tl.float32)
rp = tl.zeros((BM, BH), tl.float32)
lg = tl.zeros((BM, BH), tl.float32)
⋯ 3 unchanged lines
for kd in range(0, D, BD):
offs_d = kd + tl.arange(0, BD)
m_d = offs_d < D
- x = tl.load(x_ptr + offs_m[:, None] * s_xm + offs_d[None, :] * s_xd,
- mask=m_m[:, None] & m_d[None, :], other=0.0).to(tl.float32)
- w = tl.load(ln_w + offs_d, mask=m_d, other=0.0).to(tl.float32)
- b = tl.load(ln_b + offs_d, mask=m_d, other=0.0).to(tl.float32)
- x = ((x - mean[:, None]) * rstd[:, None] * w[None, :] + b[None, :]).to(tl.float16)
+ x = tl.load(
+ x_ptr + offs_m[:, None] * s_xm + offs_d[None, :] * s_xd,
+ mask=m_m[:, None] & m_d[None, :],
+ other=0.0,
+ ).to(tl.float32)
+
+ # LN affine in fp16 (stats kept in fp32)
+ w = tl.load(ln_w + offs_d, mask=m_d, other=0.0).to(tl.float16)
+ b = tl.load(ln_b + offs_d, mask=m_d, other=0.0).to(tl.float16)
+ x16 = ((x - mean[:, None]) * rstd[:, None]).to(tl.float16)
+ x16 = x16 * w[None, :] + b[None, :]
+
w_off = offs_d[:, None] * s_wk + offs_h[None, :] * s_wh
wm = m_d[:, None] & m_h[None, :]
- lp += tl.dot(x, tl.load(w_lp + w_off, mask=wm, other=0.0).to(tl.float16))
- rp += tl.dot(x, tl.load(w_rp + w_off, mask=wm, other=0.0).to(tl.float16))
- lg += tl.dot(x, tl.load(w_lg + w_off, mask=wm, other=0.0).to(tl.float16))
- rg += tl.dot(x, tl.load(w_rg + w_off, mask=wm, other=0.0).to(tl.float16))
- og += tl.dot(x, tl.load(w_og + w_off, mask=wm, other=0.0).to(tl.float16))
+ lp += tl.dot(x16, tl.load(w_lp + w_off, mask=wm, other=0.0).to(tl.float16))
+ rp += tl.dot(x16, tl.load(w_rp + w_off, mask=wm, other=0.0).to(tl.float16))
+ lg += tl.dot(x16, tl.load(w_lg + w_off, mask=wm, other=0.0).to(tl.float16))
+ rg += tl.dot(x16, tl.load(w_rg + w_off, mask=wm, other=0.0).to(tl.float16))
+ og += tl.dot(x16, tl.load(w_og + w_off, mask=wm, other=0.0).to(tl.float16))
l = lp * tl.sigmoid(lg)
r = rp * tl.sigmoid(rg)
g = tl.sigmoid(og)
- # mask only affects left/right
if HAS_MASK:
- m = tl.load(mask_ptr + offs_m, mask=m_m, other=0.0).to(tl.float32)
- l *= m[:, None]
- r *= m[:, None]
+ mm = tl.load(mask_ptr + offs_m, mask=m_m, other=0.0).to(tl.float32)
+ l *= mm[:, None]
+ r *= mm[:, None]
st = m_m[:, None] & m_h[None, :]
tl.store(g_out_ptr + offs_m[:, None] * H + offs_h[None, :], g.to(tl.float16), mask=st)
- # pack L/R into [B*H, N*N]
+ # Pack L/R contiguously as [B*H, N*N]. Let cuBLAS handle transpose via op(B)=T.
b_idx = offs_m // NN
- rem = offs_m % NN
- addr = (b_idx[:, None] * H + offs_h[None, :] ) * NN + rem[:, None]
+ rem = offs_m - b_idx * NN
+ addr = (b_idx[:, None] * H + offs_h[None, :]) * NN + rem[:, None]
tl.store(l_out_ptr + addr, l.to(tl.float16), mask=st)
tl.store(r_out_ptr + addr, r.to(tl.float16), mask=st)
+ # NOTE: repack kernel removed from the hotpath. Repacking [B*H, N*N] -> [M, H]
+ # is a full extra bandwidth pass and was a major regression at N=768/1024.
+ # Keep a tiny stub to avoid editing more code; it is no longer invoked.
@triton.jit
+ def _repack_bh_nn_to_m_h_kernel(
+ inp_ptr, out_ptr,
+ M: tl.constexpr, H: tl.constexpr, NN: tl.constexpr,
+ BM: tl.constexpr, BH: tl.constexpr,
+ ):
+ return
+
+
+ @triton.jit
def _tail_fused_kernel(
bmm_ptr, g_ptr,
w_out, ln_w, ln_b,
⋯ 2 unchanged lines
s_wh: tl.constexpr, s_wd: tl.constexpr,
BM: tl.constexpr, BD: tl.constexpr, BH: tl.constexpr,
):
+ # Read directly from [B*H, N*N] (flattened) with address math.
+ # Avoids materializing/repacking [M, H].
pid = tl.program_id(0)
offs_m = pid * BM + tl.arange(0, BM)
m_m = offs_m < M
⋯ 2 unchanged lines
m_h = offs_h < H
b_idx = offs_m // NN
- rem = offs_m % NN
- addr = (b_idx[:, None] * H + offs_h[None, :] ) * NN + rem[:, None]
+ rem = offs_m - b_idx * NN
+ addr = (b_idx[:, None] * H + offs_h[None, :]) * NN + rem[:, None]
- v = tl.load(bmm_ptr + addr, mask=m_m[:, None] & m_h[None, :], other=0.0).to(tl.float32)
- g = tl.load(g_ptr + offs_m[:, None] * H + offs_h[None, :], mask=m_m[:, None] & m_h[None, :], other=0.0).to(tl.float32)
+ v = tl.load(
+ bmm_ptr + addr,
+ mask=m_m[:, None] & m_h[None, :],
+ other=0.0,
+ ).to(tl.float32)
+ g = tl.load(
+ g_ptr + offs_m[:, None] * H + offs_h[None, :],
+ mask=m_m[:, None] & m_h[None, :],
+ other=0.0,
+ ).to(tl.float16)
+
mean = tl.sum(v, axis=1) / H
var = tl.sum(v * v, axis=1) / H - mean * mean
- rstd = 1.0 / tl.sqrt(var + 1e-5)
+ rstd = tl.math.rsqrt(var + 1e-5)
- w = tl.load(ln_w + offs_h, mask=m_h, other=0.0).to(tl.float32)
- b = tl.load(ln_b + offs_h, mask=m_h, other=0.0).to(tl.float32)
+ # fp16 LN affine + fp16 gate for tensorcore dot
+ w = tl.load(ln_w + offs_h, mask=m_h, other=0.0).to(tl.float16)
+ b = tl.load(ln_b + offs_h, mask=m_h, other=0.0).to(tl.float16)
- v = ((v - mean[:, None]) * rstd[:, None] * w[None, :] + b[None, :]) * g
- v16 = v.to(tl.float16)
+ v16 = ((v - mean[:, None]) * rstd[:, None]).to(tl.float16)
+ v16 = (v16 * w[None, :] + b[None, :]) * g
for kd in range(0, D, BD):
offs_d = kd + tl.arange(0, BD)
m_d = offs_d < D
- w_tile = tl.load(w_out + offs_h[:, None] * s_wh + offs_d[None, :] * s_wd,
- mask=m_h[:, None] & m_d[None, :], other=0.0).to(tl.float16)
+ w_tile = tl.load(
+ w_out + offs_h[:, None] * s_wh + offs_d[None, :] * s_wd,
+ mask=m_h[:, None] & m_d[None, :],
+ other=0.0,
+ ).to(tl.float16)
o = tl.dot(v16, w_tile)
- tl.store(out_ptr + offs_m[:, None] * D + offs_d[None, :], o.to(tl.float32),
- mask=m_m[:, None] & m_d[None, :])
+ tl.store(
+ out_ptr + offs_m[:, None] * D + offs_d[None, :],
+ o.to(tl.float32),
+ mask=m_m[:, None] & m_d[None, :],
+ )
- class TriMul(nn.Module):
- def __init__(self, dim: int, hidden_dim: int):
- super().__init__()
- self.dim = dim
- self.hidden_dim = hidden_dim
-
- 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)
+ # NOTE: TriMul nn.Module removed (not used by the evaluator); keeping only custom_kernel reduces code size/compile time.
- def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
- # x: [B, N, N, D]
- batch_size, seq_len, _, dim = x.shape
-
- x = self.norm(x)
- # Fuse projection and gating into PyTorch optimized ops where possible
- # We use grouped linear projections or just rely on torch.matmul efficiency
- left = self.left_proj(x) * self.left_gate(x).sigmoid()
- right = self.right_proj(x) * self.right_gate(x).sigmoid()
-
- if mask is not None:
- mask = mask.unsqueeze(-1)
- left = left * mask
- right = right * mask
-
- # --- BATCHED MATRIX MULTIPLICATION (cuBLAS) ---
- # Reshape tensors so that hidden dimension D becomes part of the batch.
- # left : [B, N, N, D] -> [B, D, N, N]
- # right: we need the transpose on the summed dimension k,
- # which corresponds to swapping the last two axes before the matmul.
- # right : [B, N, N, D] -> [B, D, N, N] and then view as transposed.
- B, N, _, D = left.shape
-
- # Bring D to the batch dimension; keep data contiguous for cuBLAS.
- left_t = left.permute(0, 3, 1, 2).contiguous() # [B, D, N, N]
- right_t = right.permute(0, 3, 2, 1).contiguous() # [B, D, N, N] (k ↔ j)
-
- # Merge batch and hidden dimensions.
- left_view = left_t.view(B * D, N, N) # [B*D, N, N]
- right_view = right_t.view(B * D, N, N) # [B*D, N, N]
-
- # Perform the batched matmul using cuBLAS (highly optimized on A100).
- out_view = torch.bmm(left_view, right_view) # [B*D, N, N]
-
- # Restore original layout: [B, N, N, D]
- out = out_view.view(B, D, N, N).permute(0, 2, 3, 1).contiguous()
- # ------------------------------------
-
- out = self.to_out_norm(out)
- out_gate = self.out_gate(x).sigmoid()
- out = out * out_gate
- return self.to_out(out)
-
-
def custom_kernel(data):
"""
- High-performance TriMul(outgoing) forward:
- - Triton head: LN(x) + 5 projections + sigmoid gates + optional mask,
- and directly pack L/R into [B*H, N*N] for tensor-core BMM.
- - torch.bmm: dominant N^3 contraction on tensor cores.
- - Triton tail: LN(out) + out_gate + final projection to dim (fp32 output).
+ Performance-oriented TriMul(outgoing) forward:
+ - Triton head: LayerNorm(x) + 5 projections + sigmoid gates (+ optional mask),
+ and directly pack L/R into [B*H, N*N] for tensor-core bmm; store out_gate.
+ This avoids materializing the massive [M,5H] 'proj' tensor (which can exceed 1GB).
+ - torch.baddbmm with a persistent output buffer to avoid per-call allocations.
+ - Triton tail: LayerNorm(out) + out_gate + final projection to dim (fp32 output).
"""
x, mask, weights, config = data
D, H = config["dim"], config["hidden_dim"]
⋯ 1 unchanged lines
NN = N * N
M = B * NN
- # cache fp16 transposed weights for tl.dot: [in, out]
+ # flatten x to [M, D]
+ x2d = x.reshape(M, D)
+ mask_flat = mask.reshape(M) if mask is not None else None
+
+ # cache fp16 transposed weights for tl.dot (shape [D,H] / [H,D])
w_lp = _get_w16_T(weights, "left_proj.weight", x)
w_rp = _get_w16_T(weights, "right_proj.weight", x)
w_lg = _get_w16_T(weights, "left_gate.weight", x)
w_rg = _get_w16_T(weights, "right_gate.weight", x)
w_og = _get_w16_T(weights, "out_gate.weight", x)
- w_to = _get_w16_T(weights, "to_out.weight", x) # [H, D] after transpose
+ w_to = _get_w16_T(weights, "to_out.weight", x) # [H, D]
- # flatten x to [M, D]
- x2d = x.reshape(M, D)
- mask_flat = mask.reshape(M) if mask is not None else None
+ # Reuse large buffers (critical for N=768/1024). Allocations here dominate otherwise.
+ scratch = weights.setdefault("_triumul_scratch", {})
+ skey = (B, N, D, H, x.device)
+ buf = scratch.get(skey, None)
+ if buf is None:
+ buf = {
+ "l_bmm": torch.empty((B * H, NN), device=x.device, dtype=torch.float16),
+ "r_bmm": torch.empty((B * H, NN), device=x.device, dtype=torch.float16),
+ "g_out": torch.empty((M, H), device=x.device, dtype=torch.float16),
+ "out_bmm": torch.empty((B * H, N, N), device=x.device, dtype=torch.float16),
+ "out2d": torch.empty((M, D), device=x.device, dtype=torch.float32),
+ # LN stats scratch (computed once per row, reused across hidden tiles)
+ "mean": torch.empty((M,), device=x.device, dtype=torch.float32),
+ "rstd": torch.empty((M,), device=x.device, dtype=torch.float32),
+ }
+ scratch[skey] = buf
- # packed for BMM: [B*H, N*N]
- l_bmm = torch.empty((B * H, NN), device=x.device, dtype=torch.float16)
- r_bmm = torch.empty((B * H, NN), device=x.device, dtype=torch.float16)
- g_out = torch.empty((M, H), device=x.device, dtype=torch.float16)
+ l_bmm = buf["l_bmm"]
+ r_bmm = buf["r_bmm"]
+ g_out = buf["g_out"]
+ # 1) LN stats once per row (avoid recomputing per hidden-tile pid_h)
+ mean = buf["mean"]
+ rstd = buf["rstd"]
+ _ln_stats_kernel[(triton.cdiv(M, 256),)](
+ x2d, mean, rstd,
+ M=M, D=D,
+ s_xm=x2d.stride(0), s_xd=x2d.stride(1),
+ BM=256, BD=64,
+ num_warps=4,
+ )
+
+ # 2) head: projections + gates (+ optional mask) + pack L/R for BMM
grid_head = lambda META: (triton.cdiv(H, META["BH"]), triton.cdiv(M, META["BM"]))
_head_fused_kernel[grid_head](
- x2d, mask_flat if mask is not None else x2d,
+ x2d, mask_flat if mask_flat is not None else x2d, mean, rstd,
w_lp, w_rp, w_lg, w_rg, w_og,
- weights["norm.weight"], weights["norm.bias"],
+ _get_f16(weights, "norm.weight", x), _get_f16(weights, "norm.bias", x),
l_bmm, r_bmm, g_out,
M=M, D=D, H=H, NN=NN,
s_xm=x2d.stride(0), s_xd=x2d.stride(1),
s_wk=w_lp.stride(0), s_wh=w_lp.stride(1),
- HAS_MASK=(mask is not None),
+ HAS_MASK=(mask_flat is not None),
)
- # tensor-core N^3 core
- out_bmm = torch.bmm(
- l_bmm.view(-1, N, N),
- r_bmm.view(-1, N, N).transpose(1, 2)
- ).contiguous()
+ # 3) tensor-core contraction; let cuBLAS handle transpose via op(B)=T (no custom repacking)
+ out_bmm = buf["out_bmm"]
+ A = l_bmm.view(B * H, N, N)
+ Bt = r_bmm.view(B * H, N, N).transpose(1, 2)
+ torch.baddbmm(out_bmm, A, Bt, beta=0.0, alpha=1.0, out=out_bmm)
- # tail: produce fp32 [M, D]
- out2d = torch.empty((M, D), device=x.device, dtype=torch.float32)
- BH = triton.next_power_of_2(H)
+ # 4) tail: read directly from [B*H, NN], LN + gate + final projection => fp32 [M, D]
+ out2d = buf["out2d"]
+ BD_TAIL = 128 if D == 128 else 64
grid_tail = (triton.cdiv(M, 64),)
_tail_fused_kernel[grid_tail](
- out_bmm.view(-1, NN), g_out,
- w_to, weights["to_out_norm.weight"], weights["to_out_norm.bias"],
+ out_bmm.view(B * H, NN), g_out,
+ w_to, _get_f16(weights, "to_out_norm.weight", x), _get_f16(weights, "to_out_norm.bias", x),
out2d,
M=M, H=H, D=D, NN=NN,
s_wh=w_to.stride(0), s_wd=w_to.stride(1),
- BM=64, BD=64, BH=BH,
+ BM=64, BD=BD_TAIL, BH=128,
num_warps=4,
)
return out2d.view(B, N, N, D)
scrolls · 415 diff lines total

Best evidence level for this revision: reported

JSON