From 5177a3ec08548b3d0f1a9eea43d0b721959116d9 Mon Sep 17 00:00:00 2001 From: Chunan Zeng Date: Tue, 8 Sep 2026 16:39:41 -0700 Subject: [PATCH] MiniMax-M3: Triton split-K router GEMV with in-kernel fixup (#36557) Co-authored-by: Kevin Mi <45493463+kevin-mii@users.noreply.github.com> --- python/sglang/kernels/ops/gemm/router_gemv.py | 175 ++++++++++++++++++ python/sglang/srt/models/minimax_m3.py | 14 ++ 2 files changed, 189 insertions(+) create mode 100644 python/sglang/kernels/ops/gemm/router_gemv.py diff --git a/python/sglang/kernels/ops/gemm/router_gemv.py b/python/sglang/kernels/ops/gemm/router_gemv.py new file mode 100644 index 000000000..ed4d76217 --- /dev/null +++ b/python/sglang/kernels/ops/gemm/router_gemv.py @@ -0,0 +1,175 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Skinny router-logits GEMV: [M, K] bf16 x [N, K] bf16 -> [M, N] fp32. + +MoE router gates are tiny (N ~ 128 experts, K ~ hidden) but hipblaslt's +solutions for M<=8, N=128 run at ~0.1 TB/s on gfx950 (~12us for 1.6MB). This +split-K kernel reads the gate weight once at near-roofline. The split-K +reduction runs in the same kernel via a last-CTA fixup with a self-cleaning +counter (reset to 0 by the reducing CTA), so no separate zero-init or reduce +launch is needed — every extra launch costs ~2us on gfx950. +""" + +from __future__ import annotations + +from typing import Dict, Tuple + +import torch +import triton +import triton.language as tl + +from sglang.srt.utils import is_gfx95_supported, is_hip + +_is_hip = is_hip() +_is_gfx95_supported = _is_hip and is_gfx95_supported() + +_MAX_M = 64 + +# (BLOCK_M, BLOCK_N, BLOCK_K, SPLIT_K, num_warps) per M bucket. +_CONFIGS = ( + (8, (1, 16, 512, 4, 16)), + (16, (16, 16, 512, 8, 4)), + (32, (16, 32, 512, 8, 4)), + (_MAX_M, (32, 32, 512, 8, 4)), +) + + +def _config(m: int): + for max_m, cfg in _CONFIGS: + if m <= max_m: + return cfg + return _CONFIGS[-1][1] + + +@triton.jit +def _router_gemv_kernel( + x_ptr, + w_ptr, + out_ptr, + partials_ptr, # [SPLIT_K, M, N] fp32 scratch (always fully overwritten) + counter_ptr, # [num_n_blocks * M] int32, all-zero between launches + M, + K, + N, + stride_xm, + stride_wn, + stride_om, + BLOCK_M: tl.constexpr, + BLOCK_K: tl.constexpr, + BLOCK_N: tl.constexpr, + SPLIT_K: tl.constexpr, +): + pid_n = tl.program_id(0) + pid_m = tl.program_id(1) + pid_k = tl.program_id(2) + offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) + if BLOCK_M == 1: + acc = tl.zeros((BLOCK_N,), tl.float32) + for k0 in range(pid_k * BLOCK_K, K, BLOCK_K * SPLIT_K): + offs_k = k0 + tl.arange(0, BLOCK_K) + xv = tl.load(x_ptr + pid_m * stride_xm + offs_k).to(tl.float32) + wv = tl.load(w_ptr + offs_n[:, None] * stride_wn + offs_k[None, :]).to( + tl.float32 + ) + acc += tl.sum(wv * xv[None, :], axis=1) + acc = acc[None, :] + offs_m = pid_m + tl.zeros((1,), tl.int32) + else: + offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) + mask_m = offs_m < M + acc = tl.zeros((BLOCK_M, BLOCK_N), tl.float32) + for k0 in range(pid_k * BLOCK_K, K, BLOCK_K * SPLIT_K): + offs_k = k0 + tl.arange(0, BLOCK_K) + xv = tl.load( + x_ptr + offs_m[:, None] * stride_xm + offs_k[None, :], + mask=mask_m[:, None], + other=0.0, + ) + wv = tl.load(w_ptr + offs_n[:, None] * stride_wn + offs_k[None, :]) + acc += tl.dot(xv, tl.trans(wv), out_dtype=tl.float32) + + part_base = (pid_k * tl.num_programs(1) + pid_m) * (BLOCK_M * N) + tl.store( + partials_ptr + part_base + tl.arange(0, BLOCK_M)[:, None] * N + offs_n, acc + ) + + # Last CTA for this (n-block, row tile) reduces all SPLIT_K partials and + # resets the counter, keeping the buffer reusable with no zeroing launch. + count = tl.atomic_add( + counter_ptr + pid_n * tl.num_programs(1) + pid_m, 1, sem="acq_rel" + ) + if count == SPLIT_K - 1: + total = tl.zeros((BLOCK_M, BLOCK_N), tl.float32) + for s in range(SPLIT_K): + total += tl.load( + partials_ptr + + (s * tl.num_programs(1) + pid_m) * (BLOCK_M * N) + + tl.arange(0, BLOCK_M)[:, None] * N + + offs_n + ) + tl.store( + out_ptr + offs_m[:, None] * stride_om + offs_n, + total, + mask=(offs_m < M)[:, None], + ) + tl.atomic_xchg(counter_ptr + pid_n * tl.num_programs(1) + pid_m, 0) + + +_scratch: Dict[Tuple[torch.device, int], Tuple[torch.Tensor, torch.Tensor]] = {} + + +def _get_scratch(device: torch.device, n: int) -> Tuple[torch.Tensor, torch.Tensor]: + key = (device, n) + if key not in _scratch: + max_split_k = max(cfg[3] for _, cfg in _CONFIGS) + _scratch[key] = ( + torch.empty(max_split_k * _MAX_M * n, dtype=torch.float32, device=device), + torch.zeros(n * _MAX_M // 16, dtype=torch.int32, device=device), + ) + return _scratch[key] + + +def router_gemv_supported(x: torch.Tensor, w: torch.Tensor) -> bool: + m, k = x.shape + n = w.shape[0] + _, block_n, block_k, _, _ = _config(m) + return ( + _is_gfx95_supported + and x.device.type == "cuda" + and w.device == x.device + and x.dtype == torch.bfloat16 + and w.dtype == torch.bfloat16 + and x.stride(1) == 1 + and w.stride(1) == 1 + and n % block_n == 0 + and k % block_k == 0 + and m <= _MAX_M + ) + + +def router_gemv(x: torch.Tensor, w: torch.Tensor) -> torch.Tensor: + """Router logits in fp32; caller guards with router_gemv_supported().""" + m, k = x.shape + n = w.shape[0] + block_m, block_n, block_k, split_k, num_warps = _config(m) + partials, counter = _get_scratch(x.device, n) + out = torch.empty(m, n, device=x.device, dtype=torch.float32) + grid = (n // block_n, triton.cdiv(m, block_m), split_k) + _router_gemv_kernel[grid]( + x, + w, + out, + partials, + counter, + m, + k, + n, + x.stride(0), + w.stride(0), + out.stride(0), + BLOCK_M=block_m, + BLOCK_K=block_k, + BLOCK_N=block_n, + SPLIT_K=split_k, + num_warps=num_warps, + ) + return out diff --git a/python/sglang/srt/models/minimax_m3.py b/python/sglang/srt/models/minimax_m3.py index 786fb11ee..fa2c30723 100644 --- a/python/sglang/srt/models/minimax_m3.py +++ b/python/sglang/srt/models/minimax_m3.py @@ -91,6 +91,7 @@ from sglang.srt.utils import ( add_prefix, get_device_sm, is_cuda, + is_gfx95_supported, is_hip, is_npu, log_info_on_rank0, @@ -101,8 +102,17 @@ from sglang.srt.utils.hf_transformers_utils import get_rope_config _is_cuda = is_cuda() _is_hip = is_hip() _is_npu = is_npu() +_is_gfx95_supported = _is_hip and is_gfx95_supported() _device_sm = get_device_sm() +if _is_gfx95_supported: + from sglang.kernels.ops.gemm.router_gemv import ( + router_gemv, + router_gemv_supported, + ) +else: + router_gemv = router_gemv_supported = None + _FP8_KV_DTYPES = ( torch.float8_e4m3fn, torch.float8_e5m2, @@ -512,6 +522,10 @@ class MiniMaxM3MoE(nn.Module): def _compute_router_logits(self, hidden_states: torch.Tensor) -> torch.Tensor: if self.bf16_router_gemm: + if router_gemv is not None and router_gemv_supported( + hidden_states, self.gate.weight + ): + return router_gemv(hidden_states, self.gate.weight) if _is_npu: # NPU lacks aten::mm.dtype; bf16 mm then cast keeps topk semantics. return torch.mm(hidden_states, self.gate.weight.t()).float()