[refactor] Config resolution pipeline: full-stack review (10-PR series, review only) (#30137)

This commit is contained in:
Cheng Wan
2026-07-05 00:00:07 -07:00
committed by GitHub
parent ce733f106b
commit 8fb99bbaf8
74 changed files with 2035 additions and 628 deletions
@@ -23,12 +23,15 @@ register_amd_ci(est_time=60, suite="stage-b-test-1-gpu-small-amd")
def _mock_global_server_args(backend="pytorch"):
from sglang.srt.layers import sampler as sampler_mod
from sglang.srt.runtime_context import get_flags
from sglang.srt.server_args import ServerArgs
sampler_mod.get_global_server_args = lambda: ServerArgs(
model_path="dummy",
sampling_backend=backend,
)
# The sampler reads the resolved backend from the flags tier.
get_flags().sampling_backend = backend
class _DummyTPGroup:
device_group = None
+2 -1
View File
@@ -7,6 +7,7 @@ import torch.nn as nn
import torch.nn.functional as F
from sglang.srt.layers.logits_processor import LogitsProcessor
from sglang.srt.runtime_context import get_flags
from sglang.srt.server_args import (
ServerArgs,
get_global_server_args,
@@ -44,7 +45,7 @@ class TestLMHeadFP32(unittest.TestCase):
def _make_logprocessor(self, vocab_size, enable_fp32):
set_global_server_args_for_scheduler(ServerArgs(model_path="dummy"))
get_global_server_args().enable_dp_lm_head = False
get_flags().enable_dp_lm_head = False
get_global_server_args().enable_fp32_lm_head = enable_fp32
cfg = SimpleNamespace(vocab_size=vocab_size, final_logit_softcapping=None)
return LogitsProcessor(cfg, skip_all_gather=True, logit_scale=None)
@@ -38,9 +38,11 @@ def _make_target_verify_batch(bs: int) -> ForwardBatch:
def _filter(batch: ForwardBatch, *, lo: int, hi: int) -> ForwardBatch:
fake_args = SimpleNamespace(moe_dense_tp_size=None, attention_backend="fa3")
from sglang.srt.runtime_context import get_flags
with get_parallel().override(attn_tp_size=1), patch.object(
tbo, "get_global_server_args", lambda: fake_args
):
), get_flags().attn.override(backend="fa3"):
return TboForwardBatchPreparer.filter_batch(
batch,
start_token_index=lo,
@@ -1,48 +1,58 @@
import unittest
from types import SimpleNamespace
from unittest.mock import patch
from sglang.srt.models import deepseek_v4 as deepseek_v4_module
from sglang.srt.models.deepseek_v4 import DeepseekV4ForCausalLM
from sglang.srt.runtime_context import get_context, get_flags, reset_context
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=4, suite="base-a-test-cpu")
class TestDeepseekV4SharedExpertFusionPolicy(unittest.TestCase):
"""The disable decision is a load-time resolution: it lands on the flags
tier through declare_load_time_override (dual-applied onto the published
config during the transition)."""
def setUp(self):
self._saved_server_args = get_context()._server_args
def tearDown(self):
if self._saved_server_args is None:
reset_context()
else:
get_context().set_server_args(self._saved_server_args)
def _make_model(self, n_shared_experts=1):
return SimpleNamespace(
config=SimpleNamespace(n_shared_experts=n_shared_experts)
)
def _publish(self, enforce):
server_args = ServerArgs(model_path="dummy")
server_args.enforce_shared_experts_fusion = enforce
get_context().set_server_args(server_args)
return server_args
def test_disables_shared_fusion_without_enforce(self):
server_args = SimpleNamespace(
disable_shared_experts_fusion=False,
enforce_shared_experts_fusion=False,
)
server_args = self._publish(enforce=False)
model = self._make_model()
with patch.object(
deepseek_v4_module, "get_global_server_args", return_value=server_args
):
DeepseekV4ForCausalLM.determine_num_fused_shared_experts(model)
DeepseekV4ForCausalLM.determine_num_fused_shared_experts(model)
self.assertEqual(model.num_fused_shared_experts, 0)
self.assertTrue(get_flags().disable_shared_experts_fusion)
# dual-apply transition: the published config carries the value too
self.assertTrue(server_args.disable_shared_experts_fusion)
def test_enables_shared_fusion_when_enforced(self):
server_args = SimpleNamespace(
disable_shared_experts_fusion=False,
enforce_shared_experts_fusion=True,
)
server_args = self._publish(enforce=True)
model = self._make_model()
with patch.object(
deepseek_v4_module, "get_global_server_args", return_value=server_args
):
DeepseekV4ForCausalLM.determine_num_fused_shared_experts(model)
DeepseekV4ForCausalLM.determine_num_fused_shared_experts(model)
self.assertEqual(model.num_fused_shared_experts, 1)
self.assertFalse(get_flags().disable_shared_experts_fusion)
self.assertFalse(server_args.disable_shared_experts_fusion)
@@ -184,47 +184,75 @@ class TestLoadBalanceMethod(unittest.TestCase):
class TestHiSparseDsaBackendPolicy(unittest.TestCase):
# The backend selection moved to the resolution pipeline; these policy
# tests drive the pass through its read-only view.
@staticmethod
def _resolve(kv_cache_dtype, **kw):
from types import SimpleNamespace
from sglang.srt.arg_groups.overrides import (
ResolvedView,
_dsa_split_backend_resolution,
)
hf = SimpleNamespace(architectures=["DeepseekV32ForCausalLM"])
defaults = dict(
kv_cache_dtype=kv_cache_dtype,
dsa_prefill_backend=None,
dsa_decode_backend=None,
enable_hisparse=True,
)
defaults.update(kw)
view = ResolvedView(
SimpleNamespace(
get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults
)
)
with (
patch("sglang.srt.configs.model_config.is_deepseek_dsa", return_value=True),
patch("sglang.srt.arg_groups.overrides.is_npu", return_value=False),
patch("sglang.srt.arg_groups.overrides.is_xpu", return_value=False),
patch("torch.cuda.get_device_capability", return_value=(9, 0)),
):
declared = _dsa_split_backend_resolution(view)
return {
"dsa_prefill_backend": declared.get(
"dsa_prefill_backend", defaults["dsa_prefill_backend"]
),
"dsa_decode_backend": declared.get(
"dsa_decode_backend", defaults["dsa_decode_backend"]
),
}
@patch("sglang.srt.server_args.is_hip", return_value=False)
def test_hisparse_defaults_to_flashmla_sparse_on_cuda_bfloat16(self, _mock_is_hip):
server_args = ServerArgs(model_path="dummy", enable_hisparse=True)
resolved = self._resolve("bfloat16")
server_args._set_default_dsa_backends(kv_cache_dtype="bfloat16", major=9)
self.assertEqual(server_args.dsa_prefill_backend, "flashmla_sparse")
self.assertEqual(server_args.dsa_decode_backend, "flashmla_sparse")
self.assertEqual(resolved["dsa_prefill_backend"], "flashmla_sparse")
self.assertEqual(resolved["dsa_decode_backend"], "flashmla_sparse")
@patch("sglang.srt.server_args.is_hip", return_value=False)
def test_hisparse_defaults_to_flashmla_kv_on_cuda_fp8(self, _mock_is_hip):
server_args = ServerArgs(model_path="dummy", enable_hisparse=True)
resolved = self._resolve("fp8_e4m3")
server_args._set_default_dsa_backends(kv_cache_dtype="fp8_e4m3", major=9)
self.assertEqual(server_args.dsa_prefill_backend, "flashmla_kv")
self.assertEqual(server_args.dsa_decode_backend, "flashmla_kv")
self.assertEqual(resolved["dsa_prefill_backend"], "flashmla_kv")
self.assertEqual(resolved["dsa_decode_backend"], "flashmla_kv")
@patch("sglang.srt.server_args.is_hip", return_value=True)
def test_hisparse_defaults_to_tilelang_on_rocm(self, _mock_is_hip):
server_args = ServerArgs(model_path="dummy", enable_hisparse=True)
resolved = self._resolve("bfloat16")
server_args._set_default_dsa_backends(kv_cache_dtype="bfloat16", major=9)
self.assertEqual(server_args.dsa_prefill_backend, "tilelang")
self.assertEqual(server_args.dsa_decode_backend, "tilelang")
self.assertEqual(resolved["dsa_prefill_backend"], "tilelang")
self.assertEqual(resolved["dsa_decode_backend"], "tilelang")
@patch("sglang.srt.server_args.is_hip", return_value=True)
def test_hisparse_preserves_rocm_user_backend_and_defaults_missing_side(
self, _mock_is_hip
):
server_args = ServerArgs(
model_path="dummy",
enable_hisparse=True,
dsa_prefill_backend="tilelang",
)
resolved = self._resolve("bfloat16", dsa_prefill_backend="tilelang")
server_args._set_default_dsa_backends(kv_cache_dtype="bfloat16", major=9)
self.assertEqual(server_args.dsa_prefill_backend, "tilelang")
self.assertEqual(server_args.dsa_decode_backend, "tilelang")
self.assertEqual(resolved["dsa_prefill_backend"], "tilelang")
self.assertEqual(resolved["dsa_decode_backend"], "tilelang")
@patch("sglang.srt.server_args.is_hip", return_value=True)
def test_hisparse_accepts_aiter_backend_on_rocm(self, _mock_is_hip):
@@ -25,7 +25,7 @@ _SRT_ROOT = Path(next(iter(sglang.srt.__path__)))
# Baselines counted over python/sglang/srt/**/*.py, including each function's
# own def line. Ratchet: decrease-only.
_RATCHETS = [
("get_global_server_args", r"\bget_global_server_args\s*\(", 346),
("get_global_server_args", r"\bget_global_server_args\s*\(", 278),
(
"set_global_server_args_for_*",
r"\bset_global_server_args_for_(?:scheduler|tokenizer)\s*\(",
@@ -78,6 +78,18 @@ class TestModelOverridableWhitelist(CustomTestCase):
"ep_size",
"moe_dense_tp_size",
"attn_cp_size",
"disable_overlap_schedule",
"uses_mamba_radix_cache",
"mamba_radix_cache_strategy",
"speculative_moe_runner_backend",
"speculative_moe_a2a_backend",
"disable_shared_experts_fusion",
"kv_cache_dtype",
"dsa_prefill_backend",
"dsa_decode_backend",
"prefill_attention_backend",
"decode_attention_backend",
"flashinfer_allreduce_fusion_backend",
}
),
)
@@ -790,6 +802,733 @@ class TestGoldenModelOverrides(_IsolatedPublish):
self.assertEqual(_dllm_page_size(_view(dllm_algorithm=None)), {})
self.assertEqual(_dllm_page_size(_view(disable_radix_cache=True)), {})
def test_overlap_disable_passes(self):
from sglang.srt.arg_groups.overrides import (
ResolvedView,
_dllm_overlap_disable,
_pipeline_parallel_overlap_disable,
_sparse_head_overlap_disable,
)
# pipeline parallelism: declares only when pp_size > 1
self.assertEqual(
_pipeline_parallel_overlap_disable(
ResolvedView(SimpleNamespace(pp_size=1))
),
{},
)
self.assertEqual(
_pipeline_parallel_overlap_disable(
ResolvedView(SimpleNamespace(pp_size=2))
),
{"disable_overlap_schedule": True},
)
# dllm: guarded on the algorithm and the current value
def _view(**kw):
defaults = dict(
dllm_algorithm="LowConfidence", disable_overlap_schedule=False
)
defaults.update(kw)
return ResolvedView(SimpleNamespace(**defaults))
self.assertEqual(_dllm_overlap_disable(_view(dllm_algorithm=None)), {})
self.assertEqual(
_dllm_overlap_disable(_view(disable_overlap_schedule=True)), {}
)
self.assertEqual(
_dllm_overlap_disable(_view()), {"disable_overlap_schedule": True}
)
# embeddings sparse head: keyed on the env var being set
from sglang.srt.environ import envs
view = ResolvedView(SimpleNamespace())
with patch.object(
envs.SGLANG_EMBEDDINGS_SPARSE_HEAD, "is_set", return_value=False
):
self.assertEqual(_sparse_head_overlap_disable(view), {})
with patch.object(
envs.SGLANG_EMBEDDINGS_SPARSE_HEAD, "is_set", return_value=True
):
self.assertEqual(
_sparse_head_overlap_disable(view), {"disable_overlap_schedule": True}
)
def test_deepseek_v4_overrides_at_callable_level(self):
from sglang.srt.arg_groups.overrides import _deepseek_v4_overrides
from sglang.srt.server_args import ServerArgs
hf = SimpleNamespace(architectures=["DeepseekV4ForCausalLM"])
def _args(**kw):
defaults = dict(
device="cuda",
swa_full_tokens_ratio=ServerArgs.swa_full_tokens_ratio,
moe_runner_backend="auto",
get_model_config=lambda: SimpleNamespace(nvfp4_moe_meta=None),
)
defaults.update(kw)
return SimpleNamespace(**defaults)
self.assertEqual(
_deepseek_v4_overrides(_args(), hf),
{
"attention_backend": "dsv4",
"page_size": 256,
"swa_full_tokens_ratio": 0.1,
},
)
# NPU pool geometry
self.assertEqual(
_deepseek_v4_overrides(_args(device="npu"), hf)["page_size"], 128
)
# user-set window ratio survives
self.assertNotIn(
"swa_full_tokens_ratio",
_deepseek_v4_overrides(_args(swa_full_tokens_ratio=0.5), hf),
)
# nvfp4 hybrid checkpoint routes the MoE runner
self.assertEqual(
_deepseek_v4_overrides(
_args(
get_model_config=lambda: SimpleNamespace(nvfp4_moe_meta=object())
),
hf,
)["moe_runner_backend"],
"flashinfer_trtllm_routed",
)
def test_deepseek_v4_sm120_moe_pass(self):
from sglang.srt.arg_groups.overrides import (
ResolvedView,
_deepseek_v4_sm120_moe,
)
def _view(arch="DeepseekV4ForCausalLM", **kw):
hf = SimpleNamespace(architectures=[arch])
defaults = dict(moe_runner_backend="auto")
defaults.update(kw)
return ResolvedView(
SimpleNamespace(
get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults
)
)
with patch.object(overrides_module, "is_sm120_supported", return_value=True):
self.assertEqual(
_deepseek_v4_sm120_moe(_view()), {"moe_runner_backend": "marlin"}
)
self.assertEqual(
_deepseek_v4_sm120_moe(_view(moe_runner_backend="triton")), {}
)
self.assertEqual(_deepseek_v4_sm120_moe(_view(arch="LlamaForCausalLM")), {})
with patch.object(overrides_module, "is_sm120_supported", return_value=False):
self.assertEqual(_deepseek_v4_sm120_moe(_view()), {})
def test_nemotron_h_overrides_at_callable_level(self):
from sglang.srt.arg_groups.overrides import _nemotron_h_overrides
def _hf(quant_algo="NVFP4"):
return SimpleNamespace(
architectures=["NemotronHForCausalLM"],
mlp_hidden_act="relu2",
quantization_config={"quant_algo": quant_algo},
)
def _args(mc_quant, hf, **kw):
mc = SimpleNamespace(quantization=mc_quant, hf_config=hf)
defaults = dict(
quantization=None,
moe_runner_backend="auto",
moe_a2a_backend="none",
attention_backend=None,
get_model_config=lambda: mc,
)
defaults.update(kw)
return SimpleNamespace(**defaults)
hf = _hf()
with patch.object(overrides_module, "is_sm100_supported", return_value=True):
# modelopt checkpoint: quant algo resolution + sm100 defaults
self.assertEqual(
_nemotron_h_overrides(_args("modelopt", hf), hf),
{
"quantization": "modelopt_fp4",
"moe_runner_backend": "flashinfer_trtllm",
"attention_backend": "flashinfer",
},
)
hf_mixed = _hf("MIXED_PRECISION")
self.assertEqual(
_nemotron_h_overrides(_args("modelopt", hf_mixed), hf_mixed)[
"quantization"
],
"modelopt_mixed",
)
with (
patch.object(overrides_module, "is_sm100_supported", return_value=False),
patch.object(overrides_module, "is_cuda", return_value=True),
patch.object(
overrides_module, "get_device_capability", return_value=(9, 0)
),
):
# SM80-SM90 fp4: marlin
self.assertEqual(
_nemotron_h_overrides(_args("modelopt_fp4", hf), hf),
{"quantization": "modelopt_fp4", "moe_runner_backend": "marlin"},
)
# unquantized checkpoint: cutlass fallback, no quant declared
self.assertEqual(
_nemotron_h_overrides(_args(None, hf), hf),
{"moe_runner_backend": "flashinfer_cutlass"},
)
# non-modelopt quantized checkpoint: nothing declared
self.assertEqual(_nemotron_h_overrides(_args("fp8", hf), hf), {})
# user-set moe backend survives
self.assertEqual(
_nemotron_h_overrides(_args(None, hf, moe_runner_backend="triton"), hf),
{},
)
def test_speculative_moe_runner_default_pass(self):
from sglang.srt.arg_groups.overrides import (
ResolvedView,
_speculative_moe_runner_default,
)
self.assertEqual(
_speculative_moe_runner_default(
ResolvedView(
SimpleNamespace(
speculative_moe_runner_backend=None, moe_runner_backend="triton"
)
)
),
{"speculative_moe_runner_backend": "triton"},
)
# user-set draft backend survives
self.assertEqual(
_speculative_moe_runner_default(
ResolvedView(
SimpleNamespace(
speculative_moe_runner_backend="deep_gemm",
moe_runner_backend="auto",
)
)
),
{},
)
def test_dsa_split_backend_resolution_pass(self):
from sglang.srt.arg_groups.overrides import (
ResolvedView,
_dsa_split_backend_resolution,
)
def _view(arch="DeepseekV32ForCausalLM", **kw):
hf = SimpleNamespace(architectures=[arch])
defaults = dict(
kv_cache_dtype="fp8_e4m3",
dsa_prefill_backend=None,
dsa_decode_backend=None,
enable_hisparse=False,
)
defaults.update(kw)
return ResolvedView(
SimpleNamespace(
get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults
)
)
with (
patch("sglang.srt.configs.model_config.is_deepseek_dsa", return_value=True),
patch.object(overrides_module, "is_npu", return_value=False),
patch.object(overrides_module, "is_xpu", return_value=False),
patch.object(overrides_module, "is_hip", return_value=False),
patch("torch.cuda.get_device_capability", return_value=(9, 0)),
):
# Hopper FP8 -> flashmla_kv both
self.assertEqual(
_dsa_split_backend_resolution(_view()),
{
"dsa_prefill_backend": "flashmla_kv",
"dsa_decode_backend": "flashmla_kv",
},
)
# Hopper bf16 -> flashmla_sparse / fa3
self.assertEqual(
_dsa_split_backend_resolution(_view(kv_cache_dtype="bfloat16")),
{
"dsa_prefill_backend": "flashmla_sparse",
"dsa_decode_backend": "fa3",
},
)
# user-set prefill survives; only decode defaulted
self.assertEqual(
_dsa_split_backend_resolution(_view(dsa_prefill_backend="trtllm")),
{"dsa_decode_backend": "flashmla_kv"},
)
# hisparse arm takes precedence (CUDA fp8 -> flashmla_kv)
self.assertEqual(
_dsa_split_backend_resolution(_view(enable_hisparse=True)),
{
"dsa_prefill_backend": "flashmla_kv",
"dsa_decode_backend": "flashmla_kv",
},
)
# non-family arch declares nothing
self.assertEqual(
_dsa_split_backend_resolution(_view(arch="LlamaForCausalLM")), {}
)
with (
patch("sglang.srt.configs.model_config.is_deepseek_dsa", return_value=True),
patch.object(overrides_module, "is_npu", return_value=False),
patch.object(overrides_module, "is_xpu", return_value=False),
patch.object(overrides_module, "is_hip", return_value=True),
patch("torch.cuda.get_device_capability", return_value=(9, 4)),
):
# ROCm with both unset -> tilelang
self.assertEqual(
_dsa_split_backend_resolution(_view(kv_cache_dtype="bfloat16")),
{
"dsa_prefill_backend": "tilelang",
"dsa_decode_backend": "tilelang",
},
)
def test_flashinfer_allreduce_fusion_passes(self):
from sglang.srt.arg_groups.overrides import (
ResolvedView,
_deterministic_allreduce_fusion_disable,
_enforce_disable_allreduce_fusion,
_flashinfer_allreduce_fusion_auto_enable,
)
def _view(arch="Qwen3MoeForCausalLM", **kw):
hf = SimpleNamespace(architectures=[arch])
defaults = dict(
flashinfer_allreduce_fusion_backend=None,
tp_size=2,
enable_dp_attention=False,
nnodes=1,
moe_a2a_backend="none",
enforce_disable_flashinfer_allreduce_fusion=False,
enable_deterministic_inference=False,
)
defaults.update(kw)
return ResolvedView(
SimpleNamespace(
get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults
)
)
with (
patch.object(overrides_module, "is_sm90_supported", return_value=True),
patch.object(overrides_module, "is_sm100_supported", return_value=False),
):
self.assertEqual(
_flashinfer_allreduce_fusion_auto_enable(_view()),
{"flashinfer_allreduce_fusion_backend": "auto"},
)
# guards: unsupported arch / tp==1 / dp attention / a2a backend
self.assertEqual(
_flashinfer_allreduce_fusion_auto_enable(
_view(arch="LlamaForCausalLM")
),
{},
)
self.assertEqual(
_flashinfer_allreduce_fusion_auto_enable(_view(tp_size=1)), {}
)
self.assertEqual(
_flashinfer_allreduce_fusion_auto_enable(
_view(enable_dp_attention=True)
),
{},
)
self.assertEqual(
_flashinfer_allreduce_fusion_auto_enable(
_view(moe_a2a_backend="deepep")
),
{},
)
# SM90 multi-node: blocked (nnodes>1 needs SM100)
self.assertEqual(
_flashinfer_allreduce_fusion_auto_enable(_view(nnodes=2)), {}
)
# user-set backend survives
self.assertEqual(
_flashinfer_allreduce_fusion_auto_enable(
_view(flashinfer_allreduce_fusion_backend="trtllm")
),
{},
)
# enforce-disable wins over everything
self.assertEqual(
_enforce_disable_allreduce_fusion(
_view(
flashinfer_allreduce_fusion_backend="auto",
enforce_disable_flashinfer_allreduce_fusion=True,
)
),
{"flashinfer_allreduce_fusion_backend": None},
)
self.assertEqual(_enforce_disable_allreduce_fusion(_view()), {})
# deterministic inference disables an enabled fusion
self.assertEqual(
_deterministic_allreduce_fusion_disable(
_view(
flashinfer_allreduce_fusion_backend="auto",
enable_deterministic_inference=True,
)
),
{"flashinfer_allreduce_fusion_backend": None},
)
self.assertEqual(
_deterministic_allreduce_fusion_disable(
_view(enable_deterministic_inference=True)
),
{},
)
def test_cutedsl_prefill_backend_fill_pass(self):
from sglang.srt.arg_groups.overrides import (
ResolvedView,
_cutedsl_prefill_backend_fill,
)
def _view(**kw):
defaults = dict(
attention_backend=None,
decode_attention_backend="cutedsl_mla",
prefill_attention_backend=None,
kv_cache_dtype="auto",
)
defaults.update(kw)
return ResolvedView(SimpleNamespace(**defaults))
with patch.object(overrides_module, "is_sm100_supported", return_value=True):
# decode-only cutedsl: prefill defaults to trtllm_mla
self.assertEqual(
_cutedsl_prefill_backend_fill(_view()),
{"prefill_attention_backend": "trtllm_mla"},
)
# user-set prefill survives
self.assertEqual(
_cutedsl_prefill_backend_fill(_view(prefill_attention_backend="fa3")),
{},
)
# cutedsl on the prefill side is rejected
with self.assertRaises(AssertionError):
_cutedsl_prefill_backend_fill(
_view(prefill_attention_backend="cutedsl_mla")
)
# unsupported kv dtype rejected
with self.assertRaises(ValueError):
_cutedsl_prefill_backend_fill(_view(kv_cache_dtype="fp8_e5m2"))
# not a cutedsl config: nothing declared
self.assertEqual(
_cutedsl_prefill_backend_fill(_view(decode_attention_backend=None)),
{},
)
with patch.object(overrides_module, "is_sm100_supported", return_value=False):
with self.assertRaises(ValueError):
_cutedsl_prefill_backend_fill(_view())
def test_moss_vl_overrides_at_callable_level(self):
from sglang.srt.arg_groups.overrides import _moss_vl_overrides
def _args(**kw):
defaults = dict(
attention_backend=None,
prefill_attention_backend=None,
decode_attention_backend=None,
)
defaults.update(kw)
ns = SimpleNamespace(**defaults)
ns.is_attention_backend_not_set = lambda: (
ns.attention_backend is None
and ns.prefill_attention_backend is None
and ns.decode_attention_backend is None
)
ns.get_attention_backends = lambda: (
ns.prefill_attention_backend or ns.attention_backend,
ns.decode_attention_backend or ns.attention_backend,
)
return ns
# nothing set: prefill defaults to flashinfer
self.assertEqual(
_moss_vl_overrides(_args(), None),
{"prefill_attention_backend": "flashinfer"},
)
# compatible user choice passes with no declaration
self.assertEqual(
_moss_vl_overrides(_args(attention_backend="flashinfer"), None), {}
)
# incompatible user choice rejected
with self.assertRaises(AssertionError):
_moss_vl_overrides(_args(attention_backend="fa3"), None)
def test_dsa_kv_cache_dtype_default_pass(self):
from sglang.srt.arg_groups.overrides import (
ResolvedView,
_dsa_kv_cache_dtype_default,
)
def _view(**kw):
hf = SimpleNamespace(architectures=["DeepseekV32ForCausalLM"])
defaults = dict(
kv_cache_dtype="auto",
dsa_prefill_backend=None,
dsa_decode_backend=None,
)
defaults.update(kw)
return ResolvedView(
SimpleNamespace(
get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults
)
)
with (
patch("sglang.srt.configs.model_config.is_deepseek_dsa", return_value=True),
patch.object(overrides_module, "is_npu", return_value=False),
patch.object(overrides_module, "is_xpu", return_value=False),
):
with patch("torch.cuda.get_device_capability", return_value=(9, 0)):
# Hopper: auto -> bfloat16
self.assertEqual(
_dsa_kv_cache_dtype_default(_view()),
{"kv_cache_dtype": "bfloat16"},
)
# alias normalization
self.assertEqual(
_dsa_kv_cache_dtype_default(_view(kv_cache_dtype="bf16")),
{"kv_cache_dtype": "bfloat16"},
)
# explicit value survives (no declaration)
self.assertEqual(
_dsa_kv_cache_dtype_default(_view(kv_cache_dtype="fp8_e4m3")), {}
)
# unsupported dtype rejected
with self.assertRaises(AssertionError):
_dsa_kv_cache_dtype_default(_view(kv_cache_dtype="fp8_e5m2"))
with patch("torch.cuda.get_device_capability", return_value=(10, 0)):
# Blackwell: auto -> fp8
self.assertEqual(
_dsa_kv_cache_dtype_default(_view()),
{"kv_cache_dtype": "fp8_e4m3"},
)
def test_deepseek_v4_kv_cache_dtype_pass(self):
from sglang.srt.arg_groups.overrides import (
ResolvedView,
_deepseek_v4_kv_cache_dtype,
)
def _view(arch="DeepseekV4ForCausalLM", **kw):
hf = SimpleNamespace(architectures=[arch])
defaults = dict(kv_cache_dtype="auto", device="cuda")
defaults.update(kw)
return ResolvedView(
SimpleNamespace(
get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults
)
)
self.assertEqual(
_deepseek_v4_kv_cache_dtype(_view()), {"kv_cache_dtype": "fp8_e4m3"}
)
# NPU pins bfloat16 regardless of the auto default
self.assertEqual(
_deepseek_v4_kv_cache_dtype(_view(device="npu")),
{"kv_cache_dtype": "bfloat16"},
)
# explicit supported value survives
self.assertEqual(
_deepseek_v4_kv_cache_dtype(_view(kv_cache_dtype="bfloat16")), {}
)
with self.assertRaises(AssertionError):
_deepseek_v4_kv_cache_dtype(_view(kv_cache_dtype="fp8_e5m2"))
self.assertEqual(
_deepseek_v4_kv_cache_dtype(_view(arch="LlamaForCausalLM")), {}
)
def test_deepseek_spec_moe_resolution_pass(self):
from sglang.srt.arg_groups.overrides import (
ResolvedView,
_deepseek_spec_moe_resolution,
)
from sglang.srt.environ import envs
def _view(**kw):
hf = SimpleNamespace(architectures=["DeepseekV3ForCausalLM"])
defaults = dict(
quantization="modelopt_fp4",
speculative_algorithm="EAGLE",
speculative_moe_runner_backend=None,
speculative_moe_a2a_backend=None,
ep_size=8,
)
defaults.update(kw)
return ResolvedView(
SimpleNamespace(
get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults
)
)
with patch.object(overrides_module, "is_hip", return_value=True):
with patch.object(
envs.SGLANG_NVFP4_CKPT_FP8_NEXTN_MOE, "get", return_value=False
):
self.assertEqual(
_deepseek_spec_moe_resolution(_view()),
{
"speculative_moe_runner_backend": "triton",
"speculative_moe_a2a_backend": "none",
},
)
# guards: quantization / algorithm / both fields user-set
self.assertEqual(
_deepseek_spec_moe_resolution(_view(quantization="fp8")), {}
)
self.assertEqual(
_deepseek_spec_moe_resolution(_view(speculative_algorithm=None)),
{},
)
self.assertEqual(
_deepseek_spec_moe_resolution(
_view(
speculative_moe_runner_backend="triton",
speculative_moe_a2a_backend="none",
)
),
{},
)
with patch.object(
envs.SGLANG_NVFP4_CKPT_FP8_NEXTN_MOE, "get", return_value=True
):
self.assertEqual(
_deepseek_spec_moe_resolution(_view()),
{
"speculative_moe_runner_backend": "deep_gemm",
"speculative_moe_a2a_backend": "deepep",
},
)
with self.assertRaises(ValueError):
_deepseek_spec_moe_resolution(_view(ep_size=1))
# the arm is HIP-only
with patch.object(overrides_module, "is_hip", return_value=False):
self.assertEqual(_deepseek_spec_moe_resolution(_view()), {})
def test_mamba_radix_cache_resolution_pass(self):
from sglang.srt.arg_groups.overrides import (
ResolvedView,
_mamba_radix_cache_resolution,
supports_mamba_cache_extra_buffer,
)
def _view(arch, layer_types=None, **kw):
hf = SimpleNamespace(architectures=[arch])
if layer_types is not None:
hf.layer_types = layer_types
defaults = dict(
disable_radix_cache=False,
mamba_radix_cache_strategy="auto",
disable_overlap_schedule=False,
page_size=None,
linear_attn_backend="triton",
)
defaults.update(kw)
return ResolvedView(
SimpleNamespace(
get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults
)
)
# arch guard: non-mamba arch declares nothing
self.assertEqual(_mamba_radix_cache_resolution(_view("LlamaForCausalLM")), {})
# radix cache disabled: nothing to resolve
self.assertEqual(
_mamba_radix_cache_resolution(
_view("Qwen3NextForCausalLM", disable_radix_cache=True)
),
{},
)
# auto + overlap wanted + extra-buffer support -> extra_buffer
self.assertEqual(
_mamba_radix_cache_resolution(_view("Qwen3NextForCausalLM")),
{
"uses_mamba_radix_cache": True,
"mamba_radix_cache_strategy": "extra_buffer",
},
)
# auto + no extra-buffer support (Lfm2) -> no_buffer + overlap disable
self.assertEqual(
_mamba_radix_cache_resolution(_view("Lfm2ForCausalLM")),
{
"uses_mamba_radix_cache": True,
"mamba_radix_cache_strategy": "no_buffer",
"disable_overlap_schedule": True,
},
)
# neither overlap nor paging wanted -> no_buffer even when supported
declared = _mamba_radix_cache_resolution(
_view("Qwen3NextForCausalLM", disable_overlap_schedule=True, page_size=1)
)
self.assertEqual(declared["mamba_radix_cache_strategy"], "no_buffer")
self.assertIs(declared["disable_overlap_schedule"], True)
# paging alone wants the extra buffer
self.assertEqual(
_mamba_radix_cache_resolution(
_view(
"Qwen3NextForCausalLM", disable_overlap_schedule=True, page_size=64
)
)["mamba_radix_cache_strategy"],
"extra_buffer",
)
# user-set strategy: only the routing marker is declared
self.assertEqual(
_mamba_radix_cache_resolution(
_view(
"Qwen3NextForCausalLM",
mamba_radix_cache_strategy="extra_buffer_lazy",
)
),
{"uses_mamba_radix_cache": True},
)
# NemotronH routes through the pass (covered by the guard union,
# not the branch chain — its hook invokes the handler)
self.assertEqual(
_mamba_radix_cache_resolution(_view("NemotronHForCausalLM")),
{
"uses_mamba_radix_cache": True,
"mamba_radix_cache_strategy": "extra_buffer",
},
)
# GraniteMoeHybrid is guarded on mamba layer types
self.assertEqual(
_mamba_radix_cache_resolution(
_view("GraniteMoeHybridForCausalLM", layer_types=["attention"])
),
{},
)
self.assertEqual(
_mamba_radix_cache_resolution(
_view("GraniteMoeHybridForCausalLM", layer_types=["mamba", "attention"])
)["mamba_radix_cache_strategy"],
"extra_buffer",
)
# extra-buffer support requires the triton linear-attn backend
self.assertFalse(
supports_mamba_cache_extra_buffer(
SimpleNamespace(linear_attn_backend="fla"), "Qwen3NextForCausalLM"
)
)
def test_page_size_leaf_materializes_end_state(self):
sa = self._construct("LlamaForCausalLM", "llama")
declared = {f for _s, d in sa._resolved_overrides for f in d}
@@ -9,6 +9,7 @@ import unittest
from unittest.mock import patch
import sglang.srt.server_args as server_args_module
from sglang.srt.arg_groups.arg_utils import A, Arg
from sglang.srt.runtime_context import (
Flags,
ParallelContext,
@@ -301,5 +302,146 @@ class TestFlagsTier(_IsolatedServerArgs):
reset_context() # never leave the singleton frozen for other tests
@dataclasses.dataclass
class _FakeResolvedArgs:
"""Publishable fixture with a resolvable whitelist (real flat leaves)."""
page_size: A[int | None, Arg(help="p", resolvable=True)] = None
sampling_backend: A[str | None, Arg(help="s", resolvable=True)] = None
_resolved_overrides: list = dataclasses.field(default_factory=list)
class TestRuntimeResolutionStages(_IsolatedServerArgs):
"""Runtime stages: post-publish declarations re-resolve the flags tier
atomically; freeze_flags() ends the resolution lifecycle."""
def _publish(self, **kw):
args = _FakeResolvedArgs(**kw)
get_context().set_server_args(args)
return args
def test_record_before_publish_raises(self):
reset_context()
with self.assertRaises(ValueError):
get_context().record_runtime_overrides([("stage", {"page_size": 64})])
def test_record_updates_leaves_and_accumulates_stages(self):
args = self._publish(page_size=1, sampling_backend="flashinfer")
self.assertEqual(get_flags().page_size, 1) # publish-time materialize
# dual-apply transition: the call site keeps its imperative write
args.page_size = 64
get_context().record_runtime_overrides([("stage.runner", {"page_size": 64})])
self.assertEqual(get_flags().page_size, 64)
args.sampling_backend = "pytorch"
get_context().record_runtime_overrides(
[("stage.load", {"sampling_backend": "pytorch"})]
)
self.assertEqual(get_flags().sampling_backend, "pytorch")
self.assertEqual(get_flags().page_size, 64) # earlier stage survives
def test_record_parity_failure_rolls_back(self):
self._publish(page_size=1)
flags_before = get_flags()
with self.assertRaises(AssertionError):
# declared value diverges from the live server_args (no dual-apply)
get_context().record_runtime_overrides([("bad", {"page_size": 64})])
self.assertIs(get_flags(), flags_before) # previous flags intact
self.assertEqual(get_context()._runtime_overrides, []) # rolled back
def test_record_whitelist_violation_rolls_back(self):
self._publish()
with self.assertRaises(ValueError):
get_context().record_runtime_overrides([("bad", {"nope": 1})])
self.assertEqual(get_context()._runtime_overrides, [])
def test_freeze_ends_the_resolution_lifecycle(self):
args = self._publish(page_size=1)
try:
get_context().freeze_flags()
self.assertTrue(get_flags().frozen)
with self.assertRaises(RuntimeError):
get_context().record_runtime_overrides([("late", {"page_size": 64})])
with self.assertRaises(RuntimeError):
get_context().set_server_args(args)
finally:
reset_context()
def test_declare_load_time_override_dual_applies_and_records(self):
from sglang.srt.arg_groups.overrides import declare_load_time_override
args = self._publish(page_size=1)
declare_load_time_override("model.load_time", {"page_size": 64})
self.assertEqual(args.page_size, 64) # dual-applied onto server_args
self.assertEqual(get_flags().page_size, 64) # resolved into the leaf
self.assertEqual(
get_context()._runtime_overrides,
[("model.load_time", {"page_size": 64})],
)
def test_failed_republish_keeps_previous_lifecycle(self):
args = self._publish(page_size=1)
args.page_size = 64
get_context().record_runtime_overrides([("stage", {"page_size": 64})])
flags_before = get_flags()
bad = _FakeResolvedArgs(page_size=1)
bad._resolved_overrides = [("bad", {"nope": 1})] # gate rejects
with self.assertRaises(ValueError):
get_context().set_server_args(bad)
# previous publish fully intact: slot, flags, and the recorded stages
self.assertIs(get_context()._server_args, args)
self.assertIs(get_flags(), flags_before)
self.assertEqual(
get_context()._runtime_overrides, [("stage", {"page_size": 64})]
)
def test_capture_tier_seeded_at_publish_and_survives_stages(self):
# seeded from the published config
args = self._publish(page_size=1)
args.enable_torch_compile = True
get_context().set_server_args(args) # re-publish picks up the value
self.assertTrue(get_flags().capture.enable_torch_compile)
# capture-time write (B4) targets the capture leaf
get_flags().capture.enable_torch_compile = False
self.assertFalse(get_flags().capture.enable_torch_compile)
# a runtime-stage re-resolve must not clobber the capture write
args.page_size = 64
get_context().record_runtime_overrides([("stage", {"page_size": 64})])
self.assertFalse(get_flags().capture.enable_torch_compile)
# capture stays writable after freeze
try:
get_context().freeze_flags()
get_flags().capture.enable_torch_compile = True
self.assertTrue(get_flags().capture.enable_torch_compile)
finally:
reset_context()
def test_bare_dataclass_publish_skips_materialization(self):
# object.__new__(ServerArgs) fixtures (no __init__, no field values)
# must publish without touching the flags tier — dataclass defaults
# live on the class, so materializing from them would clobber
# previously resolved flags with defaults.
from sglang.srt.server_args import ServerArgs
self._publish(page_size=64)
self.assertEqual(get_flags().page_size, 64)
bare = object.__new__(ServerArgs)
get_context().set_server_args(bare)
self.assertIs(get_server_args(), bare)
self.assertEqual(get_flags().page_size, 64) # not clobbered
def test_capture_tier_defaults_for_sentinel_publish(self):
get_context().set_server_args(object())
self.assertFalse(get_flags().capture.enable_torch_compile)
def test_republish_clears_runtime_overrides(self):
args = self._publish(page_size=1)
args.page_size = 64
get_context().record_runtime_overrides([("stage", {"page_size": 64})])
self.assertEqual(get_flags().page_size, 64)
self._publish(page_size=1) # fresh lifecycle
self.assertEqual(get_flags().page_size, 1)
self.assertEqual(get_context()._runtime_overrides, [])
if __name__ == "__main__":
unittest.main()