config: resolution reads the declarations, not the fields (#36253)
This commit is contained in:
@@ -1,6 +1,7 @@
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolution_result
|
||||
from sglang.srt.arg_groups.speculative_hook import (
|
||||
_handle_dspark,
|
||||
_target_checkpoint_bundles_dspark_draft,
|
||||
@@ -63,8 +64,13 @@ class TestDsparkDraftPathDefaulting(CustomTestCase):
|
||||
model_path=_BUNDLED_MODEL_PATH, hf_config=_bundled_hf_config()
|
||||
)
|
||||
_handle_dspark(server_args)
|
||||
self.assertEqual(server_args.speculative_draft_model_path, _BUNDLED_MODEL_PATH)
|
||||
self.assertEqual(server_args.speculative_num_draft_tokens, 6)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "speculative_draft_model_path"),
|
||||
_BUNDLED_MODEL_PATH,
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "speculative_num_draft_tokens"), 6
|
||||
)
|
||||
|
||||
def test_plain_target_without_draft_path_raises(self):
|
||||
server_args = _make_dspark_server_args(
|
||||
@@ -80,7 +86,7 @@ class TestDsparkDraftPathDefaulting(CustomTestCase):
|
||||
server_args.speculative_draft_model_path = "deepseek-ai/some-other-dspark-draft"
|
||||
_handle_dspark(server_args)
|
||||
self.assertEqual(
|
||||
server_args.speculative_draft_model_path,
|
||||
resolution_result(server_args, "speculative_draft_model_path"),
|
||||
"deepseek-ai/some-other-dspark-draft",
|
||||
)
|
||||
|
||||
|
||||
@@ -4,6 +4,7 @@ import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolution_result
|
||||
from sglang.srt.configs.embedding_model_spec import resolve_embedding_model_spec
|
||||
from sglang.srt.configs.model_config import (
|
||||
is_multimodal_piecewise_cuda_graph_supported,
|
||||
@@ -90,7 +91,10 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
|
||||
):
|
||||
args._apply_cuda_graph_compatibility()
|
||||
|
||||
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.TC_PIECEWISE)
|
||||
self.assertEqual(
|
||||
resolution_result(args, "cuda_graph_config").prefill.backend,
|
||||
Backend.TC_PIECEWISE,
|
||||
)
|
||||
disable_if_incompatible.assert_called_once()
|
||||
|
||||
def test_trtllm_mla_stays_on_breakable(self):
|
||||
@@ -118,7 +122,10 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
|
||||
):
|
||||
args._apply_cuda_graph_compatibility()
|
||||
|
||||
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
|
||||
self.assertEqual(
|
||||
resolution_result(args, "cuda_graph_config").prefill.backend,
|
||||
Backend.BREAKABLE,
|
||||
)
|
||||
|
||||
def test_explicit_tc_piecewise_overrides_trtllm_mla_default(self):
|
||||
args = ServerArgs(model_path="dummy")
|
||||
@@ -134,7 +141,10 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
|
||||
):
|
||||
args._apply_cuda_graph_compatibility()
|
||||
|
||||
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.TC_PIECEWISE)
|
||||
self.assertEqual(
|
||||
resolution_result(args, "cuda_graph_config").prefill.backend,
|
||||
Backend.TC_PIECEWISE,
|
||||
)
|
||||
|
||||
def test_multimodal_inputs_keep_tc_piecewise_prefill_enabled(self):
|
||||
runner = self._make_prefill_runner(Backend.TC_PIECEWISE)
|
||||
@@ -178,10 +188,16 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
|
||||
):
|
||||
args._handle_model_capability_adjustments()
|
||||
|
||||
self.assertTrue(args.disable_radix_cache)
|
||||
self.assertEqual(args.chunked_prefill_size, -1)
|
||||
self.assertEqual(args.cuda_graph_config.decode.backend, Backend.DISABLED)
|
||||
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
|
||||
self.assertTrue(resolution_result(args, "disable_radix_cache"))
|
||||
self.assertEqual(resolution_result(args, "chunked_prefill_size"), -1)
|
||||
self.assertEqual(
|
||||
resolution_result(args, "cuda_graph_config").decode.backend,
|
||||
Backend.DISABLED,
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(args, "cuda_graph_config").prefill.backend,
|
||||
Backend.BREAKABLE,
|
||||
)
|
||||
|
||||
def test_encoder_embedding_model_enables_embedding_mode_without_flag(self):
|
||||
args = ServerArgs(model_path="dummy")
|
||||
@@ -199,7 +215,7 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
|
||||
with patch.object(args, "get_model_config", return_value=args.model_config):
|
||||
args._handle_model_capability_adjustments()
|
||||
|
||||
self.assertTrue(args.is_embedding)
|
||||
self.assertTrue(resolution_result(args, "is_embedding"))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -15,6 +15,7 @@ import zmq.asyncio
|
||||
from fastapi import HTTPException
|
||||
from PIL import Image
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolution_result
|
||||
from sglang.srt.disaggregation.encoder.preprocessor import (
|
||||
EncoderPreprocessor,
|
||||
EncoderPreprocessResult,
|
||||
@@ -195,7 +196,7 @@ def test_epd_rejection_reads_the_resolved_transfer_backend():
|
||||
finally:
|
||||
shutil.rmtree(config_dir, ignore_errors=True)
|
||||
|
||||
assert resolved.encoder_transfer_backend == "zmq_to_tokenizer"
|
||||
assert resolution_result(resolved, "encoder_transfer_backend") == "zmq_to_tokenizer"
|
||||
# Publish that record: the guard reads the resolved value out of the bags,
|
||||
# so a raw record does not silently disable the rejection.
|
||||
publish(resolved, role="tokenizer")
|
||||
|
||||
@@ -6,6 +6,7 @@ from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolution_result
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.runtime_context import get_context
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
@@ -25,17 +26,20 @@ class TestMmProcessConfigValidation(CustomTestCase):
|
||||
|
||||
def test_valid_config_accepted(self):
|
||||
args = self._validate_config({"image": {"max_pixels": 5000000}})
|
||||
self.assertEqual(args.mm_process_config, {"image": {"max_pixels": 5000000}})
|
||||
self.assertEqual(
|
||||
resolution_result(args, "mm_process_config"),
|
||||
{"image": {"max_pixels": 5000000}},
|
||||
)
|
||||
|
||||
def test_empty_config_accepted(self):
|
||||
args = self._validate_config({})
|
||||
self.assertEqual(args.mm_process_config, {})
|
||||
self.assertEqual(resolution_result(args, "mm_process_config"), {})
|
||||
|
||||
def test_none_config_defaults_to_empty_dict(self):
|
||||
args = self._validate_config(None)
|
||||
# None is kept as-is for dummy models (default happens after early return)
|
||||
# but for real models it would be set to {}
|
||||
self.assertIsNone(args.mm_process_config)
|
||||
self.assertIsNone(resolution_result(args, "mm_process_config"))
|
||||
|
||||
def test_top_level_non_dict_rejected(self):
|
||||
with self.assertRaises(TypeError) as ctx:
|
||||
@@ -64,7 +68,7 @@ class TestMmProcessConfigValidation(CustomTestCase):
|
||||
"audio": {"sample_rate": 16000},
|
||||
}
|
||||
args = self._validate_config(config)
|
||||
self.assertEqual(args.mm_process_config, config)
|
||||
self.assertEqual(resolution_result(args, "mm_process_config"), config)
|
||||
|
||||
|
||||
class TestBaseProcessorConfigExtraction(CustomTestCase):
|
||||
|
||||
@@ -125,6 +125,13 @@ def _registry_collection_is_after_the_build():
|
||||
|
||||
|
||||
def _server_args_names(tree, path):
|
||||
"""Every local that names the record, including the read views over it.
|
||||
|
||||
A resolution-time reader reads through `resolving_view(server_args)` (the
|
||||
declaration stash over the fields): declaration-only resolvers write no
|
||||
field, so a field read there answers with the raw input. `cfg.dtype` after `cfg = resolving_view(sa)` is
|
||||
the same read this scan is looking for, so the local it binds counts.
|
||||
"""
|
||||
names = {"self"} if path.name == "server_args.py" else {"server_args"}
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
|
||||
@@ -142,6 +149,32 @@ def _server_args_names(tree, path):
|
||||
continue
|
||||
if text == "ServerArgs":
|
||||
names.add(arg.arg)
|
||||
# `cfg = resolving_view(server_args)` / `resolved_view(server_args)`
|
||||
for _ in range(2): # a view over a view-holding local is still one
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.Assign):
|
||||
continue
|
||||
value = node.value
|
||||
bare = (
|
||||
isinstance(value, ast.Call)
|
||||
and isinstance(value.func, ast.Name)
|
||||
and value.func.id in ("resolving_view", "resolved_view")
|
||||
and value.args
|
||||
and isinstance(value.args[0], ast.Name)
|
||||
and value.args[0].id in names
|
||||
)
|
||||
# `resolved = self._resolved()` is the same view, spelled as the
|
||||
# record's own member.
|
||||
member = (
|
||||
isinstance(value, ast.Call)
|
||||
and isinstance(value.func, ast.Attribute)
|
||||
and value.func.attr == "_resolved"
|
||||
and isinstance(value.func.value, ast.Name)
|
||||
and value.func.value.id in names
|
||||
)
|
||||
if not (bare or member):
|
||||
continue
|
||||
names |= {t.id for t in node.targets if isinstance(t, ast.Name)}
|
||||
return names
|
||||
|
||||
|
||||
@@ -191,11 +224,12 @@ def _late_resolution_fields():
|
||||
for name in (
|
||||
"server_args.py",
|
||||
"arg_groups/overrides.py",
|
||||
"utils/template_detection.py",
|
||||
"parser/template_detection.py",
|
||||
):
|
||||
path = _SRT / name
|
||||
if not path.exists():
|
||||
continue
|
||||
# A named file that moved away has to be loud; skipping it silently
|
||||
# leaves the scan believing it read a module it never opened.
|
||||
assert path.exists(), f"{name} is not where this scan looks for it"
|
||||
tree = _parsed(path)
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.Call):
|
||||
|
||||
@@ -764,7 +764,7 @@ class TestResolutionDeclarations(CustomTestCase):
|
||||
# Snapshot before publishing: the bag serves the very object the record
|
||||
# holds, so comparing them after the fact compares an object with
|
||||
# itself and passes however the projection behaves.
|
||||
expected = copy.deepcopy(server_args.cuda_graph_config)
|
||||
expected = copy.deepcopy(resolution_result(server_args, "cuda_graph_config"))
|
||||
publish(server_args, role="scheduler")
|
||||
published = get_exec().graph.cuda_graph_config
|
||||
resolved = expected
|
||||
|
||||
@@ -36,6 +36,7 @@ import unittest.mock
|
||||
import torch
|
||||
|
||||
import sglang
|
||||
from sglang.srt.arg_groups.overrides import resolution_result
|
||||
from sglang.srt.environ import EnvField, envs
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils import is_cuda
|
||||
@@ -270,7 +271,10 @@ class TestResolutionIsReproducible(_RestoresProcessState, CustomTestCase):
|
||||
for field in dataclasses.fields(server_args):
|
||||
if field.name in _NOT_COMPARABLE:
|
||||
continue
|
||||
value = getattr(server_args, field.name)
|
||||
# The resolution result, not the field: a declaration-only resolver
|
||||
# never writes the field, so comparing fields would miss exactly
|
||||
# the decisions a leak would shift.
|
||||
value = resolution_result(server_args, field.name)
|
||||
# Nested dataclasses (cuda_graph_config) compare structurally, and
|
||||
# everything else is deep-copied: a snapshot that stored the live
|
||||
# list/dict would follow an in-place mutation, which is exactly the
|
||||
@@ -382,10 +386,13 @@ class TestResolutionIsReproducible(_RestoresProcessState, CustomTestCase):
|
||||
# differs from the cpu that `default_before` resolved to.
|
||||
expected = (
|
||||
"cuda_ipc"
|
||||
if intermediate.mm_feature_transport == "cuda_ipc"
|
||||
if resolution_result(intermediate, "mm_feature_transport")
|
||||
== "cuda_ipc"
|
||||
else "cpu"
|
||||
)
|
||||
self.assertEqual(after.mm_feature_transport, expected)
|
||||
self.assertEqual(
|
||||
resolution_result(after, "mm_feature_transport"), expected
|
||||
)
|
||||
|
||||
def test_resolving_a_sibling_leaves_the_first_alone(self):
|
||||
for label, config, kwargs in _SHAPES:
|
||||
@@ -847,9 +854,9 @@ class TestACopyStaysResolved(_RestoresProcessState, CustomTestCase):
|
||||
self.assertEqual(get_parallel().config.dist_init_addr, "1.2.3.4:5000")
|
||||
self.assertEqual(
|
||||
get_schedule().chunked_prefill_size,
|
||||
parent.chunked_prefill_size,
|
||||
"publishing the copy re-ran resolution; the bag disagrees with the "
|
||||
"record the parent resolved",
|
||||
resolution_result(parent, "chunked_prefill_size"),
|
||||
"publishing the copy re-ran resolution; the bag disagrees with what "
|
||||
"the parent's resolution decided",
|
||||
)
|
||||
|
||||
def test_no_bare_replace_of_a_record_outside_the_helper(self):
|
||||
|
||||
@@ -10,6 +10,7 @@ from unittest.mock import MagicMock, patch
|
||||
|
||||
import sglang.srt.server_args as server_args_module
|
||||
from sglang.srt.arg_groups import pd_disaggregation_hook
|
||||
from sglang.srt.arg_groups.overrides import resolution_result
|
||||
from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding
|
||||
from sglang.srt.entrypoints.sidecar import (
|
||||
SGLANG_GRPC_ENDPOINT_ENV,
|
||||
@@ -61,7 +62,7 @@ class TestPrepareServerArgs(CustomTestCase):
|
||||
|
||||
args.resolve_once()
|
||||
|
||||
self.assertTrue(args.enable_w4a4_mxfp4_megamoe)
|
||||
self.assertTrue(resolution_result(args, "enable_w4a4_mxfp4_megamoe"))
|
||||
self.assertEqual(os.environ["DG_USE_FP4_ACTS"], "1")
|
||||
self.assertEqual(os.environ["DG_USE_MXF4_KIND"], "1")
|
||||
|
||||
@@ -76,14 +77,14 @@ class TestPrepareServerArgs(CustomTestCase):
|
||||
# nothing to be untouched by.
|
||||
args.resolve_once()
|
||||
|
||||
self.assertFalse(args.enable_w4a4_mxfp4_megamoe)
|
||||
self.assertFalse(resolution_result(args, "enable_w4a4_mxfp4_megamoe"))
|
||||
self.assertEqual(os.environ["DG_USE_FP4_ACTS"], "0")
|
||||
self.assertEqual(os.environ["DG_USE_MXF4_KIND"], "0")
|
||||
|
||||
def test_prefill_decode_interval(self):
|
||||
args = ServerArgs(model_path="dummy", prefill_decode_interval=16)
|
||||
args.resolve_once()
|
||||
self.assertEqual(args.prefill_decode_interval, 16)
|
||||
self.assertEqual(resolution_result(args, "prefill_decode_interval"), 16)
|
||||
|
||||
with self.assertRaisesRegex(
|
||||
ValueError, "--prefill-decode-interval must be non-negative"
|
||||
@@ -114,22 +115,24 @@ class TestPrepareServerArgs(CustomTestCase):
|
||||
return server_args
|
||||
|
||||
disabled = _resolved(model_path="dummy")
|
||||
self.assertFalse(disabled.enable_return_hidden_states)
|
||||
self.assertIsNone(disabled.return_hidden_states_mode)
|
||||
self.assertFalse(resolution_result(disabled, "enable_return_hidden_states"))
|
||||
self.assertIsNone(resolution_result(disabled, "return_hidden_states_mode"))
|
||||
|
||||
last = _resolved(
|
||||
model_path="dummy",
|
||||
return_hidden_states_mode="last",
|
||||
)
|
||||
self.assertTrue(last.enable_return_hidden_states)
|
||||
self.assertEqual(last.return_hidden_states_mode, "last")
|
||||
self.assertTrue(resolution_result(last, "enable_return_hidden_states"))
|
||||
self.assertEqual(resolution_result(last, "return_hidden_states_mode"), "last")
|
||||
|
||||
legacy_full = _resolved(
|
||||
model_path="dummy",
|
||||
enable_return_hidden_states=True,
|
||||
)
|
||||
self.assertTrue(legacy_full.enable_return_hidden_states)
|
||||
self.assertEqual(legacy_full.return_hidden_states_mode, "full")
|
||||
self.assertTrue(resolution_result(legacy_full, "enable_return_hidden_states"))
|
||||
self.assertEqual(
|
||||
resolution_result(legacy_full, "return_hidden_states_mode"), "full"
|
||||
)
|
||||
|
||||
parsed_last = prepare_server_args(
|
||||
[
|
||||
@@ -140,8 +143,10 @@ class TestPrepareServerArgs(CustomTestCase):
|
||||
]
|
||||
)
|
||||
parsed_last.resolve_once()
|
||||
self.assertTrue(parsed_last.enable_return_hidden_states)
|
||||
self.assertEqual(parsed_last.return_hidden_states_mode, "last")
|
||||
self.assertTrue(resolution_result(parsed_last, "enable_return_hidden_states"))
|
||||
self.assertEqual(
|
||||
resolution_result(parsed_last, "return_hidden_states_mode"), "last"
|
||||
)
|
||||
|
||||
# The rejection is resolution's, not the constructor's.
|
||||
with self.assertRaisesRegex(
|
||||
@@ -156,13 +161,24 @@ class TestPrepareServerArgs(CustomTestCase):
|
||||
def test_draft_quantization_explicitness_survives_asdict_round_trip(self):
|
||||
inherited = ServerArgs(model_path="dummy", quantization="modelopt_fp4")
|
||||
inherited._handle_missing_default_values()
|
||||
self.assertEqual(inherited.speculative_draft_model_quantization, "modelopt_fp4")
|
||||
self.assertFalse(inherited._speculative_draft_quantization_explicitly_set)
|
||||
self.assertEqual(
|
||||
resolution_result(inherited, "speculative_draft_model_quantization"),
|
||||
"modelopt_fp4",
|
||||
)
|
||||
self.assertFalse(
|
||||
resolution_result(
|
||||
inherited, "_speculative_draft_quantization_explicitly_set"
|
||||
)
|
||||
)
|
||||
|
||||
reconstructed = ServerArgs(**dataclasses.asdict(inherited))
|
||||
reconstructed._handle_missing_default_values()
|
||||
|
||||
self.assertFalse(reconstructed._speculative_draft_quantization_explicitly_set)
|
||||
self.assertFalse(
|
||||
resolution_result(
|
||||
reconstructed, "_speculative_draft_quantization_explicitly_set"
|
||||
)
|
||||
)
|
||||
|
||||
def test_config_nested_dict_args_are_json(self):
|
||||
with tempfile.NamedTemporaryFile(mode="w", suffix=".yaml", delete=False) as f:
|
||||
@@ -219,8 +235,10 @@ class TestImageProcessorBackend(CustomTestCase):
|
||||
|
||||
server_args._handle_deprecated_args()
|
||||
|
||||
self.assertEqual(server_args.image_processor_backend, "pil")
|
||||
self.assertFalse(server_args.disable_fast_image_processor)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "image_processor_backend"), "pil"
|
||||
)
|
||||
self.assertFalse(resolution_result(server_args, "disable_fast_image_processor"))
|
||||
|
||||
def test_legacy_flag_maps_to_pil_with_one_warning(self):
|
||||
server_args = ServerArgs(model_path="dummy", disable_fast_image_processor=True)
|
||||
@@ -228,8 +246,10 @@ class TestImageProcessorBackend(CustomTestCase):
|
||||
with self.assertLogs(server_args_module.logger, level="WARNING") as logs:
|
||||
server_args._handle_deprecated_args()
|
||||
|
||||
self.assertEqual(server_args.image_processor_backend, "pil")
|
||||
self.assertTrue(server_args.disable_fast_image_processor)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "image_processor_backend"), "pil"
|
||||
)
|
||||
self.assertTrue(resolution_result(server_args, "disable_fast_image_processor"))
|
||||
self.assertEqual(
|
||||
sum(
|
||||
"--disable-fast-image-processor is deprecated" in x for x in logs.output
|
||||
@@ -266,7 +286,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
||||
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
|
||||
server_args._handle_multimodal_feature_transport()
|
||||
|
||||
self.assertEqual(server_args.mm_feature_transport, "cuda_ipc")
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "mm_feature_transport"), "cuda_ipc"
|
||||
)
|
||||
self.assertTrue(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
||||
|
||||
output = "\n".join(logs.output)
|
||||
@@ -281,8 +303,12 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
||||
with self.assertLogs(server_args_module.logger, level="WARNING") as logs:
|
||||
server_args._handle_multimodal_feature_transport()
|
||||
|
||||
self.assertEqual(server_args.mm_feature_transport, "cuda_ipc")
|
||||
self.assertFalse(server_args.keep_mm_feature_on_device)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "mm_feature_transport"), "cuda_ipc"
|
||||
)
|
||||
self.assertFalse(
|
||||
resolution_result(server_args, "keep_mm_feature_on_device")
|
||||
)
|
||||
self.assertTrue(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
||||
|
||||
self.assertIn("deprecated", logs.output[0])
|
||||
@@ -305,7 +331,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
||||
with self.assertLogs(server_args_module.logger, level="WARNING") as logs:
|
||||
server_args._handle_multimodal_feature_transport()
|
||||
|
||||
self.assertEqual(server_args.mm_feature_transport, "cpu")
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "mm_feature_transport"), "cpu"
|
||||
)
|
||||
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
||||
|
||||
self.assertIn("overrides", logs.output[0])
|
||||
@@ -316,7 +344,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
||||
with patch.dict(os.environ, {"SGLANG_USE_CUDA_IPC_TRANSPORT": "0"}):
|
||||
server_args._handle_multimodal_feature_transport()
|
||||
|
||||
self.assertEqual(server_args.mm_feature_transport, "cpu")
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "mm_feature_transport"), "cpu"
|
||||
)
|
||||
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
||||
|
||||
@patch("sglang.srt.server_args.is_cuda", return_value=True)
|
||||
@@ -329,7 +359,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
||||
with self.assertNoLogs(server_args_module.logger, level="INFO"):
|
||||
server_args._handle_multimodal_feature_transport()
|
||||
|
||||
self.assertEqual(server_args.mm_feature_transport, "cpu")
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "mm_feature_transport"), "cpu"
|
||||
)
|
||||
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
||||
|
||||
@patch("sglang.srt.server_args.is_cuda", return_value=True)
|
||||
@@ -342,7 +374,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
||||
with self.assertNoLogs(server_args_module.logger, level="INFO"):
|
||||
server_args._handle_multimodal_feature_transport()
|
||||
|
||||
self.assertEqual(server_args.mm_feature_transport, "cpu")
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "mm_feature_transport"), "cpu"
|
||||
)
|
||||
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
||||
|
||||
@patch("sglang.srt.server_args.os.path.exists", return_value=True)
|
||||
@@ -367,7 +401,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
||||
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
|
||||
server_args._handle_multimodal_feature_transport()
|
||||
|
||||
self.assertEqual(server_args.mm_feature_transport, "cuda_vmm")
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "mm_feature_transport"), "cuda_vmm"
|
||||
)
|
||||
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
||||
|
||||
output = "\n".join(logs.output)
|
||||
@@ -394,7 +430,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
||||
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
|
||||
server_args._handle_multimodal_feature_transport()
|
||||
|
||||
self.assertEqual(server_args.mm_feature_transport, "cpu")
|
||||
self.assertEqual(resolution_result(server_args, "mm_feature_transport"), "cpu")
|
||||
self.assertIn("has not opted into CUDA VMM", "\n".join(logs.output))
|
||||
|
||||
@patch("sglang.srt.server_args.os.path.exists", return_value=False)
|
||||
@@ -411,7 +447,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
||||
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
|
||||
server_args._handle_multimodal_feature_transport()
|
||||
|
||||
self.assertEqual(server_args.mm_feature_transport, "cpu")
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "mm_feature_transport"), "cpu"
|
||||
)
|
||||
|
||||
self.assertIn("no IMEX channel", "\n".join(logs.output))
|
||||
|
||||
@@ -427,7 +465,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
||||
envs.SGLANG_USE_CUDA_IPC_TRANSPORT.clear()
|
||||
server_args._handle_multimodal_feature_transport()
|
||||
|
||||
self.assertEqual(server_args.mm_feature_transport, "cpu")
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "mm_feature_transport"), "cpu"
|
||||
)
|
||||
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
||||
|
||||
@patch("sglang.srt.server_args.is_cuda", return_value=True)
|
||||
@@ -439,7 +479,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
||||
envs.SGLANG_USE_CUDA_IPC_TRANSPORT.clear()
|
||||
server_args._handle_multimodal_feature_transport()
|
||||
|
||||
self.assertEqual(server_args.mm_feature_transport, "cpu")
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "mm_feature_transport"), "cpu"
|
||||
)
|
||||
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
||||
|
||||
@patch("sglang.srt.server_args.is_cuda", return_value=False)
|
||||
@@ -474,7 +516,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
||||
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
|
||||
server_args._handle_multimodal_feature_transport()
|
||||
|
||||
self.assertEqual(server_args.mm_feature_transport, "cuda_vmm")
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "mm_feature_transport"), "cuda_vmm"
|
||||
)
|
||||
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
||||
|
||||
output = "\n".join(logs.output)
|
||||
@@ -555,15 +599,22 @@ class TestLoadBalanceMethod(unittest.TestCase):
|
||||
|
||||
def test_non_pd_defaults_to_round_robin(self):
|
||||
server_args = self._load_balance_args(disaggregation_mode="null")
|
||||
self.assertEqual(server_args.load_balance_method, "round_robin")
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "load_balance_method"), "round_robin"
|
||||
)
|
||||
|
||||
def test_pd_prefill_defaults_to_follow_bootstrap_room(self):
|
||||
server_args = self._load_balance_args(disaggregation_mode="prefill")
|
||||
self.assertEqual(server_args.load_balance_method, "follow_bootstrap_room")
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "load_balance_method"),
|
||||
"follow_bootstrap_room",
|
||||
)
|
||||
|
||||
def test_pd_decode_defaults_to_round_robin(self):
|
||||
server_args = self._load_balance_args(disaggregation_mode="decode")
|
||||
self.assertEqual(server_args.load_balance_method, "round_robin")
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "load_balance_method"), "round_robin"
|
||||
)
|
||||
|
||||
def test_pd_prefill_dcp_warns_about_performance(self):
|
||||
server_args = ServerArgs(
|
||||
@@ -581,7 +632,7 @@ class TestLoadBalanceMethod(unittest.TestCase):
|
||||
disaggregation_transfer_backend="mooncake",
|
||||
dcp_size=4,
|
||||
)
|
||||
self.assertTrue(server_args.disable_radix_cache)
|
||||
self.assertTrue(resolution_result(server_args, "disable_radix_cache"))
|
||||
|
||||
def test_pd_decode_dcp_rejects_unsupported_transfer_backend(self):
|
||||
server_args = ServerArgs(
|
||||
@@ -601,7 +652,7 @@ class TestLoadBalanceMethod(unittest.TestCase):
|
||||
disaggregation_transfer_backend="fake",
|
||||
dcp_size=4,
|
||||
)
|
||||
self.assertTrue(server_args.disable_radix_cache)
|
||||
self.assertTrue(resolution_result(server_args, "disable_radix_cache"))
|
||||
|
||||
def test_pd_decode_dcp_rejects_radix_cache(self):
|
||||
server_args = ServerArgs(
|
||||
@@ -665,8 +716,11 @@ class TestLoadBalanceMethod(unittest.TestCase):
|
||||
disaggregation_transfer_backend="mooncake_tcp",
|
||||
)
|
||||
|
||||
self.assertFalse(server_args.disable_radix_cache)
|
||||
self.assertEqual(server_args.disaggregation_transfer_backend, "mooncake")
|
||||
self.assertFalse(resolution_result(server_args, "disable_radix_cache"))
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "disaggregation_transfer_backend"),
|
||||
"mooncake",
|
||||
)
|
||||
|
||||
|
||||
class TestSkipTokenizerInit(unittest.TestCase):
|
||||
@@ -681,8 +735,8 @@ class TestSkipTokenizerInit(unittest.TestCase):
|
||||
server_args._handle_tokenizer_batching()
|
||||
|
||||
# Tokenizer fanout preserved; detokenizer coerced to 1 (no decode work).
|
||||
self.assertEqual(server_args.tokenizer_worker_num, 4)
|
||||
self.assertEqual(server_args.detokenizer_worker_num, 1)
|
||||
self.assertEqual(resolution_result(server_args, "tokenizer_worker_num"), 4)
|
||||
self.assertEqual(resolution_result(server_args, "detokenizer_worker_num"), 1)
|
||||
|
||||
|
||||
class TestHiSparseDsaBackendPolicy(unittest.TestCase):
|
||||
@@ -870,7 +924,7 @@ class TestFa4PageSizeAutoForce(CustomTestCase):
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolved_view
|
||||
|
||||
self.assertEqual(args.page_size, 1) # dual-apply retired: pristine
|
||||
self.assertEqual(args.page_size, 1) # the field stays pristine
|
||||
self.assertEqual(resolved_view(args).page_size, 128)
|
||||
|
||||
@patch("sglang.srt.arg_groups.overrides.is_sm100_supported", return_value=True)
|
||||
@@ -883,7 +937,7 @@ class TestFa4PageSizeAutoForce(CustomTestCase):
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolved_view
|
||||
|
||||
self.assertEqual(args.page_size, 1) # dual-apply retired: pristine
|
||||
self.assertEqual(args.page_size, 1) # the field stays pristine
|
||||
self.assertEqual(resolved_view(args).page_size, 128)
|
||||
|
||||
|
||||
@@ -918,12 +972,12 @@ class TestContextParallelServerArgs(CustomTestCase):
|
||||
def test_canonical_prefill_cp_requires_strategy(self):
|
||||
args = self.parser.parse_args(["--model", "dummy", "--enable-prefill-cp"])
|
||||
|
||||
self.assertTrue(args.enable_prefill_cp)
|
||||
self.assertIsNone(args.cp_strategy)
|
||||
self.assertTrue(resolution_result(args, "enable_prefill_cp"))
|
||||
self.assertIsNone(resolution_result(args, "cp_strategy"))
|
||||
|
||||
server_args = self._new_cp_args(
|
||||
enable_prefill_cp=args.enable_prefill_cp,
|
||||
cp_strategy=args.cp_strategy,
|
||||
enable_prefill_cp=resolution_result(args, "enable_prefill_cp"),
|
||||
cp_strategy=resolution_result(args, "cp_strategy"),
|
||||
)
|
||||
with self.assertRaisesRegex(ValueError, "--cp-strategy"):
|
||||
server_args._handle_context_parallelism()
|
||||
@@ -940,16 +994,18 @@ class TestContextParallelServerArgs(CustomTestCase):
|
||||
)
|
||||
server_args = self._new_cp_args(
|
||||
enable_dsa_prefill_context_parallel=(
|
||||
args.enable_dsa_prefill_context_parallel
|
||||
resolution_result(args, "enable_dsa_prefill_context_parallel")
|
||||
),
|
||||
dsa_prefill_cp_mode=args.dsa_prefill_cp_mode,
|
||||
dsa_prefill_cp_mode=resolution_result(args, "dsa_prefill_cp_mode"),
|
||||
)
|
||||
|
||||
server_args._handle_legacy_cp_arguments()
|
||||
|
||||
self.assertTrue(server_args.enable_prefill_cp)
|
||||
self.assertEqual(server_args.cp_strategy, "interleave")
|
||||
self.assertEqual(server_args.dsa_prefill_cp_mode, "round-robin-split")
|
||||
self.assertTrue(resolution_result(server_args, "enable_prefill_cp"))
|
||||
self.assertEqual(resolution_result(server_args, "cp_strategy"), "interleave")
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "dsa_prefill_cp_mode"), "round-robin-split"
|
||||
)
|
||||
|
||||
def test_canonical_interleave_cp_mirrors_to_dsa_runtime_aliases(self):
|
||||
server_args = self._new_cp_args(
|
||||
@@ -961,10 +1017,18 @@ class TestContextParallelServerArgs(CustomTestCase):
|
||||
server_args._handle_legacy_cp_arguments()
|
||||
server_args._handle_context_parallelism()
|
||||
|
||||
self.assertTrue(server_args.enable_dsa_prefill_context_parallel)
|
||||
self.assertFalse(server_args.enable_prefill_context_parallel)
|
||||
self.assertEqual(server_args.dsa_prefill_cp_mode, "round-robin-split")
|
||||
self.assertEqual(server_args.prefill_cp_mode, "round-robin-split")
|
||||
self.assertTrue(
|
||||
resolution_result(server_args, "enable_dsa_prefill_context_parallel")
|
||||
)
|
||||
self.assertFalse(
|
||||
resolution_result(server_args, "enable_prefill_context_parallel")
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "dsa_prefill_cp_mode"), "round-robin-split"
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "prefill_cp_mode"), "round-robin-split"
|
||||
)
|
||||
|
||||
def test_context_parallel_handler_initializes_cp_strategy(self):
|
||||
server_args = self._new_cp_args(
|
||||
@@ -1049,15 +1113,25 @@ class TestContextParallelServerArgs(CustomTestCase):
|
||||
server_args._handle_legacy_cp_arguments()
|
||||
server_args._handle_context_parallelism()
|
||||
|
||||
self.assertTrue(server_args.enable_prefill_cp)
|
||||
self.assertEqual(server_args.cp_strategy, strategy)
|
||||
self.assertEqual(server_args.dsa_prefill_cp_mode, mode)
|
||||
self.assertEqual(server_args.prefill_cp_mode, mode)
|
||||
self.assertTrue(resolution_result(server_args, "enable_prefill_cp"))
|
||||
self.assertEqual(
|
||||
server_args.enable_dsa_prefill_context_parallel, expect_dsa
|
||||
resolution_result(server_args, "cp_strategy"), strategy
|
||||
)
|
||||
self.assertEqual(
|
||||
server_args.enable_prefill_context_parallel, expect_generic
|
||||
resolution_result(server_args, "dsa_prefill_cp_mode"), mode
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "prefill_cp_mode"), mode
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(
|
||||
server_args, "enable_dsa_prefill_context_parallel"
|
||||
),
|
||||
expect_dsa,
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "enable_prefill_context_parallel"),
|
||||
expect_generic,
|
||||
)
|
||||
|
||||
|
||||
@@ -1308,7 +1382,7 @@ class TestSSLArgs(unittest.TestCase):
|
||||
ssl_certfile="cert.pem",
|
||||
enable_ssl_refresh=True,
|
||||
)
|
||||
self.assertTrue(server_args.enable_ssl_refresh)
|
||||
self.assertTrue(resolution_result(server_args, "enable_ssl_refresh"))
|
||||
|
||||
|
||||
class TestHiCacheArgs(unittest.TestCase):
|
||||
@@ -1328,10 +1402,17 @@ class TestHiCacheArgs(unittest.TestCase):
|
||||
expected_mem_layout: str,
|
||||
expected_decode_backend: str | None = None,
|
||||
):
|
||||
self.assertEqual(args.hicache_io_backend, expected_io_backend)
|
||||
self.assertEqual(args.hicache_mem_layout, expected_mem_layout)
|
||||
self.assertEqual(
|
||||
resolution_result(args, "hicache_io_backend"), expected_io_backend
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(args, "hicache_mem_layout"), expected_mem_layout
|
||||
)
|
||||
if expected_decode_backend is not None:
|
||||
self.assertEqual(args.decode_attention_backend, expected_decode_backend)
|
||||
self.assertEqual(
|
||||
resolution_result(args, "decode_attention_backend"),
|
||||
expected_decode_backend,
|
||||
)
|
||||
|
||||
def test_hicache_io_backend_and_mem_layout_compatibility(self):
|
||||
cases = [
|
||||
@@ -1409,9 +1490,9 @@ class TestHiCacheArgs(unittest.TestCase):
|
||||
)
|
||||
args._handle_hicache()
|
||||
|
||||
self.assertEqual(args.hicache_io_backend, "kernel")
|
||||
self.assertEqual(args.hicache_mem_layout, "page_first")
|
||||
self.assertIsNone(args.decode_attention_backend)
|
||||
self.assertEqual(resolution_result(args, "hicache_io_backend"), "kernel")
|
||||
self.assertEqual(resolution_result(args, "hicache_mem_layout"), "page_first")
|
||||
self.assertIsNone(resolution_result(args, "decode_attention_backend"))
|
||||
|
||||
def test_decode_offload_rejects_host_pool_retraction(self):
|
||||
args = self._make_args(
|
||||
@@ -1494,11 +1575,19 @@ class TestDecoupledSpecArgs(CustomTestCase):
|
||||
"/tmp/tr",
|
||||
]
|
||||
)
|
||||
self.assertEqual(server_args.decoupled_spec_role, "verifier")
|
||||
self.assertEqual(server_args.decoupled_spec_bind_endpoint, "ipc:///tmp/v")
|
||||
self.assertEqual(server_args.decoupled_spec_connect_endpoints, ["ipc:///tmp/d"])
|
||||
self.assertEqual(server_args.decoupled_spec_rank, 0)
|
||||
self.assertEqual(server_args.spec_trace_dir, "/tmp/tr")
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "decoupled_spec_role"), "verifier"
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "decoupled_spec_bind_endpoint"),
|
||||
"ipc:///tmp/v",
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "decoupled_spec_connect_endpoints"),
|
||||
["ipc:///tmp/d"],
|
||||
)
|
||||
self.assertEqual(resolution_result(server_args, "decoupled_spec_rank"), 0)
|
||||
self.assertEqual(resolution_result(server_args, "spec_trace_dir"), "/tmp/tr")
|
||||
|
||||
def test_decoupled_spec_role_rejects_invalid_choice(self):
|
||||
with self.assertRaises(SystemExit):
|
||||
@@ -1533,10 +1622,10 @@ class TestAdaptiveSpecArgs(CustomTestCase):
|
||||
|
||||
handle_speculative_decoding(args)
|
||||
|
||||
self.assertTrue(args.speculative_adaptive)
|
||||
self.assertEqual(args.speculative_eagle_topk, 1)
|
||||
self.assertEqual(args.speculative_num_steps, 3)
|
||||
self.assertEqual(args.speculative_num_draft_tokens, 4)
|
||||
self.assertTrue(resolution_result(args, "speculative_adaptive"))
|
||||
self.assertEqual(resolution_result(args, "speculative_eagle_topk"), 1)
|
||||
self.assertEqual(resolution_result(args, "speculative_num_steps"), 3)
|
||||
self.assertEqual(resolution_result(args, "speculative_num_draft_tokens"), 4)
|
||||
|
||||
|
||||
class TestWaterfillArgs(CustomTestCase):
|
||||
@@ -1552,10 +1641,9 @@ class TestWaterfillArgs(CustomTestCase):
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolved_view
|
||||
|
||||
# dual-apply retired: the fields stay pristine, the declarations win
|
||||
self.assertTrue(server_args.disable_shared_experts_fusion)
|
||||
self.assertFalse(resolved_view(server_args).disable_shared_experts_fusion)
|
||||
self.assertTrue(server_args.enforce_shared_experts_fusion)
|
||||
self.assertTrue(resolution_result(server_args, "enforce_shared_experts_fusion"))
|
||||
|
||||
def test_waterfill_overrides_moe_a2a_backend_to_deepep(self):
|
||||
server_args = ServerArgs(
|
||||
@@ -1570,7 +1658,7 @@ class TestWaterfillArgs(CustomTestCase):
|
||||
|
||||
self.assertEqual(server_args.moe_a2a_backend, "none") # pristine
|
||||
self.assertEqual(resolved_view(server_args).moe_a2a_backend, "deepep")
|
||||
self.assertTrue(server_args.enforce_shared_experts_fusion)
|
||||
self.assertTrue(resolution_result(server_args, "enforce_shared_experts_fusion"))
|
||||
|
||||
def test_waterfill_keeps_megamoe_backend(self):
|
||||
server_args = ServerArgs(
|
||||
@@ -1586,7 +1674,7 @@ class TestWaterfillArgs(CustomTestCase):
|
||||
|
||||
self.assertEqual(resolved_view(server_args).moe_a2a_backend, "megamoe")
|
||||
self.assertFalse(resolved_view(server_args).disable_shared_experts_fusion)
|
||||
self.assertTrue(server_args.enforce_shared_experts_fusion)
|
||||
self.assertTrue(resolution_result(server_args, "enforce_shared_experts_fusion"))
|
||||
|
||||
def test_waterfill_supports_deepep_low_latency_mode(self):
|
||||
server_args = ServerArgs(
|
||||
@@ -1598,9 +1686,9 @@ class TestWaterfillArgs(CustomTestCase):
|
||||
# dummy-model path short-circuits __post_init__; invoke the handler directly.
|
||||
server_args._handle_a2a_moe()
|
||||
|
||||
self.assertEqual(server_args.deepep_mode, "low_latency")
|
||||
self.assertFalse(server_args.disable_cuda_graph)
|
||||
self.assertTrue(server_args.enforce_shared_experts_fusion)
|
||||
self.assertEqual(resolution_result(server_args, "deepep_mode"), "low_latency")
|
||||
self.assertFalse(resolution_result(server_args, "disable_cuda_graph"))
|
||||
self.assertTrue(resolution_result(server_args, "enforce_shared_experts_fusion"))
|
||||
|
||||
|
||||
class TestPrefillOnlyDisableKvCache(unittest.TestCase):
|
||||
@@ -1635,7 +1723,7 @@ class TestPrefillOnlyDisableKvCache(unittest.TestCase):
|
||||
|
||||
def test_valid_minimal_config_constructs(self):
|
||||
sa = self._validate_prefill_only_args()
|
||||
self.assertTrue(sa.prefill_only_disable_kv_cache)
|
||||
self.assertTrue(resolution_result(sa, "prefill_only_disable_kv_cache"))
|
||||
|
||||
def test_rejects_when_not_embedding(self):
|
||||
with self.assertRaisesRegex(ValueError, "requires --is-embedding"):
|
||||
@@ -1725,15 +1813,27 @@ class TestCudaGraphDisaggregationRoles(CustomTestCase):
|
||||
def test_cuda_graph_prefill_role_defaults_disable_decode_graph(self):
|
||||
args = self._handled_args(disaggregation_mode="prefill")
|
||||
|
||||
self.assertFalse(args.disable_cuda_graph)
|
||||
self.assertEqual(args.cuda_graph_config.decode.backend, Backend.DISABLED)
|
||||
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
|
||||
self.assertFalse(resolution_result(args, "disable_cuda_graph"))
|
||||
self.assertEqual(
|
||||
resolution_result(args, "cuda_graph_config").decode.backend,
|
||||
Backend.DISABLED,
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(args, "cuda_graph_config").prefill.backend,
|
||||
Backend.BREAKABLE,
|
||||
)
|
||||
|
||||
def test_cuda_graph_decode_role_defaults_disable_prefill_graph(self):
|
||||
args = self._handled_args(disaggregation_mode="decode")
|
||||
|
||||
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.DISABLED)
|
||||
self.assertNotEqual(args.cuda_graph_config.decode.backend, Backend.DISABLED)
|
||||
self.assertEqual(
|
||||
resolution_result(args, "cuda_graph_config").prefill.backend,
|
||||
Backend.DISABLED,
|
||||
)
|
||||
self.assertNotEqual(
|
||||
resolution_result(args, "cuda_graph_config").decode.backend,
|
||||
Backend.DISABLED,
|
||||
)
|
||||
|
||||
def test_cuda_graph_global_disable_still_disables_both_phases_for_all_roles(self):
|
||||
for disaggregation_mode in ("prefill", "decode", "null"):
|
||||
@@ -1744,10 +1844,12 @@ class TestCudaGraphDisaggregationRoles(CustomTestCase):
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
args.cuda_graph_config.decode.backend, Backend.DISABLED
|
||||
resolution_result(args, "cuda_graph_config").decode.backend,
|
||||
Backend.DISABLED,
|
||||
)
|
||||
self.assertEqual(
|
||||
args.cuda_graph_config.prefill.backend, Backend.DISABLED
|
||||
resolution_result(args, "cuda_graph_config").prefill.backend,
|
||||
Backend.DISABLED,
|
||||
)
|
||||
|
||||
def test_cuda_graph_explicit_decode_backend_survives_prefill_role(self):
|
||||
@@ -1756,7 +1858,9 @@ class TestCudaGraphDisaggregationRoles(CustomTestCase):
|
||||
cuda_graph_backend_decode=Backend.FULL,
|
||||
)
|
||||
|
||||
self.assertEqual(args.cuda_graph_config.decode.backend, Backend.FULL)
|
||||
self.assertEqual(
|
||||
resolution_result(args, "cuda_graph_config").decode.backend, Backend.FULL
|
||||
)
|
||||
self.assertIn((Phase.DECODE, "backend"), args._cuda_graph_config_locked)
|
||||
|
||||
|
||||
@@ -1782,12 +1886,18 @@ class TestPrefillCudaGraphLoRACompatibility(CustomTestCase):
|
||||
def test_enable_lora_keeps_breakable_prefill_graph(self):
|
||||
args = self._handled_args(enable_lora=True)
|
||||
|
||||
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
|
||||
self.assertEqual(
|
||||
resolution_result(args, "cuda_graph_config").prefill.backend,
|
||||
Backend.BREAKABLE,
|
||||
)
|
||||
|
||||
def test_lora_paths_keep_breakable_prefill_graph(self):
|
||||
args = self._handled_args(lora_paths=["dummy/lora-adapter"])
|
||||
|
||||
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
|
||||
self.assertEqual(
|
||||
resolution_result(args, "cuda_graph_config").prefill.backend,
|
||||
Backend.BREAKABLE,
|
||||
)
|
||||
|
||||
def test_lora_still_disables_tc_piecewise_prefill_graph(self):
|
||||
# Pin the tc_piecewise LoRA rule itself, with the hardware rule
|
||||
@@ -1811,7 +1921,10 @@ class TestPrefillCudaGraphLoRACompatibility(CustomTestCase):
|
||||
):
|
||||
args._disable_tc_piecewise_cudagraph_if_incompatible()
|
||||
|
||||
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.DISABLED)
|
||||
self.assertEqual(
|
||||
resolution_result(args, "cuda_graph_config").prefill.backend,
|
||||
Backend.DISABLED,
|
||||
)
|
||||
|
||||
|
||||
class TestBreakableCudaGraphMultimodalAllowlist(CustomTestCase):
|
||||
@@ -1840,7 +1953,10 @@ class TestBreakableCudaGraphMultimodalAllowlist(CustomTestCase):
|
||||
is_multimodal=True,
|
||||
allowlisted=False,
|
||||
)
|
||||
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.DISABLED)
|
||||
self.assertEqual(
|
||||
resolution_result(args, "cuda_graph_config").prefill.backend,
|
||||
Backend.DISABLED,
|
||||
)
|
||||
|
||||
def test_allowlisted_multimodal_arch_keeps_prefill_breakable(self):
|
||||
args = self._handled_args(
|
||||
@@ -1848,7 +1964,10 @@ class TestBreakableCudaGraphMultimodalAllowlist(CustomTestCase):
|
||||
is_multimodal=True,
|
||||
allowlisted=True,
|
||||
)
|
||||
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
|
||||
self.assertEqual(
|
||||
resolution_result(args, "cuda_graph_config").prefill.backend,
|
||||
Backend.BREAKABLE,
|
||||
)
|
||||
|
||||
def test_allowlist_membership(self):
|
||||
from sglang.srt.configs.model_config import (
|
||||
@@ -2046,20 +2165,20 @@ class TestGrpcServerArgs(CustomTestCase):
|
||||
def test_http_only_high_port_does_not_derive_grpc_port(self):
|
||||
sa = self._args(port=56000)
|
||||
sa._handle_deprecated_args()
|
||||
self.assertIsNone(sa.grpc_port)
|
||||
self.assertIsNone(resolution_result(sa, "grpc_port"))
|
||||
|
||||
def test_grpc_port_enables_native_and_env_knobs(self):
|
||||
sa = self._args(grpc_port=50051)
|
||||
with envs.SGLANG_GRPC_WORKER_THREADS.override(8):
|
||||
sa._handle_deprecated_args()
|
||||
self.assertEqual(sa.grpc_port, 50051)
|
||||
self.assertEqual(resolution_result(sa, "grpc_port"), 50051)
|
||||
self.assertEqual(sa.grpc_worker_threads, 8)
|
||||
|
||||
def test_env_grpc_port_enables_native(self):
|
||||
sa = self._args(port=30000)
|
||||
with envs.SGLANG_GRPC_PORT.override(45000):
|
||||
sa._handle_deprecated_args()
|
||||
self.assertEqual(sa.grpc_port, 45000)
|
||||
self.assertEqual(resolution_result(sa, "grpc_port"), 45000)
|
||||
|
||||
@staticmethod
|
||||
def _sidecar_parser():
|
||||
@@ -2196,20 +2315,20 @@ class TestGrpcServerArgs(CustomTestCase):
|
||||
def test_legacy_smg_derives_grpc_port_from_http_port(self):
|
||||
sa = self._args(port=30000, smg_grpc_mode=True)
|
||||
sa._handle_deprecated_args()
|
||||
self.assertEqual(sa.grpc_port, 40000)
|
||||
self.assertEqual(resolution_result(sa, "grpc_port"), 40000)
|
||||
|
||||
def test_grpc_mode_is_deprecated_alias_for_smg_grpc_mode(self):
|
||||
sa = self._args(grpc_mode=True)
|
||||
with self.assertLogs(server_args_module.logger, level="WARNING") as cm:
|
||||
sa._handle_deprecated_args()
|
||||
self.assertTrue(sa.smg_grpc_mode)
|
||||
self.assertTrue(resolution_result(sa, "smg_grpc_mode"))
|
||||
self.assertTrue(any("--grpc-mode is deprecated" in line for line in cm.output))
|
||||
|
||||
def test_legacy_smg_takes_precedence_over_grpc_port(self):
|
||||
sa = self._args(grpc_port=50051, smg_grpc_mode=True)
|
||||
sa._handle_deprecated_args()
|
||||
self.assertTrue(sa.smg_grpc_mode)
|
||||
self.assertEqual(sa.grpc_port, 50051)
|
||||
self.assertTrue(resolution_result(sa, "smg_grpc_mode"))
|
||||
self.assertEqual(resolution_result(sa, "grpc_port"), 50051)
|
||||
|
||||
def test_native_grpc_rejects_multi_tokenizer(self):
|
||||
sa = self._args(grpc_port=40000, tokenizer_worker_num=2)
|
||||
@@ -2254,7 +2373,7 @@ class TestGrpcServerArgs(CustomTestCase):
|
||||
tokenizer_manager=MagicMock(),
|
||||
template_manager=MagicMock(),
|
||||
scheduler_info={},
|
||||
grpc_port=server_args.grpc_port,
|
||||
grpc_port=resolution_result(server_args, "grpc_port"),
|
||||
)
|
||||
|
||||
self.assertEqual(handle, "handle")
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolution_result
|
||||
from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
@@ -33,18 +34,18 @@ def _make_spec_args(device: str, algorithm: str = "EAGLE", **overrides) -> Serve
|
||||
class TestSpecCPUOverlapConstraint(CustomTestCase):
|
||||
def test_cpu_eagle_forces_disable_overlap_schedule(self):
|
||||
args = _make_spec_args(device="cpu")
|
||||
self.assertFalse(args.disable_overlap_schedule)
|
||||
self.assertFalse(resolution_result(args, "disable_overlap_schedule"))
|
||||
|
||||
handle_speculative_decoding(args)
|
||||
|
||||
self.assertTrue(args.disable_overlap_schedule)
|
||||
self.assertTrue(resolution_result(args, "disable_overlap_schedule"))
|
||||
|
||||
def test_cpu_eagle3_forces_disable_overlap_schedule(self):
|
||||
args = _make_spec_args(device="cpu", algorithm="EAGLE3")
|
||||
|
||||
handle_speculative_decoding(args)
|
||||
|
||||
self.assertTrue(args.disable_overlap_schedule)
|
||||
self.assertTrue(resolution_result(args, "disable_overlap_schedule"))
|
||||
|
||||
def test_cpu_explicit_disable_overlap_is_preserved(self):
|
||||
args = _make_spec_args(device="cpu", disable_overlap_schedule=True)
|
||||
@@ -56,7 +57,7 @@ class TestSpecCPUOverlapConstraint(CustomTestCase):
|
||||
) as logs:
|
||||
handle_speculative_decoding(args)
|
||||
|
||||
self.assertTrue(args.disable_overlap_schedule)
|
||||
self.assertTrue(resolution_result(args, "disable_overlap_schedule"))
|
||||
self.assertFalse(
|
||||
any("Overlap schedule" in message for message in logs.output),
|
||||
f"hook warned about overriding an already-disabled overlap: {logs.output}",
|
||||
@@ -68,7 +69,7 @@ class TestSpecCPUOverlapConstraint(CustomTestCase):
|
||||
|
||||
handle_speculative_decoding(args)
|
||||
|
||||
self.assertFalse(args.disable_overlap_schedule)
|
||||
self.assertFalse(resolution_result(args, "disable_overlap_schedule"))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -4,6 +4,7 @@ import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolution_result
|
||||
from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
from sglang.srt.speculative.spec_registry import (
|
||||
@@ -237,7 +238,9 @@ class TestServerArgsHook(_RegistryIsolated):
|
||||
|
||||
handle_speculative_decoding(server_args)
|
||||
|
||||
self.assertEqual(server_args.speculative_algorithm, "MY_HANDLE_ARGS")
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "speculative_algorithm"), "MY_HANDLE_ARGS"
|
||||
)
|
||||
self.assertEqual(server_args.custom_spec_handle_seen, "MY_HANDLE_ARGS")
|
||||
self.assertEqual(server_args.speculative_num_draft_tokens, 7)
|
||||
|
||||
|
||||
@@ -280,10 +280,16 @@ class TestPublishInstallsSlot(_IsolatedPublish):
|
||||
set_global_server_args_for_scheduler(sa)
|
||||
self.assertIs(get_server_args(), sa)
|
||||
# Publishing is what resolved it; the handlers ahead of the dummy
|
||||
# short-circuit still declare.
|
||||
# short-circuit still declare. What they decided is the projection --
|
||||
# the fields keep what the caller passed.
|
||||
from sglang.srt.arg_groups.overrides import resolution_result
|
||||
|
||||
self.assertTrue(sa._resolved_overrides, "publishing declared nothing")
|
||||
for source, declared in sa._resolved_overrides:
|
||||
for field, value in declared.items():
|
||||
self.assertEqual(getattr(sa, field), value, f"{source}: {field}")
|
||||
self.assertEqual(
|
||||
resolution_result(sa, field), value, f"{source}: {field}"
|
||||
)
|
||||
|
||||
|
||||
class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
|
||||
@@ -31,11 +31,15 @@ class TestContextOverride(CustomTestCase):
|
||||
|
||||
def test_override_writes_bag_not_server_args(self):
|
||||
sa = self._publish()
|
||||
before = sa.hicache_ratio
|
||||
# The published leaf, not the field: `hicache_ratio` is resolved by
|
||||
# declaration, so the field still holds what the caller passed.
|
||||
before = rc.get_memory().hicache_ratio
|
||||
pristine = sa.hicache_ratio
|
||||
rc.get_context().override("test", hicache_ratio=before + 1.0)
|
||||
self.assertEqual(rc.get_memory().hicache_ratio, before + 1.0)
|
||||
# server_args stays the pristine startup record.
|
||||
self.assertEqual(sa.hicache_ratio, before)
|
||||
# server_args stays the pristine startup record: the override does not
|
||||
# touch it, and neither did resolution.
|
||||
self.assertEqual(sa.hicache_ratio, pristine)
|
||||
|
||||
def test_override_routes_across_namespaces(self):
|
||||
self._publish()
|
||||
|
||||
@@ -7,6 +7,7 @@ translates field annotations into argparse arguments.
|
||||
import argparse
|
||||
import unittest
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolution_result
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils.common import configure_media_url_security
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
@@ -95,8 +96,10 @@ class TestServerArgsAnnotatedCli(CustomTestCase):
|
||||
"32",
|
||||
]
|
||||
)
|
||||
# The normalization is a declaration.
|
||||
self.assertEqual(
|
||||
sa.allowed_media_domains, ["127.0.0.1", "media.example.com"]
|
||||
resolution_result(sa, "allowed_media_domains"),
|
||||
["127.0.0.1", "media.example.com"],
|
||||
)
|
||||
self.assertEqual(sa.media_url_max_file_size_mb, 32)
|
||||
finally:
|
||||
|
||||
@@ -134,10 +134,10 @@ _ENV_MATRIX = (({}, {"SGLANG_IS_IN_CI": "true"}),)
|
||||
# are step-12 exposure like any other pair.
|
||||
_PASSED = frozenset({"model_path", "device", "random_seed"})
|
||||
|
||||
# The reads that still take a value off the supplied instance. `initialize_moe_config`
|
||||
# is handed the record until the replay goes away; the rest are pre-publish launcher
|
||||
# reads.
|
||||
_EXPOSED = {
|
||||
("dllm/config.py", "max_running_requests"),
|
||||
("dllm/config.py", "model_path"),
|
||||
("speculative/spec_registry.py", "disable_overlap_schedule"),
|
||||
("disaggregation/encoder/server.py", "model_loader_extra_config"),
|
||||
("layers/moe/utils.py", "deepep_mode"),
|
||||
("layers/moe/utils.py", "disable_shared_experts_fusion"),
|
||||
@@ -150,40 +150,16 @@ _EXPOSED = {
|
||||
("configs/embedding_model_spec.py", "disable_radix_cache"),
|
||||
("configs/embedding_model_spec.py", "is_embedding"),
|
||||
("configs/embedding_model_spec.py", "prefill_only_disable_kv_cache"),
|
||||
("configs/model_config.py", "_speculative_draft_quantization_explicitly_set"),
|
||||
("configs/model_config.py", "disable_hybrid_swa_memory"),
|
||||
("configs/model_config.py", "dtype"),
|
||||
("configs/model_config.py", "enable_multi_layer_eagle"),
|
||||
("configs/model_config.py", "is_embedding"),
|
||||
("configs/model_config.py", "model_path"),
|
||||
("configs/model_config.py", "quantization"),
|
||||
("configs/model_config.py", "speculative_algorithm"),
|
||||
("configs/model_config.py", "speculative_draft_model_quantization"),
|
||||
("dllm/config.py", "max_running_requests"),
|
||||
("dllm/config.py", "model_path"),
|
||||
("entrypoints/engine.py", "enable_symm_mem"),
|
||||
("entrypoints/engine.py", "reasoning_parser"),
|
||||
("entrypoints/engine.py", "tool_call_parser"),
|
||||
("layers/cp/base.py", "attn_cp_size"),
|
||||
("layers/cp/base.py", "cp_strategy"),
|
||||
("layers/cp/base.py", "enable_prefill_cp"),
|
||||
("layers/cp/bcg.py", "cp_strategy"),
|
||||
("layers/cp/bcg.py", "enable_prefill_cp"),
|
||||
("layers/flashinfer_comm_fusion.py", "flashinfer_allreduce_fusion_backend"),
|
||||
("layers/moe/utils.py", "deepep_mode"),
|
||||
("layers/moe/utils.py", "moe_a2a_backend"),
|
||||
("layers/moe/utils.py", "moe_runner_backend"),
|
||||
("layers/moe/utils.py", "quantization"),
|
||||
("layers/moe/utils.py", "speculative_moe_runner_backend"),
|
||||
("lora/marlin_lora_temp/policy.py", "lora_paths"),
|
||||
("model_loader/expert_pack_runtime.py", "model_path"),
|
||||
("model_loader/expert_pack_runtime.py", "tokenizer_path"),
|
||||
("parser/template_detection.py", "model_path"),
|
||||
("speculative/adaptive_spec_params.py", "speculative_algorithm"),
|
||||
("speculative/adaptive_spec_params.py", "speculative_eagle_topk"),
|
||||
("speculative/draft_worker_common.py", "speculative_draft_attention_backend"),
|
||||
("speculative/spec_info.py", "enable_multi_layer_eagle"),
|
||||
("speculative/spec_registry.py", "disable_overlap_schedule"),
|
||||
("utils/common.py", "speculative_num_draft_tokens"),
|
||||
("utils/common.py", "speculative_num_steps"),
|
||||
("utils/hf_transformers/processor.py", "image_processor_backend"),
|
||||
@@ -216,26 +192,19 @@ _EXPOSED_CUDA_ONLY: frozenset = frozenset()
|
||||
# some code overrides post-publish. Each needs an ordering judgment, not a blanket
|
||||
# conversion; the list exists so a new one is a decision made when it is written.
|
||||
_OVERRIDDEN_AND_READ = {
|
||||
("configs/model_config.py", "dtype"),
|
||||
("configs/model_config.py", "model_path"),
|
||||
("dllm/config.py", "model_path"),
|
||||
("entrypoints/engine.py", "reasoning_parser"),
|
||||
("entrypoints/engine.py", "tool_call_parser"),
|
||||
("model_loader/expert_pack_runtime.py", "model_path"),
|
||||
("weight_cache/daemon.py", "dp_size"),
|
||||
("weight_cache/daemon.py", "dtype"),
|
||||
("weight_cache/daemon.py", "ep_size"),
|
||||
("weight_cache/daemon.py", "load_format"),
|
||||
("weight_cache/daemon.py", "model_path"),
|
||||
("configs/model_config.py", "dtype"),
|
||||
("configs/model_config.py", "model_path"),
|
||||
("mem_cache/pool_host/common.py", "hicache_storage_backend"),
|
||||
("mem_cache/pool_host/common.py", "hicache_storage_backend_extra_config"),
|
||||
("mem_cache/unified_radix_cache.py", "hicache_storage_backend"),
|
||||
("mem_cache/unified_radix_cache.py", "hicache_storage_backend_extra_config"),
|
||||
("mem_cache/unified_radix_cache.py", "hicache_storage_prefetch_policy"),
|
||||
("mem_cache/unified_radix_cache.py", "hicache_write_policy"),
|
||||
("parser/template_detection.py", "model_path"),
|
||||
("utils/common.py", "speculative_num_draft_tokens"),
|
||||
("utils/common.py", "speculative_num_steps"),
|
||||
("weight_cache/daemon.py", "dp_size"),
|
||||
|
||||
Reference in New Issue
Block a user