[refactor] Config resolution pipeline: full-stack review (10-PR series, review only) (#30137)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user