"""Sonic-MoE up-projection: variable-length grouped GEMM with fused SwiGLU. Per expert e (rows [offsets[e], offsets[e+1]) of the permuted hidden states): h_e = silu(x_e @ W_gate[e]) * (x_e @ W_up[e]) Design ------ * One *persistent* Triton kernel (grid = #SMs) walks a flat list of output tiles. A tiny prologue kernel turns `expert_offsets` into a (row_start, n_rows, expert) tile table on device, so the ragged group boundaries never require a host synchronisation. * Each tile keeps **two** fp32 accumulators (gate and up). The A-tile is fetched once and feeds two wgmma chains, i.e. the pair is one GEMM of N = 2*I with the SwiGLU folded into the epilogue. 128x128 with 8 warps is the largest tile that still fits 2 x BM x BN / 256 = 128 accumulator registers per thread without spilling. * Operands are moved with Hopper TMA (`cp.async.bulk.tensor`). Besides the cheaper address math this makes the mainloop *mask-free*: TMA zero-fills out-of-bounds elements, so the ragged last tile of an expert and a non-multiple-of-BLOCK_K reduction dim both come out right (garbage rows are dropped by the masked store; garbage k-rows of B are multiplied by zero-filled columns of A). The weights are described as one 2D (E*H, I) tensor -- legal because W[e] is contiguous -- with row e*H + k. """ from __future__ import annotations import torch import torch.nn as nn import triton import triton.language as tl try: # Hopper tensor-memory-accelerator descriptors from triton.tools.experimental_descriptor import create_2d_tma_descriptor as _mk_desc except Exception: # pragma: no cover - falls back to the cp.async kernel _mk_desc = None TILE_M = 128 # M-tiling shared by the table builder and the GEMM _NEG_LOG2E = tl.constexpr(-1.4426950408889634) # --------------------------------------------------------------------------- # # Prologue: expert_offsets -> flat tile table (one program per expert) # --------------------------------------------------------------------------- # @triton.jit def _tile_table_kernel( offsets_ptr, # (E+1,) int32 tile_row_ptr, # (MAX_TILES,) int32 first row of the tile tile_nrow_ptr, # (MAX_TILES,) int32 valid rows in the tile tile_exp_ptr, # (MAX_TILES,) int32 expert id nm_ptr, # (1,) int32 total number of m-tiles E: tl.constexpr, E_POW2: tl.constexpr, BLOCK_M: tl.constexpr, CHUNK: tl.constexpr, ): e = tl.program_id(0) ar = tl.arange(0, E_POW2) v = ar < E s = tl.load(offsets_ptr + ar, mask=v, other=0) t = tl.load(offsets_ptr + ar + 1, mask=v, other=0) cnt = tl.where(v, t - s, 0) ntl = (cnt + BLOCK_M - 1) // BLOCK_M if e == 0: tl.store(nm_ptr, tl.sum(ntl, axis=0)) # Exclusive prefix sum, recomputed per program: E is small, this is free. base = tl.sum(tl.where(ar < e, ntl, 0), axis=0) my_cnt = tl.sum(tl.where(ar == e, cnt, 0), axis=0) my_row = tl.sum(tl.where(ar == e, s, 0), axis=0) my_ntl = (my_cnt + BLOCK_M - 1) // BLOCK_M for j0 in range(0, my_ntl, CHUNK): j = j0 + tl.arange(0, CHUNK) m = j < my_ntl tl.store(tile_row_ptr + base + j, my_row + j * BLOCK_M, mask=m) tl.store(tile_nrow_ptr + base + j, tl.minimum(my_cnt - j * BLOCK_M, BLOCK_M), mask=m) tl.store(tile_exp_ptr + base + j, j * 0 + e, mask=m) @triton.jit def _swiglu(g, u): """silu(g) * u = g*u / (1 + exp(-g)). ex2.approx + div.full are both *relative*-error accurate (~2^-22), which matters here: for g << 0 the true result decays like g*u*exp(g), and a formulation with only absolute accuracy (e.g. tanh.approx) loses those elements entirely once |g*u| is large. """ return tl.fdiv(g * u, 1.0 + tl.exp2(g * _NEG_LOG2E), ieee_rounding=False) # --------------------------------------------------------------------------- # # Main kernel: TMA operands, dual accumulator, fused SwiGLU epilogue # --------------------------------------------------------------------------- # @triton.jit def _moe_up_swiglu_tma( a_desc, wg_desc, wu_desc, OUT, tile_row_ptr, tile_nrow_ptr, tile_exp_ptr, nm_ptr, H: tl.constexpr, I: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, GROUP_M: tl.constexpr, ): pid = tl.program_id(0) nprog = tl.num_programs(0) n_tiles: tl.constexpr = (I + BLOCK_N - 1) // BLOCK_N num_m = tl.load(nm_ptr) total = num_m * n_tiles pid_in_group: tl.constexpr = GROUP_M * n_tiles rel = tl.arange(0, BLOCK_M) cn = tl.arange(0, BLOCK_N) for tid in range(pid, total, nprog): # ---- flat tile id -> (m-tile, n-tile) with an L2-friendly swizzle ---- group_id = tid // pid_in_group first_m = group_id * GROUP_M gsm = tl.minimum(num_m - first_m, GROUP_M) r = tid % pid_in_group pid_m = first_m + (r % gsm) pid_n = r // gsm row0 = tl.load(tile_row_ptr + pid_m) nrow = tl.load(tile_nrow_ptr + pid_m) wrow = tl.load(tile_exp_ptr + pid_m) * H n0 = pid_n * BLOCK_N acc_g = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) acc_u = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) for k in range(0, H, BLOCK_K): a = tl._experimental_descriptor_load( a_desc, [row0, k], [BLOCK_M, BLOCK_K], tl.bfloat16) bg = tl._experimental_descriptor_load( wg_desc, [wrow + k, n0], [BLOCK_K, BLOCK_N], tl.bfloat16) bu = tl._experimental_descriptor_load( wu_desc, [wrow + k, n0], [BLOCK_K, BLOCK_N], tl.bfloat16) acc_g = tl.dot(a, bg, acc_g) acc_u = tl.dot(a, bu, acc_u) o = _swiglu(acc_g, acc_u).to(OUT.dtype.element_ty) offs_n = n0 + cn o_ptrs = OUT + row0.to(tl.int64) * I + (rel[:, None] * I + offs_n[None, :]) tl.store(o_ptrs, o, mask=(rel < nrow)[:, None] & (offs_n < I)[None, :]) # --------------------------------------------------------------------------- # # Portable fallback: same schedule, cp.async / pointer operands # --------------------------------------------------------------------------- # @triton.jit def _moe_up_swiglu_ptr( X, WG, WU, OUT, tile_row_ptr, tile_nrow_ptr, tile_exp_ptr, nm_ptr, H: tl.constexpr, I: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, GROUP_M: tl.constexpr, ): pid = tl.program_id(0) nprog = tl.num_programs(0) n_tiles: tl.constexpr = (I + BLOCK_N - 1) // BLOCK_N num_m = tl.load(nm_ptr) total = num_m * n_tiles pid_in_group: tl.constexpr = GROUP_M * n_tiles ak = tl.arange(0, BLOCK_K) rel = tl.arange(0, BLOCK_M) cn = tl.arange(0, BLOCK_N) for tid in range(pid, total, nprog): group_id = tid // pid_in_group first_m = group_id * GROUP_M gsm = tl.minimum(num_m - first_m, GROUP_M) r = tid % pid_in_group pid_m = first_m + (r % gsm) pid_n = r // gsm row0 = tl.load(tile_row_ptr + pid_m) nrow = tl.load(tile_nrow_ptr + pid_m) eidx = tl.load(tile_exp_ptr + pid_m) # Index expressions stay affine in `arange` so Triton keeps the # contiguity facts it needs for vectorised cp.async; the ragged edges # are handled with loop-invariant masks instead of clamped pointers. n0 = pid_n * BLOCK_N offs_n = n0 + cn mask_m = rel < nrow mask_n = offs_n < I a_ptrs = X + row0.to(tl.int64) * H + (rel[:, None] * H + ak[None, :]) woff = eidx.to(tl.int64) * (H * I) + (ak[:, None] * I + offs_n[None, :]) bg_ptrs = WG + woff bu_ptrs = WU + woff acc_g = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) acc_u = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) for k0 in range(0, H, BLOCK_K): mask_k = ak < H - k0 a = tl.load(a_ptrs, mask=mask_m[:, None] & mask_k[None, :], other=0.0) bg = tl.load(bg_ptrs, mask=mask_k[:, None] & mask_n[None, :], other=0.0) bu = tl.load(bu_ptrs, mask=mask_k[:, None] & mask_n[None, :], other=0.0) acc_g = tl.dot(a, bg, acc_g) acc_u = tl.dot(a, bu, acc_u) a_ptrs += BLOCK_K bg_ptrs += BLOCK_K * I bu_ptrs += BLOCK_K * I o = _swiglu(acc_g, acc_u).to(OUT.dtype.element_ty) o_ptrs = OUT + row0.to(tl.int64) * I + (rel[:, None] * I + offs_n[None, :]) tl.store(o_ptrs, o, mask=mask_m[:, None] & mask_n[None, :]) # --------------------------------------------------------------------------- # # Host side # --------------------------------------------------------------------------- # _NUM_SMS: int | None = None _SMEM_CAP = 200 * 1024 def _num_sms() -> int: global _NUM_SMS if _NUM_SMS is None: _NUM_SMS = torch.cuda.get_device_properties(0).multi_processor_count return _NUM_SMS def _pow2_le(n: int, lo: int) -> int: return max(lo, 1 << (max(n, 1).bit_length() - 1)) class Model(nn.Module): """Up-projection of a top-K MoE FFN with fused SwiGLU.""" 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._tbl: dict = {} self._desc: dict = {} # -- tile-table scratch buffers, one set per (T_perm, BLOCK_M) ---------- # def _table(self, T_perm: int, bm: int, device) -> tuple: key = (T_perm, bm) t = self._tbl.get(key) if t is None: n = (T_perm + bm - 1) // bm + self.E + 1 t = tuple(torch.empty(n, dtype=torch.int32, device=device) for _ in range(3)) + (torch.empty(1, dtype=torch.int32, device=device),) self._tbl[key] = t return t # -- TMA descriptors: content-independent, so (ptr, dims, box) is a key - # def _descriptor(self, t: torch.Tensor, dim1: int, dim0: int, box1: int, box0: int): key = (t.data_ptr(), dim1, dim0, box1, box0) d = self._desc.get(key) if d is None: if len(self._desc) > 96: self._desc.clear() d = _mk_desc(t.data_ptr(), dim1, dim0, box1, box0, t.element_size()) self._desc[key] = d return d def forward( self, hidden_states: torch.Tensor, # (T_perm, H) bf16 expert_offsets: torch.Tensor, # (E+1,) int32 ) -> torch.Tensor: T_perm, H = hidden_states.shape I = self.I # noqa: E741 E = self.E dev = hidden_states.device out = torch.empty((T_perm, I), dtype=torch.bfloat16, device=dev) if T_perm == 0 or I == 0: return out x = hidden_states if hidden_states.is_contiguous() else hidden_states.contiguous() wg, wu = self.W_gate, self.W_up if expert_offsets.device != dev: expert_offsets = expert_offsets.to(dev, non_blocking=True) # ---- tile shape: 128x128x64 unless a dimension is smaller --------- # bm = TILE_M if T_perm >= TILE_M else _pow2_le(T_perm, 16) bn = 128 if I >= 128 else _pow2_le(I, 16) bk = 64 if H >= 64 else _pow2_le(H, 16) ns = 3 while ns > 1 and (bm * bk + 2 * bk * bn) * 2 * ns > _SMEM_CAP: ns -= 1 rows, nrows, exps, nm = self._table(T_perm, bm, dev) _tile_table_kernel[(E,)]( expert_offsets, rows, nrows, exps, nm, E=E, E_POW2=triton.next_power_of_2(E), BLOCK_M=bm, CHUNK=256, num_warps=4, ) grid = (_num_sms(),) # TMA needs 16B-aligned rows in both operands (element size is 2B). use_tma = ( _mk_desc is not None and torch.cuda.get_device_capability(dev)[0] >= 9 and H % 8 == 0 and I % 8 == 0 and bk % 8 == 0 and bn % 8 == 0 and wg.is_contiguous() and wu.is_contiguous() ) if use_tma: try: a_desc = self._descriptor(x, T_perm, H, bm, bk) wg_desc = self._descriptor(wg, E * H, I, bk, bn) wu_desc = self._descriptor(wu, E * H, I, bk, bn) except Exception: use_tma = False if use_tma: _moe_up_swiglu_tma[grid]( a_desc, wg_desc, wu_desc, out, rows, nrows, exps, nm, H, I, BLOCK_M=bm, BLOCK_N=bn, BLOCK_K=bk, GROUP_M=8, num_warps=8, num_stages=ns, ) else: _moe_up_swiglu_ptr[grid]( x, wg, wu, out, rows, nrows, exps, nm, H, I, BLOCK_M=bm, BLOCK_N=bn, BLOCK_K=min(bk, 32), GROUP_M=8, num_warps=8, num_stages=4, ) return out # --------------------------------------------------------------------------- # # Harness shims (mirrors reference.py) # --------------------------------------------------------------------------- # T_total = 32768 H = 4096 I = 1536 # noqa: E741 E = 128 K = 8 def _build_routing(T_total: int, E: int, K: int, device: str = "cpu") -> torch.Tensor: T_perm = T_total * K base = T_perm // E rem = T_perm - base * E counts = torch.full((E,), base, dtype=torch.int32, device=device) counts[:rem] += 1 offsets = torch.zeros(E + 1, dtype=torch.int32, device=device) offsets[1:] = torch.cumsum(counts, dim=0) return offsets def get_inputs(): T_perm = T_total * K hidden_states = torch.randn(T_perm, H, dtype=torch.bfloat16) * 0.1 expert_offsets = _build_routing(T_total, E, K) return [hidden_states, expert_offsets] def get_init_inputs(): return [T_total, H, I, E, K]