"""Triton grouped GEMM (variable-length, expert-major) with fused SwiGLU epilogue. The MoE up-projection h_e = silu(x_e @ W_gate[e]) * (x_e @ W_up[e]) is computed by a single fused kernel. Both projections share the activation tile, so they are evaluated as one wide GEMM over a *column-interleaved* weight layout W_cat[e][k][2*j] = W_gate[e][k][j] W_cat[e][k][2*j+1] = W_up[e][k][j] which turns "two dots into two accumulators" into "one dot into one accumulator". That matters on this part: the fused kernel is register limited, and a single accumulator of the same footprint keeps the mainloop at one mma pipeline with a single B stream instead of two interleaved ones (measured ~6% faster than the two-accumulator formulation on the headline shape, and it matches the throughput of a plain single-output grouped GEMM). After the K loop the accumulator holds the gate and up values of the same output channel in adjacent columns. The m16n8k16 accumulator layout keeps consecutive column pairs in the same thread, so ``reshape -> split`` is a register rename, not a shared-memory round trip; the epilogue is then a plain elementwise ``silu(g) * u`` that is written straight to global memory. Neither GEMM result ever touches DRAM separately. Scheduling / variable-length M ------------------------------ The grid is ``(MP, I/BN, E)`` -- the expert index is the slowest-varying axis so the programs of one expert are launched back to back and keep that expert's weights hot in L2 while its activation rows stream through. M is not uniform across experts, so the kernel derives it from the routing metadata itself: every program reads ``expert_offsets[e]`` / ``expert_offsets[e+1]`` and walks ``mt = pid_m, pid_m + MP, ...`` up to ``ceil(nrows / BM)``. No host side knowledge of the routing is needed, hence no device synchronisation, and ragged/skewed routing (including empty experts, which are skipped) is handled exactly. Rows past the end of an expert are masked on the activation load and on the store, so no work is done for rows that belong to another expert. """ from __future__ import annotations import torch import torch.nn as nn import triton import triton.language as tl @triton.jit def _moe_up_swiglu_kernel( x_ptr, # (T_perm, H) bf16 w_ptr, # (E, H, 2*I) bf16, gate/up interleaved on the last axis out_ptr, # (T_perm, I) bf16 off_ptr, # (E + 1,) int32 prefix sums s_xm, s_xk, s_we, s_wk, s_wn, s_om, s_on, H: tl.constexpr, I: tl.constexpr, MP: tl.constexpr, BM: tl.constexpr, BN: tl.constexpr, # output (gate) channels per tile; weight tile is 2*BN wide BK: tl.constexpr, MASK_N: tl.constexpr, MASK_K: tl.constexpr, ): pid_m = tl.program_id(0) pid_n = tl.program_id(1) e = tl.program_id(2) start = tl.load(off_ptr + e) nrows = tl.load(off_ptr + e + 1) - start num_mt = tl.cdiv(nrows, BM) offs_m = tl.arange(0, BM) offs_w = pid_n * 2 * BN + tl.arange(0, 2 * BN) offs_k = tl.arange(0, BK) wb = w_ptr + e * s_we w_off = offs_k[:, None] * s_wk + offs_w[None, :] * s_wn for mt in range(pid_m, num_mt, MP): rm = mt * BM + offs_m row_ok = rm < nrows a_ptrs = x_ptr + (start + rm)[:, None] * s_xm + offs_k[None, :] * s_xk b_ptrs = wb + w_off acc = tl.zeros((BM, 2 * BN), dtype=tl.float32) for _k in range(0, tl.cdiv(H, BK)): if MASK_K: km = offs_k[None, :] < H - _k * BK a = tl.load(a_ptrs, mask=row_ok[:, None] & km, other=0.0) else: a = tl.load(a_ptrs, mask=row_ok[:, None], other=0.0) if MASK_N: b = tl.load(b_ptrs, mask=offs_w[None, :] < 2 * I, other=0.0) else: b = tl.load(b_ptrs) acc = tl.dot(a, b, acc) a_ptrs += BK * s_xk b_ptrs += BK * s_wk pair = tl.reshape(acc, (BM, BN, 2)) gate, up = tl.split(pair) y = (gate * tl.sigmoid(gate) * up).to(tl.bfloat16) col = pid_n * BN + tl.arange(0, BN) c_ptrs = out_ptr + (start + rm)[:, None] * s_om + col[None, :] * s_on if MASK_N: tl.store(c_ptrs, y, mask=row_ok[:, None] & (col[None, :] < I)) else: tl.store(c_ptrs, y, mask=row_ok[:, None]) # (H, I, E) -> (MP, BM, BN, BK, num_warps, num_stages) # # Every entry keeps BM * 2*BN == 16384 fp32 accumulator slots -- the largest # footprint that stays off the spill path on SM120 (254 regs, 0 spills). Larger # tiles spill and lose 30-60%; the alternative aspect ratio at the same # footprint (BM=256, BN=32) is 7% slower, so M-heavy tiles are the wrong trade. # # Stage count is *not* monotonic: 3 beats 4 by ~4% on the headline shape while # the ordering flips on the wide one. Deep pipelines only pay while the K loop # is long relative to the tile; past that they just add L2 pressure. _TUNED = { (4096, 1536, 128): (8, 128, 64, 32, 4, 3), (2048, 1024, 64): (2, 128, 64, 32, 4, 3), (2048, 4096, 64): (16, 128, 64, 32, 4, 4), } def _default_config(H: int, I: int, E: int) -> tuple[int, int, int, int, int, int]: for bn in (64, 32, 16, 8, 4, 2, 1): if I % bn == 0: break for bk in (32, 16, 8, 4, 2, 1): if H % bk == 0: break return (8, 128, bn, bk, 4, 3) class Model(nn.Module): """Up-projection of a top-K MoE FFN with fused SwiGLU (grouped GEMM).""" def __init__(self, T_total: int, H: int, I: int, E: int, K: int): # noqa: E741 super().__init__() self.T_total = T_total self.H = H self.I = I self.E = E self.K = K self.W_gate = nn.Parameter(torch.empty(E, H, I, dtype=torch.bfloat16)) self.W_up = nn.Parameter(torch.empty(E, H, I, dtype=torch.bfloat16)) nn.init.normal_(self.W_gate, std=0.02) nn.init.normal_(self.W_up, std=0.02) self._packed = None # (W_gate, W_up, ver_g, ver_u, W_cat) # -- weight packing --------------------------------------------------- def _packed_weight(self, wg: torch.Tensor, wu: torch.Tensor) -> torch.Tensor: """Interleave gate/up along the last axis. Cached across calls: the packing is a pure function of the two weight tensors, so it is rebuilt only when they are replaced or written to. Holding references to the source tensors (rather than just their data_ptrs) is what makes the identity check safe -- their storage can never be freed and recycled underneath us. """ cacheable = wg.is_contiguous() and wu.is_contiguous() cached = self._packed if (cacheable and cached is not None and cached[0] is wg and cached[1] is wu and cached[2] == wg._version and cached[3] == wu._version): return cached[4] E, H, I = wg.shape w_cat = torch.empty(E, H, 2 * I, dtype=wg.dtype, device=wg.device) view = w_cat.view(E, H, I, 2) view[..., 0].copy_(wg) # strided copy: handles non-contiguous sources view[..., 1].copy_(wu) if cacheable: self._packed = (wg, wu, wg._version, wu._version, w_cat) return w_cat def forward( self, hidden_states: torch.Tensor, expert_offsets: torch.Tensor, ) -> torch.Tensor: T_perm, H = hidden_states.shape I = self.I if hidden_states.stride(1) != 1: hidden_states = hidden_states.contiguous() wg = self.W_gate wu = self.W_up if expert_offsets.dtype != torch.int32: expert_offsets = expert_offsets.to(torch.int32) w_cat = self._packed_weight(wg, wu) out = torch.empty(T_perm, I, dtype=torch.bfloat16, device=hidden_states.device) cfg = _TUNED.get((H, I, self.E)) if cfg is None: cfg = _default_config(H, I, self.E) MP, BM, BN, BK, warps, stages = cfg grid = (MP, triton.cdiv(I, BN), self.E) _moe_up_swiglu_kernel[grid]( hidden_states, w_cat, out, expert_offsets, hidden_states.stride(0), hidden_states.stride(1), w_cat.stride(0), w_cat.stride(1), w_cat.stride(2), out.stride(0), out.stride(1), H=H, I=I, MP=MP, BM=BM, BN=BN, BK=BK, MASK_N=(I % BN) != 0, MASK_K=(H % BK) != 0, num_warps=warps, num_stages=stages, ) return out