"""Top-K MoE up-projection: variable-length grouped GEMM + fused SwiGLU. Per expert e with tokens in [offsets[e], offsets[e+1]): h_e = silu(x_e @ W_gate[e]) * (x_e @ W_up[e]) Gate/up weights are packed interleaved so a single TMA GEMM produces both halves and SwiGLU runs in-register in the epilogue. """ from __future__ import annotations from typing import Optional import torch import torch.nn as nn import triton import triton.language as tl from triton.tools.tensor_descriptor import TensorDescriptor OP_TYPE = "grouped_gemm_swiglu" SUPPORTED_PRECISIONS = ["bf16"] HARDWARE_REQUIRED = ["RTX_PRO_6000", "H100", "B200"] def _set_tma_allocator() -> None: def alloc_fn(size: int, alignment: int, stream: Optional[int]): return torch.empty(size, device="cuda", dtype=torch.int8) triton.set_allocator(alloc_fn) _set_tma_allocator() @triton.jit def _grouped_swiglu_tma_kernel( x_desc, w_desc, out_ptr, offsets_ptr, stride_om, H: tl.constexpr, I: tl.constexpr, # noqa: E741 BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, ): """One CTA per (m-tile, packed-n-tile, expert). W is interleaved [gate, up].""" pid_m = tl.program_id(0) pid_n = tl.program_id(1) e = tl.program_id(2) start = tl.load(offsets_ptr + e) end = tl.load(offsets_ptr + e + 1) n_e = end - start offs_m = pid_m * BLOCK_M if offs_m >= n_e: return offs_n = pid_n * BLOCK_N row0 = start + offs_m w_row0 = e * H acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) for ki in range(0, H // BLOCK_K): k = ki * BLOCK_K a = x_desc.load([row0, k]) b = w_desc.load([w_row0 + k, offs_n]) acc = tl.dot(a, b, acc) acc = tl.reshape(acc, (BLOCK_M, BLOCK_N // 2, 2)) gate, up = tl.split(acc) out = (gate * tl.sigmoid(gate) * up).to(tl.bfloat16) offs_om = row0 + tl.arange(0, BLOCK_M) offs_on = pid_n * (BLOCK_N // 2) + tl.arange(0, BLOCK_N // 2) mask = (offs_om[:, None] < end) & (offs_on[None, :] < I) tl.store(out_ptr + offs_om[:, None] * stride_om + offs_on[None, :], out, mask=mask) def _pack_interleaved(w_gate: torch.Tensor, w_up: torch.Tensor, dst: torch.Tensor) -> None: """dst[..., 0::2] = gate, dst[..., 1::2] = up. dst is (E, H, 2I).""" dst[..., 0::2] = w_gate dst[..., 1::2] = w_up 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.register_buffer( "_W_pack", torch.empty(E, H, 2 * I, dtype=torch.bfloat16), persistent=False, ) self._pack_ok = False self._out_buf: torch.Tensor | None = None self._cached_off_ptr = 0 self._cached_max_m = 0 self._desc_key: tuple | None = None self._x_desc = None self._w_desc = None self._graph = None self._graph_key: tuple | None = None def _ensure_pack(self) -> None: if self._pack_ok: return _pack_interleaved(self.W_gate.detach(), self.W_up.detach(), self._W_pack) self._pack_ok = True self._graph = None self._graph_key = None def _load_from_state_dict(self, *args, **kwargs): self._pack_ok = False self._graph = None self._graph_key = None return super()._load_from_state_dict(*args, **kwargs) def forward( self, hidden_states: torch.Tensor, expert_offsets: torch.Tensor, ) -> torch.Tensor: self._ensure_pack() if not hidden_states.is_contiguous(): hidden_states = hidden_states.contiguous() T_perm, H = hidden_states.shape I = self.I E = self.E device = hidden_states.device out = self._out_buf if out is None or out.shape[0] != T_perm or out.shape[1] != I or out.device != device: out = torch.empty(T_perm, I, dtype=torch.bfloat16, device=device) self._out_buf = out offsets = expert_offsets if not offsets.is_cuda: offsets = offsets.to(device, non_blocking=True) if not offsets.is_contiguous(): offsets = offsets.contiguous() off_ptr = int(offsets.data_ptr()) if off_ptr != self._cached_off_ptr: counts = offsets[1:] - offsets[:-1] self._cached_max_m = int(counts.max().item()) if E > 0 else 0 self._cached_off_ptr = off_ptr max_m = self._cached_max_m if max_m == 0: return out packed = self._W_pack # Tile pick is from SM120 sweeps. 128x256x32 keeps the tensor pipe # busiest on large intermediates; 128x128x32 / 4-stage has lower # register pressure and wins on the small-M expert case. if I >= 1536: BLOCK_M, BLOCK_N, BLOCK_K, num_stages = 128, 256, 32, 3 else: BLOCK_M, BLOCK_N, BLOCK_K, num_stages = 128, 128, 32, 4 desc_key = ( hidden_states.data_ptr(), T_perm, H, packed.data_ptr(), BLOCK_M, BLOCK_N, BLOCK_K, ) if desc_key != self._desc_key: self._x_desc = TensorDescriptor( hidden_states, [T_perm, H], [hidden_states.stride(0), hidden_states.stride(1)], [BLOCK_M, BLOCK_K], ) self._w_desc = TensorDescriptor( packed, [E * H, 2 * I], [packed.stride(1), packed.stride(2)], [BLOCK_K, BLOCK_N], ) self._desc_key = desc_key x_desc = self._x_desc w_desc = self._w_desc grid = ( triton.cdiv(max_m, BLOCK_M), triton.cdiv(2 * I, BLOCK_N), E, ) gkey = (hidden_states.data_ptr(), offsets.data_ptr(), out.data_ptr(), BLOCK_M, BLOCK_N, BLOCK_K, num_stages) def _launch(): _grouped_swiglu_tma_kernel[grid]( x_desc, w_desc, out, offsets, out.stride(0), H=H, I=I, BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K, num_warps=8, num_stages=num_stages, ) if self._graph is not None and gkey == self._graph_key: self._graph.replay() return out _launch() try: torch.cuda.synchronize() graph = torch.cuda.CUDAGraph() with torch.cuda.graph(graph): _launch() self._graph = graph self._graph_key = gkey except Exception: self._graph = None self._graph_key = None return out 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]