"""Grouped GEMM + fused SwiGLU for top-K MoE FFN up-projection (H100 / SM90). Per expert e: h_e = silu(x_e @ W_gate[e]) * (x_e @ W_up[e]) Triton dual-GEMM kernel with TMA loads: - variable-length experts scheduled via cached tile maps - A (hidden) loaded once per K-tile; both gate and up accumulate - SiLU * mul fused into the epilogue (fp32 accumulate -> bf16 store) """ from __future__ import annotations import math from typing import Optional import torch import torch.nn as nn import triton import triton.language as tl OP_TYPE = "grouped_gemm_swiglu" SUPPORTED_PRECISIONS = ["bf16"] HARDWARE_REQUIRED = ["RTX_PRO_6000", "H100", "B200"] # TMA kernels need a runtime allocator for descriptor scratch. def _triton_allocator(size: int, alignment: int, stream: int | None): return torch.empty(size, device="cuda", dtype=torch.int8) try: triton.set_allocator(_triton_allocator) except Exception: pass @triton.jit def _dual_gemm_swiglu_tma_kernel( x_ptr, w_gate_ptr, w_up_ptr, out_ptr, expert_offsets_ptr, tile_expert_ptr, tile_m_ptr, tile_n_ptr, H, I, # noqa: E741 stride_xm, stride_we, stride_wk, stride_om, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, ): pid = tl.program_id(0) e = tl.load(tile_expert_ptr + pid) pid_m = tl.load(tile_m_ptr + pid) pid_n = tl.load(tile_n_ptr + pid) row_start = tl.load(expert_offsets_ptr + e) row_end = tl.load(expert_offsets_ptr + e + 1) n_tokens = row_end - row_start if n_tokens <= 0: return x_base = x_ptr + row_start.to(tl.int64) * stride_xm out_base = out_ptr + row_start.to(tl.int64) * stride_om w_gate_base = w_gate_ptr + e.to(tl.int64) * stride_we w_up_base = w_up_ptr + e.to(tl.int64) * stride_we a_desc = tl.make_tensor_descriptor( x_base, shape=[n_tokens, H], strides=[stride_xm, 1], block_shape=[BLOCK_M, BLOCK_K], ) bg_desc = tl.make_tensor_descriptor( w_gate_base, shape=[H, I], strides=[stride_wk, 1], block_shape=[BLOCK_K, BLOCK_N], ) bu_desc = tl.make_tensor_descriptor( w_up_base, shape=[H, I], strides=[stride_wk, 1], block_shape=[BLOCK_K, BLOCK_N], ) d_desc = tl.make_tensor_descriptor( out_base, shape=[n_tokens, I], strides=[stride_om, 1], block_shape=[BLOCK_M, BLOCK_N], ) offs_m = pid_m * BLOCK_M offs_n = 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 tl.range(0, H, BLOCK_K): a = a_desc.load([offs_m, k]) bg = bg_desc.load([k, offs_n]) bu = bu_desc.load([k, offs_n]) acc_g = tl.dot(a, bg, acc_g) acc_u = tl.dot(a, bu, acc_u) # SiLU(gate) * up out = (acc_g * tl.sigmoid(acc_g) * acc_u).to(tl.bfloat16) d_desc.store([offs_m, offs_n], out) @triton.jit def _dual_gemm_swiglu_kernel( x_ptr, w_gate_ptr, w_up_ptr, out_ptr, expert_offsets_ptr, tile_expert_ptr, tile_m_ptr, tile_n_ptr, H, I, # noqa: E741 stride_xm, stride_xk, stride_we, stride_wk, stride_wn, stride_om, stride_on, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, ): pid = tl.program_id(0) e = tl.load(tile_expert_ptr + pid) pid_m = tl.load(tile_m_ptr + pid) pid_n = tl.load(tile_n_ptr + pid) row_start = tl.load(expert_offsets_ptr + e) row_end = tl.load(expert_offsets_ptr + e + 1) n_tokens = row_end - row_start if n_tokens == 0: return offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) offs_k = tl.arange(0, BLOCK_K) x_base = x_ptr + row_start.to(tl.int64) * stride_xm w_gate_base = w_gate_ptr + e.to(tl.int64) * stride_we w_up_base = w_up_ptr + e.to(tl.int64) * stride_we out_base = out_ptr + row_start.to(tl.int64) * stride_om acc_g = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) acc_u = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) mask_m = offs_m < n_tokens mask_n = offs_n < I for k in tl.range(0, H, BLOCK_K): a_ptrs = x_base + ( offs_m[:, None].to(tl.int64) * stride_xm + (k + offs_k)[None, :].to(tl.int64) * stride_xk ) a = tl.load( a_ptrs, mask=mask_m[:, None] & ((k + offs_k)[None, :] < H), other=0.0, ) b_mask = ((k + offs_k)[:, None] < H) & mask_n[None, :] bg = tl.load( w_gate_base + (k + offs_k)[:, None].to(tl.int64) * stride_wk + offs_n[None, :].to(tl.int64) * stride_wn, mask=b_mask, other=0.0, ) bu = tl.load( w_up_base + (k + offs_k)[:, None].to(tl.int64) * stride_wk + offs_n[None, :].to(tl.int64) * stride_wn, mask=b_mask, other=0.0, ) acc_g = tl.dot(a, bg, acc_g) acc_u = tl.dot(a, bu, acc_u) out = (acc_g * tl.sigmoid(acc_g) * acc_u).to(tl.bfloat16) tl.store( out_base + offs_m[:, None].to(tl.int64) * stride_om + offs_n[None, :].to(tl.int64) * stride_on, out, mask=mask_m[:, None] & mask_n[None, :], ) # Best configs found by sweep on H100 PCIe (dual TMA). _CONFIGS = [ (128, 128, 64, 8, 4, True), (128, 128, 64, 8, 3, True), (128, 128, 32, 8, 5, True), (256, 64, 64, 8, 4, True), (128, 64, 64, 4, 3, True), (64, 128, 64, 4, 5, True), (128, 64, 64, 4, 4, True), (64, 64, 64, 4, 4, True), (128, 128, 64, 8, 4, False), (128, 64, 64, 4, 4, False), ] _best_cfg: dict = {} _tile_cache: dict = {} def _build_tile_maps( expert_offsets: torch.Tensor, I: int, # noqa: E741 BLOCK_M: int, BLOCK_N: int, ): """Build (expert, m_tile, n_tile) maps. Cached by offsets content.""" # Fast path: if offsets look balanced, build without full bytes hash every time # using (ptr, last_value, E, BM, BN, I) — invalidated when offsets change. off_cpu = expert_offsets.detach().to(device="cpu", dtype=torch.int32).contiguous() key = (off_cpu.numpy().tobytes(), I, BLOCK_M, BLOCK_N, str(expert_offsets.device)) cached = _tile_cache.get(key) if cached is not None: return cached E = off_cpu.numel() - 1 experts, ms, ns = [], [], [] num_n = (I + BLOCK_N - 1) // BLOCK_N # L2-friendly order: for each expert, iterate n-tiles outer so weight tiles # stay hot while walking m (A changes, B fixed for a given n-tile). for e in range(E): n_tok = int(off_cpu[e + 1]) - int(off_cpu[e]) if n_tok <= 0: continue num_m = (n_tok + BLOCK_M - 1) // BLOCK_M for pid_n in range(num_n): for pid_m in range(num_m): experts.append(e) ms.append(pid_m) ns.append(pid_n) device = expert_offsets.device if not experts: empty = torch.empty(0, dtype=torch.int32, device=device) result = (empty, empty, empty, 0) else: result = ( torch.tensor(experts, dtype=torch.int32, device=device), torch.tensor(ms, dtype=torch.int32, device=device), torch.tensor(ns, dtype=torch.int32, device=device), len(experts), ) if len(_tile_cache) > 32: _tile_cache.clear() _tile_cache[key] = result return result def _can_use_tma(hs, wg, I, BLOCK_M, BLOCK_N, BLOCK_K) -> bool: # noqa: E741 if hs.stride(-1) != 1 or wg.stride(-1) != 1: return False if (hs.stride(0) % 8) != 0 or (wg.stride(1) % 8) != 0: return False if (I % 8) != 0: return False if (BLOCK_K % 8) != 0 or (BLOCK_N % 8) != 0 or (BLOCK_M % 8) != 0: return False return True def grouped_gemm_swiglu( hidden_states: torch.Tensor, W_gate: torch.Tensor, W_up: torch.Tensor, expert_offsets: torch.Tensor, cfg: Optional[tuple] = None, ) -> torch.Tensor: T_perm, H = hidden_states.shape E, H_w, I = W_gate.shape # noqa: E741 assert H_w == H and W_up.shape == (E, H, I) if cfg is None: avg = max(int(T_perm) // max(E, 1), 1) if avg >= 1024 and I >= 1024: cfg = (128, 128, 64, 8, 4, True) elif avg >= 256: cfg = (128, 64, 64, 4, 3, True) else: cfg = (64, 64, 64, 4, 4, True) BLOCK_M, BLOCK_N, BLOCK_K, num_warps, num_stages, use_tma = cfg tile_expert, tile_m, tile_n, total_tiles = _build_tile_maps( expert_offsets, I, BLOCK_M, BLOCK_N ) out = torch.empty(T_perm, I, dtype=torch.bfloat16, device=hidden_states.device) if total_tiles == 0: return out hs = hidden_states if hidden_states.is_contiguous() else hidden_states.contiguous() wg = W_gate if W_gate.is_contiguous() else W_gate.contiguous() wu = W_up if W_up.is_contiguous() else W_up.contiguous() eo = expert_offsets if expert_offsets.is_contiguous() else expert_offsets.contiguous() use_tma = bool(use_tma) and _can_use_tma(hs, wg, I, BLOCK_M, BLOCK_N, BLOCK_K) if use_tma: _dual_gemm_swiglu_tma_kernel[(total_tiles,)]( hs, wg, wu, out, eo, tile_expert, tile_m, tile_n, H, I, hs.stride(0), wg.stride(0), wg.stride(1), out.stride(0), BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K, num_warps=num_warps, num_stages=num_stages, ) else: _dual_gemm_swiglu_kernel[(total_tiles,)]( hs, wg, wu, out, eo, tile_expert, tile_m, tile_n, H, I, hs.stride(0), hs.stride(1), wg.stride(0), wg.stride(1), wg.stride(2), out.stride(0), out.stride(1), BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K, num_warps=num_warps, num_stages=num_stages, ) return out def _autotune_once(hidden_states, W_gate, W_up, expert_offsets) -> tuple: T_perm, H = hidden_states.shape E, _, I = W_gate.shape # noqa: E741 avg = max(T_perm // max(E, 1), 1) bucket = 1 << int(math.floor(math.log2(max(avg, 1)))) key = (H, I, E, bucket) if key in _best_cfg: return _best_cfg[key] ok = [] for c in _CONFIGS: try: grouped_gemm_swiglu(hidden_states, W_gate, W_up, expert_offsets, cfg=c) torch.cuda.synchronize() ok.append(c) except Exception: continue if not ok: best = (64, 64, 64, 4, 3, False) _best_cfg[key] = best return best best = ok[0] best_ms = float("inf") for c in ok: try: start = torch.cuda.Event(enable_timing=True) end = torch.cuda.Event(enable_timing=True) torch.cuda.synchronize() start.record() for _ in range(10): grouped_gemm_swiglu(hidden_states, W_gate, W_up, expert_offsets, cfg=c) end.record() torch.cuda.synchronize() ms = start.elapsed_time(end) / 10.0 if ms < best_ms: best_ms = ms best = c except Exception: continue _best_cfg[key] = best return best class Model(nn.Module): 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._tuned = False self._cfg = None def forward( self, hidden_states: torch.Tensor, expert_offsets: torch.Tensor, ) -> torch.Tensor: if not self._tuned and hidden_states.is_cuda: try: triton.set_allocator(_triton_allocator) except Exception: pass self._cfg = _autotune_once( hidden_states, self.W_gate, self.W_up, expert_offsets ) self._tuned = True return grouped_gemm_swiglu( hidden_states, self.W_gate, self.W_up, expert_offsets, cfg=self._cfg ) 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]