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
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 H100from utils import make_match_reference, DisableCuDNNTF32from task import input_t, output_t+ import os+ import json+ import mathimport torchimport torch.nn.functional as F- import triton- import triton.language as tl- import math# Keep harness globals unchangedtorch.backends.cuda.matmul.allow_tf32 = Truetorch.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 = dataB, N, _, D = input_tensor.shapeH = 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_tf32torch.backends.cuda.matmul.allow_tf32 = Trueprev_prec = torch.get_float32_matmul_precision() if hasattr(torch, "get_float32_matmul_precision") else Noneif 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_tf32if 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