Skip to content
KernelIndex
Search⌘K

submission 36153

msuiche · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

trimul_streamed_v4pp_planner.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-36153?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
11.8ms
#28 of 69
2025-09-08

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:048afb78377c4493b0c1d499f125af101f20c4b8bebf6f8b34cb1ebbcd7afdc0
license declaredunknown
license concludedunknown
authorsmsuiche
imported2026-08-15

Kernel source

trimul_streamed_v4pp_planner.py378 lines

#!POPCORN leaderboard trimul
#!POPCORN gpu H100
from utils import make_match_reference, DisableCuDNNTF32
from task import input_t, output_t

import os
import json
import math
import torch
import torch.nn.functional as F

# Keep harness globals unchanged
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = False

# -----------------------------------------------------------------------------
# Lightweight plan cache (default: disabled to avoid any extra startup cost)
# Enable with env:
#   TRIMUL_TUNE=1                -> time a couple of variants once per shape
#   TRIMUL_PLAN_CACHE=/path.json -> persist best plans across runs
# -----------------------------------------------------------------------------
_PLAN_CACHE = {}
_PLAN_FILE = os.getenv("TRIMUL_PLAN_CACHE", "")
_TUNE = os.getenv("TRIMUL_TUNE", "0") != "0"

def _load_plan_file():
    if _PLAN_FILE and os.path.isfile(_PLAN_FILE):
        try:
            with open(_PLAN_FILE, "r") as f:
                _PLAN_CACHE.update(json.load(f))
        except Exception:
            pass

def _save_plan_file():
    if _PLAN_FILE:
        try:
            with open(_PLAN_FILE, "w") as f:
                json.dump(_PLAN_CACHE, f)
        except Exception:
            pass

def _time_once(fn):
    s = torch.cuda.Event(enable_timing=True)
    e = torch.cuda.Event(enable_timing=True)
    s.record(); fn(); e.record(); e.synchronize()
    return s.elapsed_time(e)  # milliseconds (float)

# Optional tiny buffer cache to reduce repeated large allocations (accuracy-neutral)
_BUF = {}
def _get(key, shape, dtype, device):
    t = _BUF.get(key)
    if t is None or tuple(t.shape) != tuple(shape) or t.dtype != dtype or t.device != device:
        t = torch.empty(shape, device=device, dtype=dtype)
        _BUF[key] = t
    return t

def _pick_plan(B, N, D, H, device, runner=None):
    """Return a dict plan with keys:
         wf: 1 -> weight-first projection ([5H,D]@[D,M]), 0 -> input-first ([M,D]@[D,5H])
         th: H-chunk size
         lhs_contig: whether to make LHS contiguous before bmm (1=yes)
       Default: heuristic; if TRIMUL_TUNE=1 and runner is provided, time a few variants once.
    """
    key = f"{B}-{N}-{D}-{H}"
    if key in _PLAN_CACHE:
        return _PLAN_CACHE[key]

    # Default heuristic (fast, no timing)
    # H multiples of 32 are ideal (we won't assert to keep compatibility)
    M = B * N * N
    plan = {}
    plan["wf"] = 1 if (M >= 8 * D or N >= 768) else 0
    if H >= 256:   th = 128
    elif H >= 128: th = 128
    elif H >= 64:  th = 64
    else:          th = H
    if (H == 128) and (D >= 384) and (N >= 1024):
        th = 64
    plan["th"] = th
    plan["lhs_contig"] = 1

    if not _TUNE or runner is None:
        _PLAN_CACHE[key] = plan
        return plan

    # Load persisted plans if any
    _load_plan_file()
    if key in _PLAN_CACHE:
        return _PLAN_CACHE[key]

    # Try a tiny set of candidates; warm up, then time once each
    cands = []
    for wf in (0, 1):
        for th in ((64, 128) if H >= 128 else (H,)):
            cands.append({"wf": wf, "th": th, "lhs_contig": 1})

    # Warmup all
    for c in cands:
        runner(c, warmup=True)
    torch.cuda.synchronize()

    best = None
    best_ms = 1e9
    for c in cands:
        ms = _time_once(lambda: runner(c, warmup=False))
        if ms < best_ms:
            best, best_ms = c, ms

    _PLAN_CACHE[key] = best
    _save_plan_file()
    return best


