"""Grouped-GEMM + fused SwiGLU up-projection for top-K MoE (SM90 / H100). Per expert e: out_e = silu(x_e @ W_gate[e]) * (x_e @ W_up[e]), where x_e is a contiguous slice of the permuted hidden states given by expert_offsets. Design ------ * Single fused Triton kernel. Each CTA owns one (BM x BN) output tile of one expert and keeps TWO fp32 accumulators (gate and up). The shared activation tile is loaded once per K-step and feeds two tl.dot calls against the gate/up weight tiles, halving activation traffic vs. two separate GEMMs. SwiGLU runs in the epilogue straight from the accumulators; the gate product never touches HBM. * Work distribution: a tiny kernel computes the per-expert m-tile counts' prefix sum (cumulative ceil(n_e / BM)); each main-kernel CTA then finds its (expert, row) with a branchless binary search over that prefix sum. No host<->device syncs anywhere, and only two launches per forward call. * Optional persistent scheduling cooperates with a GROUP_M swizzle so concurrently-running CTAs share weight slices in L2 (weights dominate DRAM traffic). * TMA tensor descriptors for A/B/C give free boundary handling on loads and stores; ragged per-expert row tails fall back to a masked pointer store so the next expert's rows are never clobbered. """ from __future__ import annotations import torch import torch.nn as nn import triton import triton.language as tl from triton.tools.tensor_descriptor import TensorDescriptor _LOG_SMEM_LIMIT = 232448 # SM90 shared-memory budget per CTA in bytes _NUM_SMS: int | None = None # --------------------------------------------------------------------------- # tiny schedule kernel: cum[e] = sum_{i= E, power of 2 ): idx = tl.arange(0, BLOCK) mask = idx < E o0 = tl.load(offsets_ptr + idx, mask=mask, other=0) o1 = tl.load(offsets_ptr + idx + 1, mask=mask, other=0) t = (o1 - o0 + (BM - 1)) // BM c = tl.cumsum(t, axis=0) tl.store(cum_ptr + 1 + idx, c, mask=mask) zero = tl.zeros((), dtype=tl.int32) tl.store(cum_ptr, zero) total = tl.sum(t, axis=0) tl.store(total_ptr, total) # --------------------------------------------------------------------------- # main fused grouped-GEMM + SwiGLU kernel # --------------------------------------------------------------------------- @triton.jit def _find_expert(cum_ptr, t, E: tl.constexpr, LOG2E: tl.constexpr): """Largest e in [0, E-1] with cum[e] <= t. (cum has E+1 entries, cum[0]=0)""" pos = tl.zeros((), dtype=tl.int32) step: tl.constexpr = 1 << (LOG2E - 1) for _ in tl.static_range(LOG2E): nxt = pos + step ok = nxt <= E c = tl.load(cum_ptr + nxt, mask=ok, other=2147483647) pos = tl.where(ok & (c <= t), nxt, pos) step = step // 2 if step > 1 else 0 # type: ignore[const-overflow] return pos @triton.jit def _moe_up_swiglu_kernel( a_desc, g_desc, # W_gate flattened as (E*H, I) u_desc, # W_up flattened as (E*H, I) c_desc, c_ptr, # raw output pointer for the masked tail-store path offsets_ptr, # (E+1,) int32 cum_ptr, # (E+1,) int32 m-tile prefix sums total_ptr, # (1,) int32 total m-tiles H: tl.constexpr, I: tl.constexpr, E: tl.constexpr, LOG2E: tl.constexpr, NUM_N_TILES: tl.constexpr, BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr, GROUP_M: tl.constexpr, PERSISTENT: tl.constexpr, ): pid = tl.program_id(0) num_m_tiles = tl.load(total_ptr) if PERSISTENT: num_progs = tl.num_programs(0) total_tiles = num_m_tiles * NUM_N_TILES for tile_id in tl.range(pid, total_tiles, num_progs): _do_tile( a_desc, g_desc, u_desc, c_desc, c_ptr, offsets_ptr, cum_ptr, tile_id, num_m_tiles, H, I, E, LOG2E, NUM_N_TILES, BM, BN, BK, GROUP_M, ) else: num_in_group = GROUP_M * NUM_N_TILES if pid < num_in_group * ((num_m_tiles + GROUP_M - 1) // GROUP_M): _do_tile( a_desc, g_desc, u_desc, c_desc, c_ptr, offsets_ptr, cum_ptr, pid, num_m_tiles, H, I, E, LOG2E, NUM_N_TILES, BM, BN, BK, GROUP_M, ) @triton.jit def _do_tile( a_desc, g_desc, u_desc, c_desc, c_ptr, offsets_ptr, cum_ptr, pid, num_m_tiles, H: tl.constexpr, I: tl.constexpr, E: tl.constexpr, LOG2E: tl.constexpr, NUM_N_TILES: tl.constexpr, BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr, GROUP_M: tl.constexpr, ): # swizzled decode pid -> (m_tile, n_tile), GROUP_M m-tiles per group for L2 num_in_group = GROUP_M * NUM_N_TILES group_id = pid // num_in_group first_m = group_id * GROUP_M group_size = tl.minimum(num_m_tiles - first_m, GROUP_M) m_tile = first_m + (pid % num_in_group) % group_size n_tile = (pid % num_in_group) // group_size e = _find_expert(cum_ptr, m_tile, E, LOG2E) tiles_before = tl.load(cum_ptr + e) start = tl.load(offsets_ptr + e) e_end = tl.load(offsets_ptr + e + 1) m0 = start + (m_tile - tiles_before) * BM n0 = n_tile * BN w_row0 = e * H # W flattened as (E*H, I) acc_g = tl.zeros((BM, BN), dtype=tl.float32) acc_u = tl.zeros((BM, BN), dtype=tl.float32) for k in range(0, H, BK): a = a_desc.load([m0, k]) wg = g_desc.load([w_row0 + k, n0]) wu = u_desc.load([w_row0 + k, n0]) acc_g = tl.dot(a, wg, acc_g) acc_u = tl.dot(a, wu, acc_u) out = acc_u * (acc_g * tl.sigmoid(acc_g)) out_bf16 = out.to(tl.bfloat16) if m0 + BM <= e_end: c_desc.store([m0, n0], out_bf16) else: # ragged tail: mask off rows that belong to the next expert rows = m0 + tl.arange(0, BM) cols = n0 + tl.arange(0, BN) ptrs = c_ptr + rows.to(tl.int64)[:, None] * I + cols[None, :] mask = (rows < e_end)[:, None] & (cols < I)[None, :] tl.store(ptrs, out_bf16, mask=mask) # --------------------------------------------------------------------------- # host side # --------------------------------------------------------------------------- def _next_pow2(x: int) -> int: return max(1, 1 << (x - 1).bit_length()) # tuned on H100 PCIe (see scratch/); keyed by (H, I) _CONFIGS = { (4096, 1536): {"BM": 128, "BN": 128, "BK": 64, "GROUP_M": 8, "num_warps": 8, "num_stages": 3, "PERSISTENT": True, "GRID_MULT": 2}, (2048, 1024): {"BM": 128, "BN": 128, "BK": 64, "GROUP_M": 2, "num_warps": 8, "num_stages": 3, "PERSISTENT": True, "GRID_MULT": 2}, (2048, 4096): {"BM": 128, "BN": 128, "BK": 64, "GROUP_M": 8, "num_warps": 8, "num_stages": 3, "PERSISTENT": True, "GRID_MULT": 2}, } def _pick_config(H: int, I: int) -> dict: cfg = _CONFIGS.get((H, I)) if cfg is None: cfg = {"BM": 128, "BN": 128, "BK": 64, "GROUP_M": 8, "num_warps": 8, "num_stages": 3, "PERSISTENT": True} return cfg def _smem_fits(BM: int, BN: int, BK: int, num_stages: int) -> bool: per_stage = (BM * BK + 2 * BK * BN) * 2 # a tile + two weight tiles, bf16 return per_stage * num_stages < _LOG_SMEM_LIMIT def _run( hidden_states: torch.Tensor, expert_offsets: torch.Tensor, W_gate: torch.Tensor, W_up: torch.Tensor, E: int, ) -> torch.Tensor: T_perm, H = hidden_states.shape I = W_gate.shape[-1] device = hidden_states.device out = torch.empty(T_perm, I, dtype=torch.bfloat16, device=device) if T_perm == 0: return out offsets = expert_offsets if offsets.dtype != torch.int32: offsets = offsets.to(torch.int32) cfg = _pick_config(H, I) BM, BN, BK = cfg["BM"], cfg["BN"], cfg["BK"] num_stages = cfg["num_stages"] while not _smem_fits(BM, BN, BK, num_stages) and num_stages > 2: num_stages -= 1 cum = torch.empty(E + 1, dtype=torch.int32, device=device) total = torch.empty(1, dtype=torch.int32, device=device) BLOCK = max(_next_pow2(E), 16) _cum_tiles_kernel[(1,)](offsets, cum, total, E=E, BM=BM, BLOCK=BLOCK, num_warps=4) a_desc = TensorDescriptor.from_tensor(hidden_states, [BM, BK]) g_desc = TensorDescriptor.from_tensor(W_gate.view(E * H, I), [BK, BN]) u_desc = TensorDescriptor.from_tensor(W_up.view(E * H, I), [BK, BN]) c_desc = TensorDescriptor.from_tensor(out, [BM, BN]) num_n_tiles = triton.cdiv(I, BN) max_tiles = T_perm // BM + E if cfg["PERSISTENT"]: global _NUM_SMS if _NUM_SMS is None: _NUM_SMS = torch.cuda.get_device_properties(device).multi_processor_count grid = (min(_NUM_SMS * cfg.get("GRID_MULT", 2), max_tiles * num_n_tiles),) else: grid = (max_tiles * num_n_tiles,) _moe_up_swiglu_kernel[grid]( a_desc, g_desc, u_desc, c_desc, out, offsets, cum, total, H=H, I=I, E=E, LOG2E=_next_pow2(E).bit_length() - 1 if E > 1 else 1, NUM_N_TILES=num_n_tiles, BM=BM, BN=BN, BK=BK, GROUP_M=cfg["GROUP_M"], PERSISTENT=cfg["PERSISTENT"], num_warps=cfg["num_warps"], num_stages=num_stages, ) return out 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) def forward( self, hidden_states: torch.Tensor, # (T_perm, H) bf16 expert_offsets: torch.Tensor, # (E+1,) int32 ) -> torch.Tensor: hidden_states = hidden_states.contiguous() return _run( hidden_states, expert_offsets, self.W_gate, self.W_up, self.E, ) # --- module-level shape shims, mirroring 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]