MiniMax-M3: Triton split-K router GEMV with in-kernel fixup (#36557)

Co-authored-by: Kevin Mi <45493463+kevin-mii@users.noreply.github.com>
This commit is contained in:
Chunan Zeng
2026-09-08 16:39:41 -07:00
committed by GitHub
co-authored by Kevin Mi
parent a25bbca8ed
commit 5177a3ec08
2 changed files with 189 additions and 0 deletions
@@ -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
+14
View File
@@ -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()