def custom_kernel(data: input_t) -> output_t:
    """
    Two-pass streamed TriMul with a lightweight shape planner (off by default):
      • One big projection GEMM (orientation auto-picked per shape or via tiny search)
      • Mask applied once (left only), with an all-ones fast-path
      • PASS 1: contraction per H-chunk to accumulate mean/var (no EIN writes)
      • PASS 2: recompute contraction chunk, apply LN(g), accumulate directly into OUT via addmm_
      • Chunk size TH kept “fat” (K large) with a small exception for (H=128,D=384,N>=1024)
      • FP32 math, LayerNorm eps=1e-5, no clamping; DisableCuDNNTF32() untouched
      • cuBLAS/cuBLASLt used for heavy GEMMs, TF32 allowed (as in harness)
    """
    with DisableCuDNNTF32():
        input_tensor, mask, weights, config = data
        B, N, _, D = input_tensor.shape
        H = config["hidden_dim"]
        device = input_tensor.device

        # Prefer Tensor Cores / TF32 for cuBLAS/cuBLASLt (fast on H100)
        prev_tf32 = torch.backends.cuda.matmul.allow_tf32
        torch.backends.cuda.matmul.allow_tf32 = True
        prev_prec = torch.get_float32_matmul_precision() if hasattr(torch, "get_float32_matmul_precision") else None
        if hasattr(torch, "set_float32_matmul_precision"):
            torch.set_float32_matmul_precision("high")

        try:
            # 0) Input LayerNorm (FP32; eps=1e-5; no clamping)
            x = F.layer_norm(
                input_tensor, (D,),
                weight=weights["norm.weight"],
                bias=weights["norm.bias"],
                eps=1e-5,
            )

            # Optional tiny runner used only when TRIMUL_TUNE=1:
            # runs a single small iteration to pick wf/th; avoids big copies.
            def _runner(plan, warmup=True):
                wf = plan["wf"]; th = plan["th"]
                M = B * N * N
                # Projections (one GEMM)
                if wf:
                    x2dT = x.view(M, D).t().contiguous()               # [D, M]
                    Wcat_key = "__proj_Wcat__"                         # [5H, D]
                    Wcat = weights.get(Wcat_key)
                    if (Wcat is None) or (Wcat.shape != (5 * H, D)) or (Wcat.device != device):
                        Wcat = torch.cat([
                            weights['left_proj.weight' ],
                            weights['right_proj.weight'],
                            weights['left_gate.weight' ],
                            weights['right_gate.weight'],
                            weights['out_gate.weight'  ],
                        ], dim=0).contiguous()
                        weights[Wcat_key] = Wcat
                    PT = torch.matmul(Wcat, x2dT).view(5, H, M)        # [5,H,M]
                    Lpre_T, Rpre_T, LGpre_T, RGpre_T, OGpre_T = PT[0], PT[1], PT[2], PT[3], PT[4]
                else:
                    x2d = x.view(M, D)                                 # [M, D]
                    WcatT_key = "__proj_Wcat_T__"                      # [D, 5H]
                    Wcat_T = weights.get(WcatT_key)
                    if (Wcat_T is None) or (Wcat_T.shape != (D, 5 * H)) or (Wcat_T.device != device):
                        Wcat_T = torch.cat([
                            weights['left_proj.weight' ].t().contiguous(),
                            weights['right_proj.weight'].t().contiguous(),
                            weights['left_gate.weight' ].t().contiguous(),
                            weights['right_gate.weight'].t().contiguous(),
                            weights['out_gate.weight'  ].t().contiguous(),
                        ], dim=1).contiguous()
                        weights[WcatT_key] = Wcat_T
                    P = torch.matmul(x2d, Wcat_T)                      # [M,5H]
                    Lpre, Rpre, LGpre, RGpre, OGpre = torch.split(P, H, dim=1)
                    Lpre_T, Rpre_T, LGpre_T, RGpre_T, OGpre_T = Lpre.t(), Rpre.t(), LGpre.t(), RGpre.t(), OGpre.t()

                # Nomask fast-path
                all_ones = False
                try:
                    mn = float(mask.min().item()); mx = float(mask.max().item())
                    all_ones = (mn == 1.0 and mx == 1.0)
                except Exception:
                    pass

                if all_ones:
                    LEFT_T  = torch.sigmoid(LGpre_T) * Lpre_T
                else:
                    mrow = mask.to(torch.float32).view(1, M)
                    LEFT_T  = torch.sigmoid(LGpre_T) * Lpre_T * mrow
                RIGHT_T = torch.sigmoid(RGpre_T) * Rpre_T
                # One tiny contraction chunk to get a timing signal
                t = th if th <= H else H
                Lbt = LEFT_T.view(H, B, N, N)[:t].reshape(t * B, N, N).contiguous()
                Rbt = RIGHT_T.view(H, B, N, N)[:t].reshape(t * B, N, N)
                _ = torch.bmm(Lbt, Rbt.transpose(1, 2))               # discard
                if not warmup:
                    torch.cuda.synchronize()

            # Select plan
            plan = _pick_plan(B, N, D, H, device, runner=_runner if _TUNE else None)
            wf = plan["wf"]; TH = plan["th"]; lhs_contig = plan["lhs_contig"]

            # 1) Projections (one GEMM), obeying plan["wf"]
            M = B * N * N
            if wf:
                x2dT = x.view(M, D).t().contiguous()                   # [D, M]
                Wcat_key = "__proj_Wcat__"                             # [5H, D]
                Wcat = weights.get(Wcat_key)
                if (Wcat is None) or (Wcat.shape != (5 * H, D)) or (Wcat.device != device):
                    Wcat = torch.cat([
                        weights['left_proj.weight' ],
                        weights['right_proj.weight'],
                        weights['left_gate.weight' ],
                        weights['right_gate.weight'],
                        weights['out_gate.weight'  ],
                    ], dim=0).contiguous()
                    weights[Wcat_key] = Wcat
                PT = torch.matmul(Wcat, x2dT).view(5, H, M)            # [5,H,M]
                Lpre_T, Rpre_T, LGpre_T, RGpre_T, OGpre_T = PT[0], PT[1], PT[2], PT[3], PT[4]
            else:
                x2d = x.view(M, D)                                     # [M, D]
                WcatT_key = "__proj_Wcat_T__"                          # [D, 5H]
                Wcat_T = weights.get(WcatT_key)
                if (Wcat_T is None) or (Wcat_T.shape != (D, 5 * H)) or (Wcat_T.device != device):
                    Wcat_T = torch.cat([
                        weights['left_proj.weight' ].t().contiguous(),
                        weights['right_proj.weight'].t().contiguous(),
                        weights['left_gate.weight' ].t().contiguous(),
                        weights['right_gate.weight'].t().contiguous(),
                        weights['out_gate.weight'  ].t().contiguous(),
                    ], dim=1).contiguous()
                    weights[WcatT_key] = Wcat_T
                P = torch.matmul(x2d, Wcat_T)                           # [M,5H]
                Lpre, Rpre, LGpre, RGpre, OGpre = torch.split(P, H, dim=1)
                Lpre_T, Rpre_T, LGpre_T, RGpre_T, OGpre_T = Lpre.t(), Rpre.t(), LGpre.t(), RGpre.t(), OGpre.t()

            # 2) Gates + mask once (left only) with an all-ones fast-path
            all_ones = False
            try:
                mn = float(mask.min().item()); mx = float(mask.max().item())
                all_ones = (mn == 1.0 and mx == 1.0)
            except Exception:
                pass

            if all_ones:
                LEFT_T  = torch.sigmoid(LGpre_T) * Lpre_T
            else:
                mrow = mask.to(torch.float32).view(1, M)
                LEFT_T  = torch.sigmoid(LGpre_T) * Lpre_T * mrow
            RIGHT_T = torch.sigmoid(RGpre_T) * Rpre_T
            OG_T    = torch.sigmoid(OGpre_T)

            # Views as [H, B, N, N] (no copies)
            LEFT_HBNN  = LEFT_T.view(H, B, N, N)
            RIGHT_HBNN = RIGHT_T.view(H, B, N, N)
            OG_HBNN    = OG_T.view(H, B, N, N)

            # 3) PASS 1: accumulate mean/var over H (no EIN/G materialization)
            S  = _get(("S",  B, N, N, device), (B, N, N), torch.float32, device);  S.zero_()
            S2 = _get(("S2", B, N, N, device), (B, N, N), torch.float32, device); S2.zero_()

            for h0 in range(0, H, TH):
                h1 = min(H, h0 + TH); t = h1 - h0
                Lbt = LEFT_HBNN [h0:h1].view(t * B, N, N)
                if lhs_contig: Lbt = Lbt.contiguous()
                Rbt = RIGHT_HBNN[h0:h1].view(t * B, N, N)
                Cbt = torch.bmm(Lbt, Rbt.transpose(1, 2))                 # [B*t, N, N]
                C   = Cbt.view(t, B, N, N)
                S  += C.sum(dim=0)
                S2 += (C * C).sum(dim=0)

            Hf = float(H)
            mean = S / Hf
            var  = S2 / Hf - mean * mean
            inv_std = torch.rsqrt(var + 1e-5)                             # [B, N, N]

            # 4) PASS 2: recompute contraction chunks, apply LN(g), accumulate into OUT
            Wt_key = "__to_out_wT__"  # [H, D]
            Wt_full = weights.get(Wt_key)
            if (Wt_full is None) or (Wt_full.shape != (H, D)) or (Wt_full.device != device):
                Wt_full = weights['to_out.weight'].t().contiguous()
                weights[Wt_key] = Wt_full

            OUT2D = _get(("OUT2D", M, D, device), (M, D), torch.float32, device)
            # Use beta=0 on first addmm to avoid a large memset
            LNw = weights['to_out_norm.weight']   # [H]
            LNb = weights['to_out_norm.bias']     # [H]

            mean_ = mean.unsqueeze(0)  # [1, B, N, N]
            inv_  = inv_std.unsqueeze(0)

            first = True
            for h0 in range(0, H, TH):
                h1 = min(H, h0 + TH); t = h1 - h0
                Lbt = LEFT_HBNN [h0:h1].view(t * B, N, N)
                if lhs_contig: Lbt = Lbt.contiguous()
                Rbt = RIGHT_HBNN[h0:h1].view(t * B, N, N)
                Cbt = torch.bmm(Lbt, Rbt.transpose(1, 2))                 # [B*t, N, N]
                C   = Cbt.view(t, B, N, N)                                 # [t, B, N, N]

                lnw = LNw[h0:h1].view(t, 1, 1, 1)
                lnb = LNb[h0:h1].view(t, 1, 1, 1)
                Cn  = ((C - mean_) * inv_) * lnw + lnb

                OGc = OG_HBNN[h0:h1]                                       # [t, B, N, N]
                G   = Cn * OGc                                             # [t, B, N, N]

                GflatT = G.view(t, M)                                      # [t, M]
                Wt     = Wt_full[h0:h1, :]                                  # [t, D]
                OUT2D.addmm_(GflatT.t(), Wt, beta=(0.0 if first else 1.0), alpha=1.0)
                first = False

            return OUT2D.view(B, N, N, D)

        finally:
            torch.backends.cuda.matmul.allow_tf32 = prev_tf32
            if hasattr(torch, "set_float32_matmul_precision") and prev_prec is not None:
                torch.set_float32_matmul_precision(prev_prec)


