[refactor] Migrate the first override families: Mistral/Pixtral dtype, MiniMaxM2, MiMoV2 (stack 7/15) (#30069)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-07-04 02:21:32 -07:00
committed by GitHub
co-authored by Claude Fable 5
parent df6491d80c
commit 8d8f17e28c
4 changed files with 189 additions and 26 deletions
+44 -1
View File
@@ -29,13 +29,21 @@ Two declaration forms, keyed on ``hf_config.architectures[0]``:
from __future__ import annotations from __future__ import annotations
import dataclasses import dataclasses
import logging
from typing import Any, Callable, Dict, Iterable, List, Optional, Sequence, Tuple from typing import Any, Callable, Dict, Iterable, List, Optional, Sequence, Tuple
from sglang.srt.arg_groups.arg_utils import model_overridable_fields from sglang.srt.arg_groups.arg_utils import model_overridable_fields
from sglang.srt.runtime_context import resolve_flag_leaf from sglang.srt.runtime_context import resolve_flag_leaf
logger = logging.getLogger(__name__)
# Constant per-architecture overrides (populated by the migration sweeps). # Constant per-architecture overrides (populated by the migration sweeps).
MODEL_OVERRIDES: Dict[str, Dict[str, Any]] = {} MODEL_OVERRIDES: Dict[str, Dict[str, Any]] = {
# These models run in bfloat16 regardless of the requested dtype
# (faithful port of the legacy unconditional arch branch).
"MistralLarge3ForCausalLM": {"dtype": "bfloat16"},
"PixtralForConditionalGeneration": {"dtype": "bfloat16"},
}
# Derived per-architecture override providers, in registration order. # Derived per-architecture override providers, in registration order.
_MODEL_OVERRIDE_FNS: Dict[str, List[Callable[..., dict]]] = {} _MODEL_OVERRIDE_FNS: Dict[str, List[Callable[..., dict]]] = {}
@@ -83,6 +91,41 @@ def collect_model_override_declarations(
return declarations return declarations
# ---------------------------------------------------------------------------
# Derived per-family declarations (faithful ports of legacy arch branches).
# Callables read the PRISTINE server_args, never write; logging is kept
# verbatim from the legacy branch for operator-visible fidelity.
# ---------------------------------------------------------------------------
def _register_for(*architectures: str):
"""Register one provider for several architectures (family lists)."""
def decorator(fn: Callable[..., dict]) -> Callable[..., dict]:
for architecture in architectures:
register_model_override(architecture)(fn)
return fn
return decorator
# 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:
if server_args.speculative_algorithm == "EAGLE":
logger.info("Enable multi-layer EAGLE speculative decoding for MiMoV2 model.")
return {"enable_multi_layer_eagle": True}
return {}
@_register_for("MiniMaxM2ForCausalLM")
def _minimax_m2_overrides(server_args: Any, hf_config: Any) -> dict:
logger.info(
"Enable TF32 matmul for MiniMaxM2ForCausalLM model to improve gate gemm performance."
)
return {"enable_tf32_matmul": True}
@dataclasses.dataclass(frozen=True) @dataclasses.dataclass(frozen=True)
class OverrideRecord: class OverrideRecord:
"""Provenance of one resolved write: ``base`` is the value before this """Provenance of one resolved write: ``base`` is the value before this
+7
View File
@@ -306,6 +306,13 @@ class Flags(_StaticFlags):
moe: MoeFlags = dataclasses.field(default_factory=MoeFlags) moe: MoeFlags = dataclasses.field(default_factory=MoeFlags)
capture: CaptureFlags = dataclasses.field(default_factory=CaptureFlags) capture: CaptureFlags = dataclasses.field(default_factory=CaptureFlags)
# -- resolved config leaves (flat; materialized at publish) --------------
# Pristine user requests stay on the matching server_args fields; these
# leaves carry the model-resolved values.
dtype: str = "auto"
enable_tf32_matmul: bool = False
enable_multi_layer_eagle: bool = False
def freeze(self) -> None: def freeze(self) -> None:
for field in dataclasses.fields(self): for field in dataclasses.fields(self):
value = getattr(self, field.name) value = getattr(self, field.name)
+13 -18
View File
@@ -572,6 +572,7 @@ class ServerArgs:
'* "float32" for FP32 precision.' '* "float32" for FP32 precision.'
), ),
choices=["auto", "half", "float16", "bfloat16", "float", "float32"], choices=["auto", "half", "float16", "bfloat16", "float", "float32"],
model_overridable=True,
), ),
] = "auto" ] = "auto"
quantization: A[ quantization: A[
@@ -653,7 +654,10 @@ class ServerArgs:
] = None # For flash_rl load format ] = None # For flash_rl load format
enable_tf32_matmul: A[ enable_tf32_matmul: A[
bool, bool,
"Enable float32 matmuls to use TensorFloat32 precision for better performance (via torch.set_float32_matmul_precision). CUDA only.", Arg(
help="Enable float32 matmuls to use TensorFloat32 precision for better performance (via torch.set_float32_matmul_precision). CUDA only.",
model_overridable=True,
),
] = False ] = False
# ------------------------------------------------------------------------- # -------------------------------------------------------------------------
@@ -1595,7 +1599,10 @@ class ServerArgs:
] = False ] = False
enable_multi_layer_eagle: A[ enable_multi_layer_eagle: A[
bool, bool,
"Enable multi-layer Eagle speculative decoding.", Arg(
help="Enable multi-layer Eagle speculative decoding.",
model_overridable=True,
),
] = False ] = False
speculative_adaptive: A[ speculative_adaptive: A[
bool, bool,
@@ -3741,12 +3748,6 @@ class ServerArgs:
) )
apply_declarations_to_server_args(self, self._resolved_overrides) apply_declarations_to_server_args(self, self._resolved_overrides)
if model_arch in [
"MistralLarge3ForCausalLM",
"PixtralForConditionalGeneration",
]:
self.dtype = "bfloat16"
if model_arch in [ if model_arch in [
"DeepseekV4ForCausalLM", "DeepseekV4ForCausalLM",
]: ]:
@@ -4205,11 +4206,8 @@ class ServerArgs:
f"attention TP size is {expected_attn_tp_size}." f"attention TP size is {expected_attn_tp_size}."
) )
if self.speculative_algorithm == "EAGLE": # enable_multi_layer_eagle for EAGLE moved to the override registry
self.enable_multi_layer_eagle = True # (arg_groups/overrides.py: _mimo_v2_overrides).
logger.info(
"Enable multi-layer EAGLE speculative decoding for MiMoV2 model."
)
if self.enable_hierarchical_cache: if self.enable_hierarchical_cache:
if not envs.SGLANG_ENABLE_UNIFIED_RADIX_TREE.get(): if not envs.SGLANG_ENABLE_UNIFIED_RADIX_TREE.get():
@@ -4518,11 +4516,8 @@ class ServerArgs:
elif model_arch in ["ZayaForCausalLM"]: elif model_arch in ["ZayaForCausalLM"]:
self._handle_mamba_radix_cache(model_arch=model_arch) self._handle_mamba_radix_cache(model_arch=model_arch)
elif model_arch in ["MiniMaxM2ForCausalLM"]: # MiniMaxM2ForCausalLM (enable_tf32_matmul) moved to the override registry
self.enable_tf32_matmul = True # (arg_groups/overrides.py: _minimax_m2_overrides).
logger.info(
"Enable TF32 matmul for MiniMaxM2ForCausalLM model to improve gate gemm performance."
)
if ( if (
model_arch in ["Qwen3VLForConditionalGeneration"] model_arch in ["Qwen3VLForConditionalGeneration"]
+125 -7
View File
@@ -1,11 +1,15 @@
"""Unit tests for the model-override machinery: whitelist metadata, registry """Unit tests for the model-override machinery: whitelist metadata, registry,
(V3a — declarations only, nothing calls this in production yet).""" gate, publish wiring, and the per-arch golden diffs for migrated families."""
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu") register_cpu_ci(est_time=30, suite="base-a-test-cpu")
import dataclasses import dataclasses
import json
import os
import shutil
import tempfile
import unittest import unittest
from types import SimpleNamespace from types import SimpleNamespace
from typing import Optional from typing import Optional
@@ -48,12 +52,16 @@ class TestModelOverridableWhitelist(CustomTestCase):
frozenset({"resolved_by_model", "also_resolved"}), frozenset({"resolved_by_model", "also_resolved"}),
) )
def test_server_args_whitelist_empty_at_skeleton(self): def test_server_args_whitelist_is_exactly_the_migrated_fields(self):
# No ServerArgs field is tagged yet: the V3 sweeps whitelist fields # Fields are whitelisted one family at a time by the migration
# one family at a time. This pin makes accidental tagging visible. # sweeps. This pin makes accidental tagging visible — extend it in
# the same commit that tags a new field.
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
self.assertEqual(model_overridable_fields(ServerArgs), frozenset()) self.assertEqual(
model_overridable_fields(ServerArgs),
frozenset({"dtype", "enable_tf32_matmul", "enable_multi_layer_eagle"}),
)
def test_non_dataclass_yields_empty_whitelist(self): def test_non_dataclass_yields_empty_whitelist(self):
self.assertEqual(model_overridable_fields(SimpleNamespace), frozenset()) self.assertEqual(model_overridable_fields(SimpleNamespace), frozenset())
@@ -260,6 +268,116 @@ class TestPublishResolvesFlags(_IsolatedPublish):
self.assertIs(get_flags(), flags_before) self.assertIs(get_flags(), flags_before)
class TestGoldenModelOverrides(_IsolatedPublish):
"""Per-arch golden diff for migrated families: the declarative path must
reproduce the legacy imperative writes byte-identically on server_args
(dual-apply) and materialize the same values on the flags tier at
publish."""
_MINI_CONFIG = {
"hidden_size": 64,
"intermediate_size": 128,
"num_attention_heads": 4,
"num_hidden_layers": 2,
"num_key_value_heads": 2,
"vocab_size": 512,
"max_position_embeddings": 128,
"rms_norm_eps": 1e-5,
"torch_dtype": "bfloat16",
# MLA shape fields (required by the MistralLarge3/Pixtral arch
# family; inert extras for non-MLA control arches).
"kv_lora_rank": 32,
"qk_nope_head_dim": 16,
"qk_rope_head_dim": 8,
"v_head_dim": 16,
}
def _construct(self, arch, model_type, **server_kwargs):
from sglang.srt.server_args import ServerArgs
# Golden resolution must be host-independent: accelerator-less CI
# runners resolve only the base platform, where get_device() raises.
server_kwargs.setdefault("device", "cuda")
config = dict(self._MINI_CONFIG, architectures=[arch], model_type=model_type)
config_dir = tempfile.mkdtemp(prefix="golden_override_")
self.addCleanup(shutil.rmtree, config_dir, ignore_errors=True)
with open(os.path.join(config_dir, "config.json"), "w") as f:
json.dump(config, f)
return ServerArgs(model_path=config_dir, **server_kwargs)
def _publish(self, server_args):
from sglang.srt.runtime_context import get_flags
from sglang.srt.server_args import set_global_server_args_for_scheduler
set_global_server_args_for_scheduler(server_args)
return get_flags()
def test_mistral_large3_forces_bfloat16(self):
sa = self._construct("MistralLarge3ForCausalLM", "mistral")
self.assertEqual(sa.dtype, "bfloat16") # dual-apply == legacy write
self.assertEqual(
sa._resolved_overrides,
[("MODEL_OVERRIDES['MistralLarge3ForCausalLM']", {"dtype": "bfloat16"})],
)
self.assertEqual(self._publish(sa).dtype, "bfloat16")
def test_pixtral_forces_bfloat16(self):
sa = self._construct("PixtralForConditionalGeneration", "pixtral")
self.assertEqual(sa.dtype, "bfloat16")
self.assertEqual(self._publish(sa).dtype, "bfloat16")
def test_user_requested_dtype_is_still_overridden(self):
# Legacy fidelity: the arch branch overwrote dtype unconditionally,
# so the declaration must too. The pristine request survives only on
# provenance (and, post-V3, as the un-overridden server_args field).
sa = self._construct("MistralLarge3ForCausalLM", "mistral", dtype="float16")
self.assertEqual(sa.dtype, "bfloat16")
self.assertEqual(self._publish(sa).dtype, "bfloat16")
def test_control_arch_keeps_pristine_dtype(self):
sa = self._construct("LlamaForCausalLM", "llama")
self.assertEqual(sa.dtype, "auto")
self.assertEqual(sa._resolved_overrides, [])
# publish still materializes the whitelisted leaf with the pristine
# value: readers only ever read flags.
self.assertEqual(self._publish(sa).dtype, "auto")
def test_minimax_m2_enables_tf32_matmul(self):
sa = self._construct("MiniMaxM2ForCausalLM", "llama")
self.assertTrue(sa.enable_tf32_matmul) # dual-apply == legacy write
self.assertEqual(
sa._resolved_overrides,
[("_minimax_m2_overrides", {"enable_tf32_matmul": True})],
)
flags = self._publish(sa)
self.assertTrue(flags.enable_tf32_matmul)
self.assertFalse(flags.enable_multi_layer_eagle) # pristine materialize
def test_mimo_v2_declarations(self):
# Callable-level golden: MiMoV2 archs are hybrid (config-shape heavy),
# 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 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})],
)
class TestDualApplyParity(CustomTestCase): class TestDualApplyParity(CustomTestCase):
def test_dual_apply_replays_and_parity_holds(self): def test_dual_apply_replays_and_parity_holds(self):
flags, args = _FakeFlags(), _FakeArgs() flags, args = _FakeFlags(), _FakeArgs()