[Fix] Route asymmetric-KV models to fa4 on SM100 and pin MiMoV2 FP8 MoE to flashinfer_trtllm (#32818)

This commit is contained in:
Liangsheng Yin
2026-07-29 16:37:43 -07:00
committed by GitHub
parent 3c1717d9b6
commit e5c46ff07d
6 changed files with 88 additions and 21 deletions
@@ -79,7 +79,9 @@ import { MiMoV25Deployment } from '/src/snippets/autoregressive/mimo-v25-deploym
**MiMo-V2.5 (310B):**
- The checkpoint has a TP=4-interleaved fused `qkv_proj`; attention-TP per DP group **must** be 4. Use `--dp = TP / 4`; for TP > 4 this also requires DP-attention. Total GPUs must be a multiple of 4. A bare `--tp 8` without `--dp 2` will fail to load with `MiMoV2 fused qkv_proj checkpoint is TP=4-interleaved; got attention tp_size=8`.
- Single-node deployments: H100/H200 8× GPUs (`--tp 8 --dp 2`), B200 4× GPUs (`--tp 4`, dp=1, no DP-attn flag needed), GB300 4× GPUs (`--tp 4`, single NVL4 node). FP8 quantization.
- On Blackwell, pass `--attention-backend fa4`: MiMoV2's asymmetric KV (`head_dim` 192 / `v_head_dim` 128) fails on the SM100 default `trtllm_mha`, which requires equal K/V widths.
- On Blackwell, pass `--mm-attention-backend fa4` for the V2.5 vision encoder. The checkpoint config requests FlashAttention-3 internally, but SGLang rejects FA3 on Blackwell and expects FA4 for multimodal attention.
- On Blackwell, pass `--moe-runner-backend flashinfer_trtllm`; the default `auto` falls through to the triton fused-MoE runner, ~12% slower at bs=1 decode.
- `--enable-dp-lm-head` and `--mm-enable-dp-encoder` are required whenever `--enable-dp-attention` is on, to keep LM head and encoder sharding consistent.
- EAGLE MTP uses the checkpoint's MTP weights. Enable with `--speculative-algorithm EAGLE` and `--enable-multi-layer-eagle` (both Hopper and Blackwell).
- **Multimodal**: Supports image, video, and audio understanding; see Section 4.3 for invocation examples.
@@ -321,9 +321,9 @@ export const MiMoV25Deployment = () => {
if (multinode) flags.push(...multiNodeFlags(nnodes));
// MoE backend: Blackwell uses flashinfer_trtllm (hardware-driven); Hopper
// MoE backend: Blackwell uses flashinfer_trtllm (both variants); Hopper
// optionally uses DeepEP (toggle).
if (isPro && blackwell) {
if (blackwell) {
flags.push(" --moe-runner-backend flashinfer_trtllm");
} else if (useDeepep) {
flags.push(" --moe-a2a-backend deepep");
@@ -349,7 +349,11 @@ export const MiMoV25Deployment = () => {
flags.push(` --model-loader-extra-config '{"enable_multithread_load": true, "num_threads": 64}'`);
}
} else {
if (blackwell) flags.push(" --mm-attention-backend fa4");
if (blackwell) {
// fa4 is required, not tuning: trtllm_mha rejects MiMoV2's 192/128 KV.
flags.push(" --attention-backend fa4");
flags.push(" --mm-attention-backend fa4");
}
flags.push(" --mem-fraction-static 0.65");
flags.push(" --chunked-prefill-size 16384");
}
+13 -2
View File
@@ -448,10 +448,21 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict:
# Keep in sync with MIMO_V2_MODEL_ARCHS (server_args.py / configs/hf_config.py).
@_register_for("MiMoV2ForCausalLM", "MiMoV2FlashForCausalLM")
def _mimo_v2_overrides(server_args: Any, hf_config: Any) -> dict:
overrides: Dict[str, Any] = {}
if server_args.speculative_algorithm == "EAGLE":
logger.info("Enable multi-layer EAGLE speculative decoding for MiMoV2 model.")
return {"enable_multi_layer_eagle": True}
return {}
overrides["enable_multi_layer_eagle"] = True
# On Blackwell "auto" falls through to the triton fused-MoE runner, ~12%
# slower at bs=1 decode. FP4 checkpoints use flashinfer_mxfp4 instead.
if (
is_sm100_supported()
and server_args.moe_runner_backend == "auto"
and get_quantization_config(hf_config) == "fp8"
):
overrides["moe_runner_backend"] = "flashinfer_trtllm"
logger.info("MiMoV2 FP8 on SM100: moe_runner_backend=flashinfer_trtllm.")
return overrides
@_register_for("MiniMaxM2ForCausalLM")
+11
View File
@@ -713,6 +713,17 @@ class ModelConfig:
def linear_attn_registry_result(self) -> Any:
return get_linear_attn_config(self.hf_config)
@property
def has_asymmetric_kv(self) -> bool:
"""Whether K and V rows differ in width (MiMoV2 is 192 / 128).
Not an ``__init__`` field because the MLA special-casing below still
rewrites ``v_head_dim``.
"""
return (
self.head_dim != self.v_head_dim or self.swa_head_dim != self.swa_v_head_dim
)
def _detect_attention_sinks(self) -> bool:
"""Check whether the model uses learned attention sinks.
+4
View File
@@ -5565,6 +5565,10 @@ class ServerArgs:
or self.speculative_eagle_topk is not None
)
):
# trtllm_mha requires equal K/V row widths; fa4 carries
# v_head_dim through.
if model_config.has_asymmetric_kv:
return "fa4"
return "trtllm_mha"
elif is_hip():
return "aiter"
+51 -16
View File
@@ -348,24 +348,59 @@ class TestGoldenModelOverrides(_IsolatedPublish):
# so the declaration is pinned directly for both provider inputs.
from sglang.srt.arg_groups.overrides import _mimo_v2_overrides
self.assertEqual(
_mimo_v2_overrides(SimpleNamespace(speculative_algorithm="EAGLE"), None),
{"enable_multi_layer_eagle": True},
)
self.assertEqual(
_mimo_v2_overrides(SimpleNamespace(speculative_algorithm=None), None),
{},
)
def _args(**kw):
defaults = dict(speculative_algorithm=None, moe_runner_backend="auto")
defaults.update(kw)
return SimpleNamespace(**defaults)
# Non-SM100: the MoE pin must not fire, so hf_config is never inspected.
with patch.object(overrides_module, "is_sm100_supported", return_value=False):
self.assertEqual(
_mimo_v2_overrides(_args(speculative_algorithm="EAGLE"), None),
{"enable_multi_layer_eagle": True},
)
self.assertEqual(_mimo_v2_overrides(_args(), None), {})
def test_mimo_v2_sm100_fp8_pins_flashinfer_trtllm_moe(self):
"""Blackwell FP8 must not be left on the triton fused-MoE runner."""
from sglang.srt.arg_groups.overrides import _mimo_v2_overrides
def _args(**kw):
defaults = dict(speculative_algorithm=None, moe_runner_backend="auto")
defaults.update(kw)
return SimpleNamespace(**defaults)
with patch.object(overrides_module, "is_sm100_supported", return_value=True):
with patch.object(
overrides_module, "get_quantization_config", return_value="fp8"
):
self.assertEqual(
_mimo_v2_overrides(_args(), None),
{"moe_runner_backend": "flashinfer_trtllm"},
)
# An explicit user choice is never overwritten.
self.assertEqual(
_mimo_v2_overrides(_args(moe_runner_backend="triton"), None), {}
)
# FP4 checkpoints run through flashinfer_mxfp4, so they must not be
# pinned to flashinfer_trtllm.
with patch.object(
overrides_module, "get_quantization_config", return_value="mxfp4"
):
self.assertEqual(_mimo_v2_overrides(_args(), None), {})
def test_mimo_v2_family_is_registered(self):
self.assertEqual(
collect_model_override_declarations(
"MiMoV2FlashForCausalLM",
SimpleNamespace(speculative_algorithm="EAGLE"),
None,
),
[("_mimo_v2_overrides", {"enable_multi_layer_eagle": True})],
)
with patch.object(overrides_module, "is_sm100_supported", return_value=False):
self.assertEqual(
collect_model_override_declarations(
"MiMoV2FlashForCausalLM",
SimpleNamespace(
speculative_algorithm="EAGLE", moe_runner_backend="auto"
),
None,
),
[("_mimo_v2_overrides", {"enable_multi_layer_eagle": True})],
)
def test_step3p_hierarchical_cache_golden(self):
# SWA-hybrid arch: the mini config needs layer_types/sliding_window.