# ============================================================
# Input generation (unchanged)
def generate_input(seqlen: int, bs: int, dim: int, hiddendim: int,
                   seed: int, nomask: bool, distribution: str) -> input_t:
    batch_size = bs
    seq_len = seqlen
    hidden_dim = hiddendim
    no_mask = nomask

    config = {"hidden_dim": hidden_dim, "dim": dim}

    gen = torch.Generator(device='cuda')
    gen.manual_seed(seed)

    weights = {}

    if distribution == "cauchy":
        input_tensor = torch.distributions.Cauchy(0, 2).sample(
            (batch_size, seq_len, seq_len, dim)
        ).to(device='cuda', dtype=torch.float32)
    else:
        input_tensor = torch.randn(
            (batch_size, seq_len, seq_len, dim),
            device='cuda', dtype=torch.float32, generator=gen
        ).contiguous()

    if no_mask:
        mask = torch.ones(batch_size, seq_len, seq_len, device=input_tensor.device)
    else:
        mask = torch.randint(0, 2, (batch_size, seq_len, seq_len),
                             device=input_tensor.device, generator=gen)

    weights["norm.weight"] = torch.randn(dim, device="cuda", dtype=torch.float32)
    weights["norm.bias"] = torch.randn(dim, device="cuda", dtype=torch.float32)
    weights["left_proj.weight"]  = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
    weights["right_proj.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
    weights["left_gate.weight"]  = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
    weights["right_gate.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
    weights["out_gate.weight"]   = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
    weights["to_out_norm.weight"] = torch.randn(hidden_dim, device="cuda", dtype=torch.float32)
    weights["to_out.weight"]      = torch.randn(dim, hidden_dim, device="cuda", dtype=torch.float32) / math.sqrt(dim)
    weights["to_out_norm.bias"]   = torch.randn(hidden_dim, device="cuda", dtype=torch.float32)

    return (input_tensor, mask, weights, config)


# Correctness check
check_implementation = make_match_reference(custom_kernel, rtol=2e-2, atol=2e-2)
scrolls · 378 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 35764.

+
#!POPCORN leaderboard trimul
#!POPCORN gpu H100
from utils import make_match_reference, DisableCuDNNTF32
from task import input_t, output_t
+ import os
+ import json
+ import math
import torch
import torch.nn.functional as F
- import triton
- import triton.language as tl
- import math
# Keep harness globals unchanged
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = False
+ # -----------------------------------------------------------------------------
+ # Lightweight plan cache (default: disabled to avoid any extra startup cost)
+ # Enable with env:
+ # TRIMUL_TUNE=1 -> time a couple of variants once per shape
+ # TRIMUL_PLAN_CACHE=/path.json -> persist best plans across runs
+ # -----------------------------------------------------------------------------
+ _PLAN_CACHE = {}
+ _PLAN_FILE = os.getenv("TRIMUL_PLAN_CACHE", "")
+ _TUNE = os.getenv("TRIMUL_TUNE", "0") != "0"
- # ============================================================
- # 1) Fused 5× projections + gates + mask
- # ============================================================
- @triton.jit
- def proj5_gated_mask_kernel(
- X_ptr, # float32 [M, D]
- LW_ptr, RW_ptr, LGW_ptr, RGW_ptr, OGW_ptr, # float32 [H, D]
- MASK_ptr, # float32 [M] (0/1)
- LEFT_ptr, RIGHT_ptr, OG_ptr, # float32 [M, H]
- M, D, H,
- stride_x_m, stride_x_d,
- stride_w_h, stride_w_d,
- stride_o_m, stride_o_h,
- BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
- ):
- pid_m = tl.program_id(0)
- pid_n = tl.program_id(1)
+ def _load_plan_file():
+ if _PLAN_FILE and os.path.isfile(_PLAN_FILE):
+ try:
+ with open(_PLAN_FILE, "r") as f:
+ _PLAN_CACHE.update(json.load(f))
+ except Exception:
+ pass
- offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
- offs_h = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
- offs_k = tl.arange(0, BLOCK_K)
+ def _save_plan_file():
+ if _PLAN_FILE:
+ try:
+ with open(_PLAN_FILE, "w") as f:
+ json.dump(_PLAN_CACHE, f)
+ except Exception:
+ pass
- m_mask = offs_m < M
- h_mask = offs_h < H
+ def _time_once(fn):
+ s = torch.cuda.Event(enable_timing=True)
+ e = torch.cuda.Event(enable_timing=True)
+ s.record(); fn(); e.record(); e.synchronize()
+ return s.elapsed_time(e) # milliseconds (float)
- acc_l = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
- acc_r = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
- acc_lg = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
- acc_rg = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
- acc_og = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
+ # Optional tiny buffer cache to reduce repeated large allocations (accuracy-neutral)
+ _BUF = {}
+ def _get(key, shape, dtype, device):
+ t = _BUF.get(key)
+ if t is None or tuple(t.shape) != tuple(shape) or t.dtype != dtype or t.device != device:
+ t = torch.empty(shape, device=device, dtype=dtype)
+ _BUF[key] = t
+ return t
- tl.multiple_of(offs_k, 16)
- tl.multiple_of(offs_h, 16)
+ def _pick_plan(B, N, D, H, device, runner=None):
+ """Return a dict plan with keys:
+ wf: 1 -> weight-first projection ([5H,D]@[D,M]), 0 -> input-first ([M,D]@[D,5H])
+ th: H-chunk size
+ lhs_contig: whether to make LHS contiguous before bmm (1=yes)
+ Default: heuristic; if TRIMUL_TUNE=1 and runner is provided, time a few variants once.
+ """
+ key = f"{B}-{N}-{D}-{H}"
+ if key in _PLAN_CACHE:
+ return _PLAN_CACHE[key]
- num_k = tl.cdiv(D, BLOCK_K)
- for kb in range(num_k):
- k = kb * BLOCK_K + offs_k
- k_mask = k < D
+ # Default heuristic (fast, no timing)
+ # H multiples of 32 are ideal (we won't assert to keep compatibility)
+ M = B * N * N
+ plan = {}
+ plan["wf"] = 1 if (M >= 8 * D or N >= 768) else 0
+ if H >= 256: th = 128
+ elif H >= 128: th = 128
+ elif H >= 64: th = 64
+ else: th = H
+ if (H == 128) and (D >= 384) and (N >= 1024):
+ th = 64
+ plan["th"] = th
+ plan["lhs_contig"] = 1
- # X tile [M, K]
- x_ptrs = X_ptr + offs_m[:, None] * stride_x_m + k[None, :] * stride_x_d
- X_blk = tl.load(x_ptrs, mask=(m_mask[:, None] & k_mask[None, :]), other=0.0)
+ if not _TUNE or runner is None:
+ _PLAN_CACHE[key] = plan
+ return plan
- # Five weight tiles [H, K]
- lw_ptrs = LW_ptr + offs_h[:, None] * stride_w_h + k[None, :] * stride_w_d
- rw_ptrs = RW_ptr + offs_h[:, None] * stride_w_h + k[None, :] * stride_w_d
- lgw_ptrs = LGW_ptr + offs_h[:, None] * stride_w_h + k[None, :] * stride_w_d
- rgw_ptrs = RGW_ptr + offs_h[:, None] * stride_w_h + k[None, :] * stride_w_d
- ogw_ptrs = OGW_ptr + offs_h[:, None] * stride_w_h + k[None, :] * stride_w_d
+ # Load persisted plans if any
+ _load_plan_file()
+ if key in _PLAN_CACHE:
+ return _PLAN_CACHE[key]
- LW_blk = tl.load(lw_ptrs, mask=(h_mask[:, None] & k_mask[None, :]), other=0.0)
- RW_blk = tl.load(rw_ptrs, mask=(h_mask[:, None] & k_mask[None, :]), other=0.0)
- LGW_blk = tl.load(lgw_ptrs, mask=(h_mask[:, None] & k_mask[None, :]), other=0.0)
- RGW_blk = tl.load(rgw_ptrs, mask=(h_mask[:, None] & k_mask[None, :]), other=0.0)
- OGW_blk = tl.load(ogw_ptrs, mask=(h_mask[:, None] & k_mask[None, :]), other=0.0)
+ # Try a tiny set of candidates; warm up, then time once each
+ cands = []
+ for wf in (0, 1):
+ for th in ((64, 128) if H >= 128 else (H,)):
+ cands.append({"wf": wf, "th": th, "lhs_contig": 1})
- # FP32 matmul (no TF32)
- acc_l += tl.dot(X_blk, tl.trans(LW_blk), allow_tf32=True)
- acc_r += tl.dot(X_blk, tl.trans(RW_blk), allow_tf32=True)
- acc_lg += tl.dot(X_blk, tl.trans(LGW_blk), allow_tf32=True)
- acc_rg += tl.dot(X_blk, tl.trans(RGW_blk), allow_tf32=True)
- acc_og += tl.dot(X_blk, tl.trans(OGW_blk), allow_tf32=True)
+ # Warmup all
+ for c in cands:
+ runner(c, warmup=True)
+ torch.cuda.synchronize()
- # Gates + mask
- lgate = tl.sigmoid(acc_lg)
- rgate = tl.sigmoid(acc_rg)
- ogate = tl.sigmoid(acc_og)
+ best = None
+ best_ms = 1e9
+ for c in cands:
+ ms = _time_once(lambda: runner(c, warmup=False))
+ if ms < best_ms:
+ best, best_ms = c, ms
- mval = tl.load(MASK_ptr + offs_m, mask=m_mask, other=0.0) # [M]
- mval = mval[:, None] # [M,1]
+ _PLAN_CACHE[key] = best
+ _save_plan_file()
+ return best
- left = acc_l * lgate * mval
- right = acc_r * rgate * mval
- # Stores
- left_ptrs = LEFT_ptr + offs_m[:, None] * stride_o_m + offs_h[None, :] * stride_o_h
- right_ptrs = RIGHT_ptr + offs_m[:, None] * stride_o_m + offs_h[None, :] * stride_o_h
- og_ptrs = OG_ptr + offs_m[:, None] * stride_o_m + offs_h[None, :] * stride_o_h
-
- tl.store(left_ptrs, left, mask=(m_mask[:, None] & h_mask[None, :]))
- tl.store(right_ptrs, right, mask=(m_mask[:, None] & h_mask[None, :]))
- tl.store(og_ptrs, ogate, mask=(m_mask[:, None] & h_mask[None, :]))
-
-
- # ============================================================
- # 2) Contraction: EIN[b,i,j,h] = sum_k LEFT[b,i,k,h] * RIGHT[b,j,k,h]
- # Vectorized: broadcast over I/J, reduce over K (no per-h indexing)
- # ============================================================
- @triton.jit
- def contraction_kernel(
- LEFT_ptr, RIGHT_ptr, OUT_ptr, # float32
- B, N, H,
- stride_l_b, stride_l_i, stride_l_k, stride_l_h,
- stride_r_b, stride_r_j, stride_r_k, stride_r_h,
- stride_o_b, stride_o_i, stride_o_j, stride_o_h,
- BLOCK_I: tl.constexpr, BLOCK_J: tl.constexpr,
- BLOCK_K: tl.constexpr, BLOCK_H: tl.constexpr,
- ):
- # Grid mapping: x-dim covers (b, h-tile); y -> i-tiles; z -> j-tiles
- pid_bh = tl.program_id(0)
- pid_i = tl.program_id(1)
- pid_j = tl.program_id(2)
-
- # Decode i/j tiles
- 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
-
- # Decode (b, h_start) from pid_bh
- tiles_h = (H + BLOCK_H - 1) // BLOCK_H # runtime integer ok
- b = pid_bh // tiles_h
- h_tile = pid_bh % tiles_h
- h_start = h_tile * BLOCK_H
-
- # Iterate over the H micro-tile with compile-time unrolling
- for h_rel in tl.static_range(0, BLOCK_H):
- h = h_start + h_rel
- h_valid = h < H
-
- # Accumulator for this single h
- acc = tl.zeros((BLOCK_I, BLOCK_J), dtype=tl.float32)
-
- # Stream over K dimension
- for k0 in range(0, N, BLOCK_K):
- offs_k = k0 + tl.arange(0, BLOCK_K)
- mask_k = offs_k < N
-
- # LEFT[b, i, k, h] -> [I, K]
- l_ptrs = (LEFT_ptr
- + b * stride_l_b
- + offs_i[:, None] * stride_l_i
- + offs_k[None, :] * stride_l_k
- + h * stride_l_h)
- L = tl.load(l_ptrs,
- mask=(mask_i[:, None] & mask_k[None, :] & h_valid),
- other=0.0)
-
- # RIGHT[b, j, k, h] -> [J, K]
- r_ptrs = (RIGHT_ptr
- + b * stride_r_b
- + offs_j[:, None] * stride_r_j
- + offs_k[None, :] * stride_r_k
- + h * stride_r_h)
- R = tl.load(r_ptrs,
- mask=(mask_j[:, None] & mask_k[None, :] & h_valid),
- other=0.0)
-
- # Use TF32 on tensor cores where available
- acc += tl.dot(L, tl.trans(R), allow_tf32=True)
-
- # Store EIN[b, i, j, h] for this h
- o_ptrs = (OUT_ptr
- + b * stride_o_b
- + offs_i[:, None] * stride_o_i
- + offs_j[None, :] * stride_o_j
- + h * stride_o_h)
- tl.store(o_ptrs, acc, mask=(mask_i[:, None] & mask_j[None, :] & h_valid))
-
-
- # ============================================================
- # 3) Epilogue: LN over H (no clamp; eps=1e-5) -> * out_gate_sigmoid -> final W[D,H]
- # ============================================================
- @triton.jit
- def epilogue_ln_gate_kernel(
- EIN_ptr, OG_ptr, # float32 [B, N, N, H]
- LN_w_ptr, LN_b_ptr, # float32 [H]
- G_ptr, # float32 [B, N, N, H] (output: ln(ein)*og)
- B, N, H,
- stride_e_b, stride_e_i, stride_e_j, stride_e_h,
- stride_g_b, stride_g_i, stride_g_j, stride_g_h,
- BLOCK_H: tl.constexpr,
- ):
- pid_pos = tl.program_id(0) # (b,i,j)
-
- total_pos = B * N * N
- if pid_pos >= total_pos:
- return
-
- b = pid_pos // (N * N)
- rem = pid_pos % (N * N)
- i = rem // N
- j = rem % N
-
- # Stats over H (no clamping)
- sum_x = 0.0
- sum_x2 = 0.0
- for h0 in range(0, H, BLOCK_H):
- offs_h = h0 + tl.arange(0, BLOCK_H)
- mask_h = offs_h < H
- e_ptrs = EIN_ptr + b * stride_e_b + i * stride_e_i + j * stride_e_j + offs_h * stride_e_h
- vals = tl.load(e_ptrs, mask=mask_h, other=0.0) # [BLOCK_H]
- sum_x += tl.sum(vals)
- sum_x2 += tl.sum(vals * vals)
-
- Hf = tl.full((1,), H, tl.float32)
- mean = sum_x / Hf
- var = sum_x2 / Hf - mean * mean
- inv_std = tl.rsqrt(var + 1e-5)
-
- # Write normalized-and-gated vector to G
- for h0 in range(0, H, BLOCK_H):
- offs_h = h0 + tl.arange(0, BLOCK_H)
- mask_h = offs_h < H
-
- e_ptrs = EIN_ptr + b * stride_e_b + i * stride_e_i + j * stride_e_j + offs_h * stride_e_h
- og_ptrs = OG_ptr + b * stride_e_b + i * stride_e_i + j * stride_e_j + offs_h * stride_e_h
- g_ptrs = G_ptr + b * stride_g_b + i * stride_g_i + j * stride_g_j + offs_h * stride_g_h
-
- lnw = tl.load(LN_w_ptr + offs_h, mask=mask_h, other=1.0)
- lnb = tl.load(LN_b_ptr + offs_h, mask=mask_h, other=0.0)
- ein = tl.load(e_ptrs, mask=mask_h, other=0.0)
- og = tl.load(og_ptrs, mask=mask_h, other=0.0) # already sigmoid'd
-
- normed = ((ein - mean) * inv_std) * lnw + lnb
- gated = normed * og # [BLOCK_H]
-
- tl.store(g_ptrs, gated, mask=mask_h)
-
-
- # ============================================================
- # Python wrapper
- # ============================================================
def custom_kernel(data: input_t) -> output_t:
+ """
+ Two-pass streamed TriMul with a lightweight shape planner (off by default):
+ • One big projection GEMM (orientation auto-picked per shape or via tiny search)
+ • Mask applied once (left only), with an all-ones fast-path
+ • PASS 1: contraction per H-chunk to accumulate mean/var (no EIN writes)
+ • PASS 2: recompute contraction chunk, apply LN(g), accumulate directly into OUT via addmm_
+ • Chunk size TH kept “fat” (K large) with a small exception for (H=128,D=384,N>=1024)
+ • FP32 math, LayerNorm eps=1e-5, no clamping; DisableCuDNNTF32() untouched
+ • cuBLAS/cuBLASLt used for heavy GEMMs, TF32 allowed (as in harness)
+ """
with DisableCuDNNTF32():
input_tensor, mask, weights, config = data
B, N, _, D = input_tensor.shape
H = config["hidden_dim"]
+ device = input_tensor.device
- # Prefer Tensor Cores / TF32 for speed on Ampere+/Hopper
+ # Prefer Tensor Cores / TF32 for cuBLAS/cuBLASLt (fast on H100)
prev_tf32 = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = True
prev_prec = torch.get_float32_matmul_precision() if hasattr(torch, "get_float32_matmul_precision") else None
if hasattr(torch, "set_float32_matmul_precision"):
torch.set_float32_matmul_precision("high")
+
try:
- # 0) Input LayerNorm (fused), FP32, eps=1e-5 (no clamp)
+ # 0) Input LayerNorm (FP32; eps=1e-5; no clamping)
x = F.layer_norm(
input_tensor, (D,),
weight=weights["norm.weight"],
bias=weights["norm.bias"],
eps=1e-5,
- ).contiguous()
+ )
- # Flatten to [M, D]
+ # Optional tiny runner used only when TRIMUL_TUNE=1:
+ # runs a single small iteration to pick wf/th; avoids big copies.
+ def _runner(plan, warmup=True):
+ wf = plan["wf"]; th = plan["th"]
+ M = B * N * N
+ # Projections (one GEMM)
+ if wf:
+ x2dT = x.view(M, D).t().contiguous() # [D, M]
+ Wcat_key = "__proj_Wcat__" # [5H, D]
+ Wcat = weights.get(Wcat_key)
+ if (Wcat is None) or (Wcat.shape != (5 * H, D)) or (Wcat.device != device):
+ Wcat = torch.cat([
+ weights['left_proj.weight' ],
+ weights['right_proj.weight'],
+ weights['left_gate.weight' ],
+ weights['right_gate.weight'],
+ weights['out_gate.weight' ],
+ ], dim=0).contiguous()
+ weights[Wcat_key] = Wcat
+ PT = torch.matmul(Wcat, x2dT).view(5, H, M) # [5,H,M]
+ Lpre_T, Rpre_T, LGpre_T, RGpre_T, OGpre_T = PT[0], PT[1], PT[2], PT[3], PT[4]
+ else:
+ x2d = x.view(M, D) # [M, D]
+ WcatT_key = "__proj_Wcat_T__" # [D, 5H]
+ Wcat_T = weights.get(WcatT_key)
+ if (Wcat_T is None) or (Wcat_T.shape != (D, 5 * H)) or (Wcat_T.device != device):
+ Wcat_T = torch.cat([
+ weights['left_proj.weight' ].t().contiguous(),
+ weights['right_proj.weight'].t().contiguous(),
+ weights['left_gate.weight' ].t().contiguous(),
+ weights['right_gate.weight'].t().contiguous(),
+ weights['out_gate.weight' ].t().contiguous(),
+ ], dim=1).contiguous()
+ weights[WcatT_key] = Wcat_T
+ P = torch.matmul(x2d, Wcat_T) # [M,5H]
+ Lpre, Rpre, LGpre, RGpre, OGpre = torch.split(P, H, dim=1)
+ Lpre_T, Rpre_T, LGpre_T, RGpre_T, OGpre_T = Lpre.t(), Rpre.t(), LGpre.t(), RGpre.t(), OGpre.t()
+
+ # Nomask fast-path
+ all_ones = False
+ try:
+ mn = float(mask.min().item()); mx = float(mask.max().item())
+ all_ones = (mn == 1.0 and mx == 1.0)
+ except Exception:
+ pass
+
+ if all_ones:
+ LEFT_T = torch.sigmoid(LGpre_T) * Lpre_T
+ else:
+ mrow = mask.to(torch.float32).view(1, M)
+ LEFT_T = torch.sigmoid(LGpre_T) * Lpre_T * mrow
+ RIGHT_T = torch.sigmoid(RGpre_T) * Rpre_T
+ # One tiny contraction chunk to get a timing signal
+ t = th if th <= H else H
+ Lbt = LEFT_T.view(H, B, N, N)[:t].reshape(t * B, N, N).contiguous()
+ Rbt = RIGHT_T.view(H, B, N, N)[:t].reshape(t * B, N, N)
+ _ = torch.bmm(Lbt, Rbt.transpose(1, 2)) # discard
+ if not warmup:
+ torch.cuda.synchronize()
+
+ # Select plan
+ plan = _pick_plan(B, N, D, H, device, runner=_runner if _TUNE else None)
+ wf = plan["wf"]; TH = plan["th"]; lhs_contig = plan["lhs_contig"]
+
+ # 1) Projections (one GEMM), obeying plan["wf"]
M = B * N * N
- x2d = x.view(M, D)
- mask_f = mask.to(dtype=torch.float32).reshape(M).contiguous()
+ if wf:
+ x2dT = x.view(M, D).t().contiguous() # [D, M]
+ Wcat_key = "__proj_Wcat__" # [5H, D]
+ Wcat = weights.get(Wcat_key)
+ if (Wcat is None) or (Wcat.shape != (5 * H, D)) or (Wcat.device != device):
+ Wcat = torch.cat([
+ weights['left_proj.weight' ],
+ weights['right_proj.weight'],
+ weights['left_gate.weight' ],
+ weights['right_gate.weight'],
+ weights['out_gate.weight' ],
+ ], dim=0).contiguous()
+ weights[Wcat_key] = Wcat
+ PT = torch.matmul(Wcat, x2dT).view(5, H, M) # [5,H,M]
+ Lpre_T, Rpre_T, LGpre_T, RGpre_T, OGpre_T = PT[0], PT[1], PT[2], PT[3], PT[4]
+ else:
+ x2d = x.view(M, D) # [M, D]
+ WcatT_key = "__proj_Wcat_T__" # [D, 5H]
+ Wcat_T = weights.get(WcatT_key)
+ if (Wcat_T is None) or (Wcat_T.shape != (D, 5 * H)) or (Wcat_T.device != device):
+ Wcat_T = torch.cat([
+ weights['left_proj.weight' ].t().contiguous(),
+ weights['right_proj.weight'].t().contiguous(),
+ weights['left_gate.weight' ].t().contiguous(),
+ weights['right_gate.weight'].t().contiguous(),
+ weights['out_gate.weight' ].t().contiguous(),
+ ], dim=1).contiguous()
+ weights[WcatT_key] = Wcat_T
+ P = torch.matmul(x2d, Wcat_T) # [M,5H]
+ Lpre, Rpre, LGpre, RGpre, OGpre = torch.split(P, H, dim=1)
+ Lpre_T, Rpre_T, LGpre_T, RGpre_T, OGpre_T = Lpre.t(), Rpre.t(), LGpre.t(), RGpre.t(), OGpre.t()
- # Contiguous weights
- LW = weights['left_proj.weight' ].contiguous() # [H,D]
- RW = weights['right_proj.weight'].contiguous()
- LGW = weights['left_gate.weight' ].contiguous()
- RGW = weights['right_gate.weight'].contiguous()
- OGW = weights['out_gate.weight' ].contiguous()
+ # 2) Gates + mask once (left only) with an all-ones fast-path
+ all_ones = False
+ try:
+ mn = float(mask.min().item()); mx = float(mask.max().item())
+ all_ones = (mn == 1.0 and mx == 1.0)
+ except Exception:
+ pass
- # Outputs of projection kernel
- LEFT2D = torch.empty((M, H), device=x2d.device, dtype=torch.float32)
- RIGHT2D = torch.empty((M, H), device=x2d.device, dtype=torch.float32)
- OG2D = torch.empty((M, H), device=x2d.device, dtype=torch.float32)
+ if all_ones:
+ LEFT_T = torch.sigmoid(LGpre_T) * Lpre_T
+ else:
+ mrow = mask.to(torch.float32).view(1, M)
+ LEFT_T = torch.sigmoid(LGpre_T) * Lpre_T * mrow
+ RIGHT_T = torch.sigmoid(RGpre_T) * Rpre_T
+ OG_T = torch.sigmoid(OGpre_T)
- # Launch fused projections (small tiles to fit SMEM)
- grid_proj = (triton.cdiv(M, 64), triton.cdiv(H, 64))
- proj5_gated_mask_kernel[grid_proj](
- x2d, LW, RW, LGW, RGW, OGW, mask_f,
- LEFT2D, RIGHT2D, OG2D,
- M, D, H,
- x2d.stride(0), x2d.stride(1),
- LW.stride(0), LW.stride(1),
- LEFT2D.stride(0), LEFT2D.stride(1),
- BLOCK_M=64, BLOCK_N=64, BLOCK_K=32,
- num_warps=4, num_stages=2,
- )
+ # Views as [H, B, N, N] (no copies)
+ LEFT_HBNN = LEFT_T.view(H, B, N, N)
+ RIGHT_HBNN = RIGHT_T.view(H, B, N, N)
+ OG_HBNN = OG_T.view(H, B, N, N)
- LEFT = LEFT2D.view(B, N, N, H)
- RIGHT = RIGHT2D.view(B, N, N, H)
- OG = OG2D.view(B, N, N, H)
+ # 3) PASS 1: accumulate mean/var over H (no EIN/G materialization)
+ S = _get(("S", B, N, N, device), (B, N, N), torch.float32, device); S.zero_()
+ S2 = _get(("S2", B, N, N, device), (B, N, N), torch.float32, device); S2.zero_()
- # Contraction via batched GEMM over (b,h): for each h, L[i,k] @ R[j,k]^T -> [i,j]
- Left_h = LEFT.permute(0, 3, 1, 2).contiguous().view(B * H, N, N)
- Right_h = RIGHT.permute(0, 3, 2, 1).contiguous().view(B * H, N, N)
- EIN_h = torch.bmm(Left_h, Right_h) # [B*H, N, N]
- EIN = EIN_h.view(B, H, N, N).permute(0, 2, 3, 1).contiguous()
+ for h0 in range(0, H, TH):
+ h1 = min(H, h0 + TH); t = h1 - h0
+ Lbt = LEFT_HBNN [h0:h1].view(t * B, N, N)
+ if lhs_contig: Lbt = Lbt.contiguous()
+ Rbt = RIGHT_HBNN[h0:h1].view(t * B, N, N)
+ Cbt = torch.bmm(Lbt, Rbt.transpose(1, 2)) # [B*t, N, N]
+ C = Cbt.view(t, B, N, N)
+ S += C.sum(dim=0)
+ S2 += (C * C).sum(dim=0)
+ Hf = float(H)
+ mean = S / Hf
+ var = S2 / Hf - mean * mean
+ inv_std = torch.rsqrt(var + 1e-5) # [B, N, N]
- # Epilogue split: (1) Triton LN+gate per (b,i,j,h) -> G; (2) cuBLAS GEMM G @ W^T
- W = weights['to_out.weight' ].contiguous() # [D,H]
- LNw = weights['to_out_norm.weight' ].contiguous() # [H]
- LNb = weights['to_out_norm.bias' ].contiguous() # [H]
+ # 4) PASS 2: recompute contraction chunks, apply LN(g), accumulate into OUT
+ Wt_key = "__to_out_wT__" # [H, D]
+ Wt_full = weights.get(Wt_key)
+ if (Wt_full is None) or (Wt_full.shape != (H, D)) or (Wt_full.device != device):
+ Wt_full = weights['to_out.weight'].t().contiguous()
+ weights[Wt_key] = Wt_full
- G = torch.empty_like(OG) # [B,N,N,H]
- grid_epi = (B * N * N,)
- epilogue_ln_gate_kernel[grid_epi](
- EIN, OG, LNw, LNb, G,
- B, N, H,
- EIN.stride(0), EIN.stride(1), EIN.stride(2), EIN.stride(3),
- G.stride(0), G.stride(1), G.stride(2), G.stride(3),
- BLOCK_H=64,
- num_warps=4, num_stages=2,
- )
+ OUT2D = _get(("OUT2D", M, D, device), (M, D), torch.float32, device)
+ # Use beta=0 on first addmm to avoid a large memset
+ LNw = weights['to_out_norm.weight'] # [H]
+ LNb = weights['to_out_norm.bias'] # [H]
- M = B * N * N
- OUT2D = torch.matmul(G.view(M, H), W.t()) # [M,D]
- OUT = OUT2D.view(B, N, N, D)
- return OUT
+ mean_ = mean.unsqueeze(0) # [1, B, N, N]
+ inv_ = inv_std.unsqueeze(0)
+
+ first = True
+ for h0 in range(0, H, TH):
+ h1 = min(H, h0 + TH); t = h1 - h0
+ Lbt = LEFT_HBNN [h0:h1].view(t * B, N, N)
+ if lhs_contig: Lbt = Lbt.contiguous()
+ Rbt = RIGHT_HBNN[h0:h1].view(t * B, N, N)
+ Cbt = torch.bmm(Lbt, Rbt.transpose(1, 2)) # [B*t, N, N]
+ C = Cbt.view(t, B, N, N) # [t, B, N, N]
+
+ lnw = LNw[h0:h1].view(t, 1, 1, 1)
+ lnb = LNb[h0:h1].view(t, 1, 1, 1)
+ Cn = ((C - mean_) * inv_) * lnw + lnb
+
+ OGc = OG_HBNN[h0:h1] # [t, B, N, N]
+ G = Cn * OGc # [t, B, N, N]
+
+ GflatT = G.view(t, M) # [t, M]
+ Wt = Wt_full[h0:h1, :] # [t, D]
+ OUT2D.addmm_(GflatT.t(), Wt, beta=(0.0 if first else 1.0), alpha=1.0)
+ first = False
+
+ return OUT2D.view(B, N, N, D)
+
finally:
torch.backends.cuda.matmul.allow_tf32 = prev_tf32
if hasattr(torch, "set_float32_matmul_precision") and prev_prec is not None:
scrolls · 593 diff lines total

Best evidence level for this revision: reported

JSON