config: the per-instance families read the bags (#35026)

This commit is contained in:
Cheng Wan
2026-08-17 16:17:53 -07:00
committed by GitHub
parent a97bc8db32
commit cba3c5d5ac
45 changed files with 909 additions and 530 deletions
@@ -2,13 +2,13 @@ from __future__ import annotations
import os
import unittest
from types import SimpleNamespace
os.environ["SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE"] = "1"
from sglang.srt.kv_canary.token_oracle.install import install_token_oracle_from_env
from sglang.srt.kv_canary.token_oracle.oracle import HashOracle
from sglang.srt.layers.sampler import _CUSTOM_SAMPLER_FACTORIES
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.test_utils import CustomTestCase
@@ -16,23 +16,26 @@ register_cuda_ci(est_time=60, stage="extra-a", runner_config="1-gpu-small")
register_amd_ci(est_time=60, suite="extra-a-test-1-gpu-small-amd")
def _make_server_args(*, sampling_backend: str) -> SimpleNamespace:
return SimpleNamespace(sampling_backend=sampling_backend)
def _publish(case, *, sampling_backend: str) -> None:
"""The gate reads the published config, so the test publishes one."""
override = get_context().override_server_args(sampling_backend=sampling_backend)
override.install()
case.addCleanup(override.restore)
class TestInstallTokenOracleFromEnv(CustomTestCase):
def test_install_token_oracle_from_env_disabled_returns_none(self) -> None:
"""Verify server-arg-disabled token oracle installation (sampling_backend != 'token_oracle') returns no TokenOracleManager."""
server_args = _make_server_args(sampling_backend="auto")
hook = install_token_oracle_from_env(server_args=server_args, vocab_size=1000)
_publish(self, sampling_backend="auto")
hook = install_token_oracle_from_env(vocab_size=1000)
self.assertIsNone(hook)
def test_install_token_oracle_from_env_enabled_registers_oracle_backend(
self,
) -> None:
"""Verify token oracle installation via sampling_backend='token_oracle' registers the oracle backend."""
server_args = _make_server_args(sampling_backend="token_oracle")
hook = install_token_oracle_from_env(server_args=server_args, vocab_size=512)
_publish(self, sampling_backend="token_oracle")
hook = install_token_oracle_from_env(vocab_size=512)
self.assertIsNotNone(hook)
self.assertIn("token_oracle", _CUSTOM_SAMPLER_FACTORIES)
@@ -40,8 +43,8 @@ class TestInstallTokenOracleFromEnv(CustomTestCase):
self,
) -> None:
"""Verify token oracle installation via sampling_backend='token_oracle' returns a TokenOracleManager wrapping a HashOracle."""
server_args = _make_server_args(sampling_backend="token_oracle")
hook = install_token_oracle_from_env(server_args=server_args, vocab_size=256)
_publish(self, sampling_backend="token_oracle")
hook = install_token_oracle_from_env(vocab_size=256)
self.assertIsNotNone(hook)
self.assertIsInstance(hook.oracle, HashOracle)
self.assertEqual(hook.oracle.vocab_size, 256)
@@ -28,6 +28,7 @@ from sglang.srt.constrained.base_grammar_backend import (
create_grammar_backend,
register_grammar_backend,
)
from sglang.srt.runtime_context import get_context # noqa: E402
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(2.0, "base-a-test-cpu")
@@ -231,18 +232,39 @@ class TestCreateGrammarBackend(unittest.TestCase):
GRAMMAR_BACKEND_REGISTRY.clear()
GRAMMAR_BACKEND_REGISTRY.update(self._saved)
def _publish(self, **fields):
"""Set the config the factory reads.
The factory takes every config value off the published bags, so a test
that sets one on the handed object would be setting something the
factory does not read -- which is how a mismatch between the two used
to stay invisible here.
"""
override = get_context().override_server_args(**fields)
override.install()
self.addCleanup(override.restore)
def _make_server_args(
self, backend="none", reasoning_parser=None, enable_strict_thinking=False
self,
backend="none",
reasoning_parser=None,
enable_strict_thinking=False,
**fields
):
published = {
"grammar_backend": backend,
"reasoning_parser": reasoning_parser,
"enable_strict_thinking": enable_strict_thinking,
"constrained_json_whitespace_pattern": None,
"constrained_json_disable_any_whitespace": False,
}
published.update(fields)
self._publish(**published)
# Handed on to plugin-registered backends; not a config source.
args = MagicMock()
args.override = lambda source, **updates: [
setattr(args, key, value) for key, value in updates.items()
]
args.grammar_backend = backend
args.reasoning_parser = reasoning_parser
args.enable_strict_thinking = enable_strict_thinking
args.constrained_json_whitespace_pattern = None
args.constrained_json_disable_any_whitespace = False
return args
def test_none_backend_returns_none(self):
@@ -293,8 +315,9 @@ class TestCreateGrammarBackend(unittest.TestCase):
def test_outlines_backend(self, mock_outlines_cls):
mock_backend = MagicMock(spec=BaseGrammarBackend)
mock_outlines_cls.return_value = mock_backend
args = self._make_server_args("outlines")
args.constrained_json_whitespace_pattern = r"\s*"
args = self._make_server_args(
"outlines", constrained_json_whitespace_pattern=r"\s*"
)
result = create_grammar_backend(args, "tok", 32000)
mock_outlines_cls.assert_called_once_with("tok", whitespace_pattern=r"\s*")
@@ -304,8 +327,9 @@ class TestCreateGrammarBackend(unittest.TestCase):
def test_xgrammar_backend(self, mock_xgrammar_cls):
mock_backend = MagicMock(spec=BaseGrammarBackend)
mock_xgrammar_cls.return_value = mock_backend
args = self._make_server_args("xgrammar")
args.constrained_json_disable_any_whitespace = True
args = self._make_server_args(
"xgrammar", constrained_json_disable_any_whitespace=True
)
result = create_grammar_backend(args, "tok", 32000, {1, 2})
mock_xgrammar_cls.assert_called_once_with(
@@ -336,9 +360,11 @@ class TestCreateGrammarBackend(unittest.TestCase):
def test_llguidance_backend(self, mock_guidance_cls):
mock_backend = MagicMock(spec=BaseGrammarBackend)
mock_guidance_cls.return_value = mock_backend
args = self._make_server_args("llguidance")
args.constrained_json_disable_any_whitespace = False
args.constrained_json_whitespace_pattern = r"\s+"
args = self._make_server_args(
"llguidance",
constrained_json_disable_any_whitespace=False,
constrained_json_whitespace_pattern=r"\s+",
)
result = create_grammar_backend(args, "tok", 32000, {1, 2})
mock_guidance_cls.assert_called_once_with(
@@ -42,7 +42,7 @@ from sglang.srt.multimodal.kimi_k3_image_processing import (
materialize_kimi_k3_cpu_features,
prepare_kimi_k3_encoder_inputs,
)
from sglang.srt.runtime_context import get_context
from sglang.srt.runtime_context import get_context, publish, reset_context
from sglang.srt.server_args import resolve_encoder_transfer_backend
from sglang.srt.utils import ImageData
from sglang.test.ci.ci_register import register_cpu_ci
@@ -78,29 +78,32 @@ def test_kimi_k3_encoder_transfer_backend_auto_avoids_tp_fanout():
def test_epd_language_only_rejects_missing_dispatched_embedding():
server_args = SimpleNamespace(
override = get_context().override_server_args(
language_only=True,
encoder_transfer_backend="zmq_to_tokenizer",
)
request = SimpleNamespace(need_wait_for_mm_inputs=True)
override.install()
try:
request = SimpleNamespace(need_wait_for_mm_inputs=True)
with pytest.raises(HTTPException) as exc_info:
_reject_missing_dispatched_encoder_embedding(server_args, request, None)
with pytest.raises(HTTPException) as exc_info:
_reject_missing_dispatched_encoder_embedding(request, None)
assert getattr(exc_info.value, "status_code", None) == 503
assert getattr(exc_info.value, "status_code", None) == 503
finally:
override.restore()
def test_epd_rejection_reads_the_resolved_transfer_backend():
"""Tripwire for step 12: this guard fires on the *resolved* backend.
"""This guard fires on the *resolved* backend.
The record is produced by actual resolution -- a language-only Kimi-K3
launch at TP2, whose `encoder_transfer_backend` starts at the argument
default `"auto"` (`ENCODER_TRANSFER_BACKEND_CHOICES[0]`) and is filled in
by `resolve_encoder_transfer_backend` to `"zmq_to_tokenizer"`. Today the
guard therefore rejects. When step 12 makes the instance raw, this same
launch hands the guard a record still at `"auto"`, the rejection silently
stops, and *this test fails* -- which is the signal to give this reader
the resolved value (per-engine overlay or bag) rather than the record.
by `resolve_encoder_transfer_backend` to `"zmq_to_tokenizer"`. The guard
reads that resolved value out of the published bags, so the rejection
survives the record going raw: what a reader must never do is go back to
the record for this field.
Fixed doubles cannot trip on that change, so the record here must come
from resolution, not a SimpleNamespace.
"""
@@ -188,20 +191,30 @@ def test_epd_rejection_reads_the_resolved_transfer_backend():
shutil.rmtree(config_dir, ignore_errors=True)
assert resolved.encoder_transfer_backend == "zmq_to_tokenizer"
request = SimpleNamespace(need_wait_for_mm_inputs=True)
with pytest.raises(HTTPException) as exc_info:
_reject_missing_dispatched_encoder_embedding(resolved, request, None)
assert getattr(exc_info.value, "status_code", None) == 503
# 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")
try:
request = SimpleNamespace(need_wait_for_mm_inputs=True)
with pytest.raises(HTTPException) as exc_info:
_reject_missing_dispatched_encoder_embedding(request, None)
assert getattr(exc_info.value, "status_code", None) == 503
finally:
reset_context()
def test_epd_allows_local_processing_when_request_was_not_dispatched():
server_args = SimpleNamespace(
override = get_context().override_server_args(
language_only=True,
encoder_transfer_backend="zmq_to_tokenizer",
)
request = SimpleNamespace(need_wait_for_mm_inputs=False)
override.install()
try:
request = SimpleNamespace(need_wait_for_mm_inputs=False)
_reject_missing_dispatched_encoder_embedding(server_args, request, None)
_reject_missing_dispatched_encoder_embedding(request, None)
finally:
override.restore()
def _encoder(model_type="kimi_k3"):
@@ -8,13 +8,21 @@ from sglang.srt.managers.scheduler_components.batch_result_processor import (
SchedulerBatchResultProcessor,
)
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
def _make_processor(server_mode: str = "full") -> SchedulerBatchResultProcessor:
def _make_processor(case, server_mode: str = "full") -> SchedulerBatchResultProcessor:
# The server-side hidden-state ceiling is a bag leaf.
override = get_context().override_server_args(
enable_return_hidden_states=True,
return_hidden_states_mode=server_mode,
)
override.install()
case.addCleanup(override.restore)
metrics_reporter = Mock()
metrics_reporter.num_generated_tokens = 0
metrics_reporter.forward_ct_decode = 0
@@ -26,8 +34,6 @@ def _make_processor(server_mode: str = "full") -> SchedulerBatchResultProcessor:
server_args=SimpleNamespace(
enable_metrics=False,
enable_hisparse=False,
enable_return_hidden_states=True,
return_hidden_states_mode=server_mode,
),
model_config=SimpleNamespace(think_end_ids=None),
token_to_kv_pool_allocator=Mock(),
@@ -139,7 +145,7 @@ class TestPrefillHiddenStateOffsets(CustomTestCase):
can_run_cuda_graph=False,
skipped_output_comm=False,
)
processor = _make_processor(server_mode)
processor = _make_processor(self, server_mode)
with (
patch(
@@ -160,7 +166,7 @@ class TestPrefillHiddenStateOffsets(CustomTestCase):
class TestDecodeHiddenStateRetention(CustomTestCase):
def test_last_mode_multi_step_storage_stays_bounded(self):
processor = _make_processor()
processor = _make_processor(self)
req = _DecodeReq()
batch = SimpleNamespace(
reqs=[req],
@@ -4,6 +4,7 @@ from unittest.mock import Mock
from sglang.srt.managers.io_struct import GenerateReqInput
from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -11,19 +12,21 @@ register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestHiddenStateServerMode(CustomTestCase):
@staticmethod
def _make_tokenizer_manager(mode):
def _make_tokenizer_manager(self, mode):
# The server-side hidden-state mode is a bag leaf.
override = get_context().override_server_args(
enable_return_hidden_states=mode is not None,
return_hidden_states_mode=mode,
)
override.install()
self.addCleanup(override.restore)
manager = TokenizerManager.__new__(TokenizerManager)
manager.context_len = 128
manager.num_reserved_tokens = 0
manager.allow_auto_truncate = False
manager.validate_total_tokens = False
manager.is_generation = True
manager.server_args = SimpleNamespace(
enable_return_hidden_states=mode is not None,
return_hidden_states_mode=mode,
enable_custom_logit_processor=False,
)
manager.server_args = SimpleNamespace(enable_custom_logit_processor=False)
manager._validate_token_ids_logprob = Mock()
return manager
@@ -7,6 +7,7 @@ from unittest.mock import MagicMock, patch
import torch
from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_context
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -80,14 +81,20 @@ class TestBaseProcessorConfigExtraction(CustomTestCase):
BaseMultimodalProcessor,
)
# The multimodal config comes from the bags.
override = get_context().override_server_args(
mm_process_config=mm_process_config,
allowed_media_domains=[],
)
override.install()
self.addCleanup(override.restore)
server_args = MagicMock()
server_args.mm_process_config = mm_process_config
server_args.mm_processor_worker_num = mm_processor_worker_num
server_args.mm_io_worker_num = mm_io_worker_num
server_args.mm_preprocess_cache_size_mb = None
server_args.tokenizer_worker_num = 1
server_args.trust_mm_content_hashes = False
server_args.allowed_media_domains = []
server_args.media_url_max_file_size_mb = 64
hf_config = MagicMock()
@@ -170,8 +177,14 @@ class TestBaseProcessorConfigExtraction(CustomTestCase):
class TestMultimodalFeatureTransportRuntime(CustomTestCase):
@staticmethod
def _server_args(mm_feature_transport):
def _server_args(self, mm_feature_transport):
override = get_context().override_server_args(
mm_feature_transport=mm_feature_transport,
mm_process_config={},
allowed_media_domains=[],
)
override.install()
self.addCleanup(override.restore)
return SimpleNamespace(
mm_feature_transport=mm_feature_transport,
image_processor_backend="auto",
@@ -197,8 +210,7 @@ class TestMultimodalFeatureTransportRuntime(CustomTestCase):
return processor
def test_cuda_ipc_pool_uses_resolved_server_arg(self):
# The processor module can be imported before this instance is built;
# transport policy must still resolve from the instance's ServerArgs.
# Transport policy resolves from the mm bag, so the test publishes it.
from sglang.srt.multimodal.processors import base_processor
with (
@@ -792,16 +804,21 @@ class TestDoubleBosGuard(CustomTestCase):
BaseMultimodalProcessor,
)
override = get_context().override_server_args(
mm_process_config={},
mm_feature_transport="cpu",
allowed_media_domains=[],
)
override.install()
self.addCleanup(override.restore)
server_args = MagicMock()
server_args.mm_process_config = {}
server_args.mm_processor_worker_num = 0
server_args.mm_io_worker_num = 0
server_args.mm_feature_transport = "cpu"
server_args.disable_fast_image_processor = True
server_args.mm_preprocess_cache_size_mb = None
server_args.tokenizer_worker_num = 1
server_args.trust_mm_content_hashes = False
server_args.allowed_media_domains = []
server_args.media_url_max_file_size_mb = 64
mock_hf_processor = MagicMock()
@@ -35,6 +35,7 @@ from sglang.srt.managers.tokenizer_manager import ( # noqa: E402
from sglang.srt.observability.req_time_stats import ( # noqa: E402
APIServerReqTimeStats,
)
from sglang.srt.runtime_context import get_context
register_cpu_ci(est_time=15, suite="base-a-test-cpu")
@@ -103,8 +104,15 @@ _PER_REQUEST_OPTIONAL_FIELDS = frozenset(
)
def _make_tokenizer_manager() -> TokenizerManager:
"""Create a TokenizerManager with mocked dependencies, bypassing __init__."""
def _make_tokenizer_manager(case) -> TokenizerManager:
"""Create a TokenizerManager with mocked dependencies, bypassing __init__.
The config it reads comes from the bags, so the stand-in needs a published
config rather than attributes on a mock.
"""
override = get_context().override_server_args(speculative_algorithm=None)
override.install()
case.addCleanup(override.restore)
tm = TokenizerManager.__new__(TokenizerManager)
tm.server_args = MagicMock()
tm._config_updates = []
@@ -212,7 +220,7 @@ class TestRidToStateCleanupOnAbort(CustomTestCase):
def test_abort_removes_rid_from_state(self):
"""After _handle_abort_req, rid should be removed from rid_to_state."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "abort_test_rid"
state = _make_req_state(rid)
tm.rid_to_state[rid] = state
@@ -224,7 +232,7 @@ class TestRidToStateCleanupOnAbort(CustomTestCase):
def test_abort_allows_resubmit_same_rid(self):
"""After abort, _init_req_state should accept the same rid again."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "resubmit_after_abort_rid"
state = _make_req_state(rid)
tm.rid_to_state[rid] = state
@@ -245,7 +253,7 @@ class TestRidToStateCleanupOnAbort(CustomTestCase):
def test_abort_sets_finished_and_notifies(self):
"""_handle_abort_req should mark state as finished and set the event."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "abort_notify_rid"
state = _make_req_state(rid)
tm.rid_to_state[rid] = state
@@ -266,7 +274,7 @@ class TestRidToStateCleanupOnBatchOutput(CustomTestCase):
def test_batch_output_removes_rid_on_finish(self):
"""When a request finishes in _handle_batch_output, rid should be removed."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "batch_finish_rid"
state = _make_req_state(rid)
tm.rid_to_state[rid] = state
@@ -278,7 +286,7 @@ class TestRidToStateCleanupOnBatchOutput(CustomTestCase):
def test_batch_output_allows_resubmit_after_finish(self):
"""After a request finishes, the same rid can be resubmitted."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "batch_resubmit_rid"
state = _make_req_state(rid)
tm.rid_to_state[rid] = state
@@ -299,7 +307,7 @@ class TestRidToStateCleanupOnBatchOutput(CustomTestCase):
def test_batch_output_keeps_rid_when_not_finished(self):
"""When a request is not yet finished, rid should remain in rid_to_state."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "batch_ongoing_rid"
state = _make_req_state(rid)
tm.rid_to_state[rid] = state
@@ -316,7 +324,7 @@ class TestInitReqStateDuplicateDetection(CustomTestCase):
def test_duplicate_rid_raises_error(self):
"""_init_req_state should raise ValueError if rid already exists."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "duplicate_rid"
state = _make_req_state(rid)
tm.rid_to_state[rid] = state
@@ -334,7 +342,7 @@ class TestInitReqStateDuplicateDetection(CustomTestCase):
def test_unique_rid_succeeds(self):
"""_init_req_state should succeed with a unique rid."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "unique_rid"
obj = Mock(spec=GenerateReqInput)
@@ -353,7 +361,7 @@ class TestResubmitAfterCompletion(CustomTestCase):
def test_complete_then_resubmit_same_rid(self):
"""A request that completes normally should allow resubmission with the same rid."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "complete_resubmit_rid"
# Phase 1: simulate a request in rid_to_state, then complete it
@@ -379,7 +387,7 @@ class TestResubmitAfterCompletion(CustomTestCase):
def test_abort_then_resubmit_same_rid(self):
"""An aborted request should allow resubmission with the same rid."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "abort_resubmit_rid"
# Phase 1: simulate a request, then abort it
@@ -413,9 +421,9 @@ class _DummyAsyncCM:
return False
def _make_tm_for_generate() -> TokenizerManager:
def _make_tm_for_generate(case) -> TokenizerManager:
"""Augment the mocked TokenizerManager with what generate_request needs."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(case)
tm.server_args.language_only = False
tm.server_args.tokenizer_worker_num = 1
tm.server_args.enable_strict_thinking = False
@@ -450,7 +458,7 @@ class TestDiscardPendingReqStates(CustomTestCase):
"""Direct tests for _discard_pending_req_states."""
def test_discard_single(self):
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "d_single"
tm.rid_to_state[rid] = _make_req_state(rid)
obj = Mock(spec=GenerateReqInput)
@@ -460,7 +468,7 @@ class TestDiscardPendingReqStates(CustomTestCase):
self.assertNotIn(rid, tm.rid_to_state)
def test_discard_batch_removes_all(self):
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rids = ["d0", "d1", "d2"]
for r in rids:
tm.rid_to_state[r] = _make_req_state(r)
@@ -473,7 +481,7 @@ class TestDiscardPendingReqStates(CustomTestCase):
def test_discard_ignores_already_removed(self):
"""Popping a rid that is no longer present must not raise."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
tm.rid_to_state["p1"] = _make_req_state("p1")
obj = Mock(spec=GenerateReqInput)
obj.is_single = False
@@ -484,7 +492,7 @@ class TestDiscardPendingReqStates(CustomTestCase):
class TestParallelStreamTaskCleanup(CustomTestCase):
def test_failing_choice_cancels_and_closes_sibling_waiters(self):
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
async def drive():
sibling_closed = asyncio.Event()
@@ -512,7 +520,7 @@ class TestParallelStreamTaskCleanup(CustomTestCase):
asyncio.run(drive())
def test_failing_non_stream_choice_cancels_and_closes_sibling_waiters(self):
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
async def drive():
sibling_closed = asyncio.Event()
@@ -546,7 +554,7 @@ class TestGenerateRequestCleanupOnDispatchFailure(CustomTestCase):
"""
def test_single_failure_before_dispatch_cleans_up(self):
tm = _make_tm_for_generate()
tm = _make_tm_for_generate(self)
rid = "single_overlen"
obj = _make_generate_obj(rid, is_single=True)
# Simulate over-length rejection during tokenization/validation.
@@ -566,7 +574,7 @@ class TestGenerateRequestCleanupOnDispatchFailure(CustomTestCase):
self.assertNotIn(rid, tm.rid_to_state)
def test_batch_failure_before_dispatch_cleans_up_all(self):
tm = _make_tm_for_generate()
tm = _make_tm_for_generate(self)
rids = ["b0", "b1", "b2"]
obj = _make_generate_obj(list(rids), is_single=False)
@@ -588,7 +596,7 @@ class TestGenerateRequestCleanupOnDispatchFailure(CustomTestCase):
self.assertNotIn(r, tm.rid_to_state)
def test_thinking_budget_rejects_runtime_without_strict_thinking(self):
tm = _make_tm_for_generate()
tm = _make_tm_for_generate(self)
obj = GenerateReqInput(
text="hello",
rid="thinking-budget",
@@ -15,6 +15,7 @@ from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import (
PrefillCudaGraphRunner,
)
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -23,31 +24,29 @@ register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestHiddenStateGraphRecapture(CustomTestCase):
def test_server_mode_sets_graph_capture_ceiling(self):
disabled = SimpleNamespace(
enable_return_hidden_states=False,
return_hidden_states_mode=None,
)
last = SimpleNamespace(
enable_return_hidden_states=True,
return_hidden_states_mode="last",
)
full = SimpleNamespace(
enable_return_hidden_states=True,
return_hidden_states_mode="full",
)
self.assertEqual(
get_server_return_hidden_states_mode(disabled),
CaptureHiddenMode.NULL,
)
self.assertEqual(
get_server_return_hidden_states_mode(last),
CaptureHiddenMode.LAST,
)
self.assertEqual(
get_server_return_hidden_states_mode(full),
CaptureHiddenMode.FULL,
cases = (
(dict(enable_return_hidden_states=False), CaptureHiddenMode.NULL),
(
dict(
enable_return_hidden_states=True, return_hidden_states_mode="last"
),
CaptureHiddenMode.LAST,
),
(
dict(
enable_return_hidden_states=True, return_hidden_states_mode="full"
),
CaptureHiddenMode.FULL,
),
)
for fields, expected in cases:
with self.subTest(**fields):
override = get_context().override_server_args(**fields)
override.install()
try:
self.assertEqual(get_server_return_hidden_states_mode(), expected)
finally:
override.restore()
@staticmethod
def _make_runner(runner_cls, capture_hidden_mode):
@@ -17,6 +17,7 @@ from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import (
PrefillCudaGraphRunner,
)
from sglang.srt.model_executor.runner.shape_key import ShapeKey
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -111,13 +112,17 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
def test_eagle_target_tc_piecewise_skips_last_mode_capture(self):
eager_runner = object()
# The server-side hidden-state ceiling is a bag leaf.
override = get_context().override_server_args(
enable_return_hidden_states=True,
return_hidden_states_mode="last",
)
override.install()
self.addCleanup(override.restore)
model_runner = SimpleNamespace(
is_draft_worker=False,
spec_algorithm=SimpleNamespace(is_eagle=lambda: True),
server_args=SimpleNamespace(
enable_return_hidden_states=True,
return_hidden_states_mode="last",
),
server_args=SimpleNamespace(),
)
with patch.object(
+105 -98
View File
@@ -66,7 +66,7 @@ from sglang.srt.multimodal.transport.cuda_ipc import (
DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY,
CudaIpcTensorTransportProxy,
)
from sglang.srt.runtime_context import get_context, get_parallel
from sglang.srt.runtime_context import get_context, get_parallel, publish, reset_context
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import ImageData
from sglang.test.ci.ci_register import register_cpu_ci
@@ -642,11 +642,9 @@ def _k3_preprocess_config(
)
def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls):
server_args = SimpleNamespace(
mm_feature_transport="cpu",
image_processor_backend="auto",
disable_fast_image_processor=False,
skip_tokenizer_init=False,
mm_process_config={},
mm_io_worker_num=0,
mm_processor_worker_num=0,
tokenizer_worker_num=1,
@@ -654,32 +652,35 @@ def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls
trust_mm_content_hashes=False,
base_gpu_id=0,
rl_on_policy_target=None,
allowed_media_domains=[],
media_url_max_file_size_mb=64,
)
processor = processor_cls(
hf_config=SimpleNamespace(media_placeholder_token_id=42),
server_args=server_args,
_processor=_HFProcessor(),
transport_mode=None,
)
try:
worker_processor = asyncio.run(
processor.mm_processor_executor.run(lambda *, processor: processor)
# The multimodal config comes from the bags.
with get_context().override_server_args(
mm_feature_transport="cpu", mm_process_config={}, allowed_media_domains=[]
):
processor = processor_cls(
hf_config=SimpleNamespace(media_placeholder_token_id=42),
server_args=server_args,
_processor=_HFProcessor(),
transport_mode=None,
)
assert isinstance(processor._processor, wrapper_cls)
assert isinstance(worker_processor, wrapper_cls)
assert worker_processor is not processor._processor
if processor_cls is KimiK3ImageProcessor:
fingerprint_config = processor.preprocess_fingerprint_payload()[
"wrapped_processor"
]
assert isinstance(fingerprint_config, KimiK3PreprocessConfig)
assert fingerprint_config.patch_size == 14
finally:
processor.mm_processor_executor.shutdown()
processor.io_executor.shutdown()
processor.cpu_executor.shutdown()
try:
worker_processor = asyncio.run(
processor.mm_processor_executor.run(lambda *, processor: processor)
)
assert isinstance(processor._processor, wrapper_cls)
assert isinstance(worker_processor, wrapper_cls)
assert worker_processor is not processor._processor
if processor_cls is KimiK3ImageProcessor:
fingerprint_config = processor.preprocess_fingerprint_payload()[
"wrapped_processor"
]
assert isinstance(fingerprint_config, KimiK3PreprocessConfig)
assert fingerprint_config.patch_size == 14
finally:
processor.mm_processor_executor.shutdown()
processor.io_executor.shutdown()
processor.cpu_executor.shutdown()
def test_kimi_k3_expands_image_placeholders_with_original_dimensions():
@@ -844,82 +845,88 @@ def test_kimi_k3_normal_cache_path_connects_real_producer_to_model_consumer():
tokenizer_worker_num=1,
mm_preprocess_cache_size_mb=1,
)
processor = KimiK3ImageProcessor(
hf_config=hf_config,
server_args=server_args,
_processor=hf_processor,
transport_mode=None,
)
image = Image.new("RGB", (28, 28), color=(1, 2, 3))
encoded_image = io.BytesIO()
image.save(encoded_image, format="PNG")
image_data = ImageData(
url="data:image/png;base64,"
+ base64.b64encode(encoded_image.getvalue()).decode()
)
request = SimpleNamespace(video_data=None, mm_content_hashes=None)
class _Tower(nn.Module):
device = torch.device("cpu")
patch_size = 14
def __init__(self):
super().__init__()
self.patch_embed = SimpleNamespace(
proj=SimpleNamespace(weight=torch.empty(1, dtype=torch.float32))
)
def forward(self, pixel_values, _grid_thws):
return pixel_values
model = KimiK3ForConditionalGeneration.__new__(KimiK3ForConditionalGeneration)
nn.Module.__init__(model)
model.use_data_parallel = False
model.vision_tower = _Tower()
model.mm_projector = _Projector()
publish(server_args, role="tokenizer")
try:
with (
patch(
"sglang.srt.multimodal.processors.kimi_k3.is_cuda", return_value=True
),
patch.object(
processor,
"prepare_artifact_batch",
wraps=processor.prepare_artifact_batch,
) as prepare_artifacts,
):
cold = asyncio.run(
processor.process_mm_data_async([image_data], [1, 42, 2], request)
)
hot = asyncio.run(
processor.process_mm_data_async([image_data], [3, 42, 4], request)
)
cold_items = pickle.loads(pickle.dumps(cold.mm_items))
hot_items = pickle.loads(pickle.dumps(hot.mm_items))
processor = KimiK3ImageProcessor(
hf_config=hf_config,
server_args=server_args,
_processor=hf_processor,
transport_mode=None,
)
image = Image.new("RGB", (28, 28), color=(1, 2, 3))
encoded_image = io.BytesIO()
image.save(encoded_image, format="PNG")
image_data = ImageData(
url="data:image/png;base64,"
+ base64.b64encode(encoded_image.getvalue()).decode()
)
request = SimpleNamespace(video_data=None, mm_content_hashes=None)
with (
patch("sglang.srt.models.kimi_k3.configured_tp_size", return_value=1),
patch(
"sglang.srt.multimodal.processors.kimi_k25._gpu_preprocess_images",
return_value=(
torch.ones((4, 3), dtype=torch.float32),
torch.tensor([[1, 2, 2]], dtype=torch.int64),
class _Tower(nn.Module):
device = torch.device("cpu")
patch_size = 14
def __init__(self):
super().__init__()
self.patch_embed = SimpleNamespace(
proj=SimpleNamespace(weight=torch.empty(1, dtype=torch.float32))
)
def forward(self, pixel_values, _grid_thws):
return pixel_values
model = KimiK3ForConditionalGeneration.__new__(KimiK3ForConditionalGeneration)
nn.Module.__init__(model)
model.use_data_parallel = False
model.vision_tower = _Tower()
model.mm_projector = _Projector()
try:
with (
patch(
"sglang.srt.multimodal.processors.kimi_k3.is_cuda",
return_value=True,
),
),
):
cold_features = model.get_image_feature(cold_items)
hot_features = model.get_image_feature(hot_items)
finally:
processor.shutdown()
patch.object(
processor,
"prepare_artifact_batch",
wraps=processor.prepare_artifact_batch,
) as prepare_artifacts,
):
cold = asyncio.run(
processor.process_mm_data_async([image_data], [1, 42, 2], request)
)
hot = asyncio.run(
processor.process_mm_data_async([image_data], [3, 42, 4], request)
)
cold_items = pickle.loads(pickle.dumps(cold.mm_items))
hot_items = pickle.loads(pickle.dumps(hot.mm_items))
assert prepare_artifacts.call_count == 1
assert cold.mm_items[0].hash == hot.mm_items[0].hash
assert cold.mm_items[0].offsets == hot.mm_items[0].offsets == [(3, 3)]
assert (
cold_items[0].model_specific_data[DEFERRED_PREPROCESSING_KEY].backend == "gpu"
)
torch.testing.assert_close(cold_features, hot_features)
with (
patch("sglang.srt.models.kimi_k3.configured_tp_size", return_value=1),
patch(
"sglang.srt.multimodal.processors.kimi_k25._gpu_preprocess_images",
return_value=(
torch.ones((4, 3), dtype=torch.float32),
torch.tensor([[1, 2, 2]], dtype=torch.int64),
),
),
):
cold_features = model.get_image_feature(cold_items)
hot_features = model.get_image_feature(hot_items)
finally:
processor.shutdown()
assert prepare_artifacts.call_count == 1
assert cold.mm_items[0].hash == hot.mm_items[0].hash
assert cold.mm_items[0].offsets == hot.mm_items[0].offsets == [(3, 3)]
assert (
cold_items[0].model_specific_data[DEFERRED_PREPROCESSING_KEY].backend
== "gpu"
)
torch.testing.assert_close(cold_features, hot_features)
finally:
reset_context()
def test_kimi_k3_model_accepts_mixed_cached_eager_and_deferred_artifacts():
@@ -21,11 +21,13 @@ maybe_stub_sgl_kernel()
from sglang.srt.multimodal.processors.qwen_vl import ( # noqa: E402
QwenVLImageProcessor,
)
from sglang.srt.runtime_context import publish, reset_context # noqa: E402
from sglang.srt.server_args import ServerArgs # noqa: E402
register_cpu_ci(est_time=0, suite="base-a-test-cpu", disabled="Qwen test fixtures")
def make_processor(config, image_processor_cls=None):
def make_processor(case, config, image_processor_cls=None):
"""A ``QwenVLImageProcessor`` over a tiny hand-built tokenizer.
``image_processor_cls`` picks the HF backend; they resample differently."""
image_processor_cls = image_processor_cls or HfQwenImageProcessor
@@ -87,6 +89,18 @@ def make_processor(config, image_processor_cls=None):
allowed_media_domains=[],
media_url_max_file_size_mb=64,
)
# The processor reads its media policy, transport and per-modality limits
# from the mm bag, so the fixture publishes before building it.
publish(
ServerArgs(
model_path="dummy",
mm_feature_transport=server_args.mm_feature_transport,
mm_process_config=server_args.mm_process_config,
allowed_media_domains=server_args.allowed_media_domains,
),
role="tokenizer",
)
case.addCleanup(reset_context)
return QwenVLImageProcessor(
hf_config, server_args, processor, None, skip_mm_pool=True
)
@@ -50,7 +50,9 @@ class TestQwenE2eParity(CustomTestCase):
import transformers.models.qwen2_vl as qwen2_vl
self.processor = make_processor(
PROCESSOR_CONFIGS["qwen2_5_vl"], getattr(qwen2_vl, self.image_processor)
self,
PROCESSOR_CONFIGS["qwen2_5_vl"],
getattr(qwen2_vl, self.image_processor),
)
def tearDown(self):
@@ -50,7 +50,7 @@ class TestQwenNativeMmHashes(CustomTestCase):
from sglang.srt.managers.multimodal_processor import import_processors
import_processors("sglang.srt.multimodal.processors")
self.processor = make_processor(PROCESSOR_CONFIGS["qwen2_5_vl"])
self.processor = make_processor(self, PROCESSOR_CONFIGS["qwen2_5_vl"])
def tearDown(self):
self.processor.io_executor.shutdown()
@@ -42,6 +42,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
def test_model_class_controls_cuda_vmm_opt_in(self):
from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.runtime_context import get_context
class SupportedModel:
supports_cuda_vmm_feature_transport = True
@@ -49,8 +50,10 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
class UnsupportedModel:
pass
override = get_context().override_server_args(mm_feature_transport="cuda_vmm")
override.install()
self.addCleanup(override.restore)
manager = object.__new__(TokenizerManager)
manager.server_args = SimpleNamespace(mm_feature_transport="cuda_vmm")
manager.model_config = object()
with patch(
@@ -70,9 +73,12 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
def test_cpu_transport_skips_model_opt_in_lookup(self):
from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.runtime_context import get_context
override = get_context().override_server_args(mm_feature_transport="cpu")
override.install()
self.addCleanup(override.restore)
manager = object.__new__(TokenizerManager)
manager.server_args = SimpleNamespace(mm_feature_transport="cpu")
manager.model_config = object()
with patch(
@@ -21,6 +21,7 @@ from sglang.srt.multimodal.cache import (
resolve_multimodal_item_hash,
snapshot_media,
)
from sglang.srt.runtime_context import publish, reset_context
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci
@@ -216,26 +217,32 @@ class TestMediaIdentity(unittest.TestCase):
return {"model_type": "vlm", "architectures": ["VLM"]}
config = Config()
args = ServerArgs(
model_path="dummy",
revision="model-revision",
disable_fast_image_processor=False,
mm_process_config={"image": {"max_pixels": 1024}},
)
base = build_processor_fingerprint(Processor("gpu"), config, args)
changed_backend = build_processor_fingerprint(Processor("cpu"), config, args)
changed_args = ServerArgs(
model_path="dummy",
revision="model-revision",
disable_fast_image_processor=False,
mm_process_config={"image": {"max_pixels": 2048}},
)
changed_config = build_processor_fingerprint(
Processor("gpu"), config, changed_args
)
def fingerprint(processor, mm_process_config):
# The digest reads the effective config, so the test publishes it
# rather than handing one in: that is the only source the function
# has, and two callers with the same effective config must agree.
publish(
ServerArgs(
model_path="dummy",
revision="model-revision",
disable_fast_image_processor=False,
mm_process_config=mm_process_config,
),
role="test",
)
return build_processor_fingerprint(processor, config)
self.addCleanup(reset_context)
small = {"image": {"max_pixels": 1024}}
base = fingerprint(Processor("gpu"), small)
changed_backend = fingerprint(Processor("cpu"), small)
changed_config = fingerprint(Processor("gpu"), {"image": {"max_pixels": 2048}})
same_again = fingerprint(Processor("gpu"), small)
self.assertNotEqual(base, changed_backend)
self.assertNotEqual(base, changed_config)
self.assertEqual(base, same_again)
def test_item_hash_namespace_covers_identity_and_processor_output(self):
digest = snapshot_media(b"image").content_digest
@@ -2161,6 +2161,7 @@ class TestGrpcServerArgs(CustomTestCase):
tokenizer_manager=MagicMock(),
template_manager=MagicMock(),
scheduler_info={},
grpc_port=server_args.grpc_port,
)
self.assertEqual(handle, "handle")
@@ -52,6 +52,23 @@ _SLOT_OWNERS = ("srt/runtime_context.py", "srt/server_args.py", "srt/arg_groups/
# live topology cannot answer there. The test below asserts this map is exactly
# the set of call sites, so the reasons cannot drift away from the code.
_CONFIGURED_SIZE_CALL_SITES = {
("srt/entrypoints/engine.py", "configured_pp_size"): (
"the launch path decides how many scheduler processes to spawn; it runs "
"before any of them exists, so there is no group to ask"
),
("srt/ray/engine.py", "configured_pp_size"): (
"the Ray driver sizes the actor placement group; the actors it is about "
"to create are the ones that will hold the process groups"
),
("srt/ray/data_parallel_controller.py", "configured_pp_size"): (
"same placement arithmetic on the DP path -- ranks per TP group, "
"computed in the driver before the actors start"
),
("srt/ray/data_parallel_controller.py", "configured_attn_cp_size"): (
"the attention-CP factor of that same placement arithmetic, and the one "
"size whose live value cannot express the configured intent when "
"attn_cp_size > moe_dp_size aliases the groups"
),
("srt/layers/attention/dsa/dsa_indexer.py", "configured_pp_size"): (
"gates `pp_size > 1 and not get_pp_group()...`; the short circuit is the "
"point, since with PP off the group is never touched, which is what lets "
@@ -91,6 +108,10 @@ _CONFIGURED_SIZE_CALL_SITES = {
"the consumer count is configured fan-out arithmetic (tp_size // "
"dp_size), which is what the record answered before"
),
("srt/disaggregation/encode_server.py", "configured_tp_size"): (
"the encode server's launch entry sizes its workers before it has "
"spawned any of them"
),
("srt/model_loader/loader.py", "configured_moe_dp_size"): (
"the same dict already carries the live moe_dp_size under 'dp'; this entry "
"is the configured intent"
@@ -0,0 +1,280 @@
"""Launch paths read the configured parallel sizes, not the live ones.
`get_parallel().pp_size` and its four siblings are read-through properties over
the process groups, so they answer only after distributed init. The launcher
decides how many processes to spawn *before* that, and a live read there raises
`Distributed environment is not initialized` -- a startup crash no unit test
reaches, because nothing short of booting a server runs the launcher.
"""
import ast
import pathlib
import unittest
import sglang
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=9, suite="base-a-test-cpu")
_PACKAGE_ROOT = pathlib.Path(sglang.__file__).resolve().parent
# Live-shadowed sizes a launch path is known to have read. ParallelContext
# shadows more properties than these (every `_v(name, ...)` one raises the same
# "Distributed environment is not initialized"); this dict carries the ones a
# `configured_*` accessor answers, so it is a remedy map, not a census.
_LIVE_SHADOWED = {
"tp_size": "configured_tp_size()",
"pp_size": "configured_pp_size()",
"moe_dp_size": "configured_moe_dp_size()",
"attn_cp_size": "configured_attn_cp_size()",
"dcp_size": "a configured accessor (none exists yet; add one beside configured_pp_size)",
}
# Launch paths that decide how many children to spawn are derived below
# from the spawn itself. These launch without a size-driven spawn, so no
# derivation reaches them and they are carried by hand.
_HAND_CARRIED = (
"srt/entrypoints/http_server.py",
"srt/entrypoints/sidecar.py",
"srt/ray/data_parallel_controller.py",
"srt/ray/engine.py",
"srt/ray/http_server.py",
)
def _multiprocessing_names(tree):
"""Names bound to multiprocessing, to one of its start contexts, or to the
process constructors themselves."""
modules, constructors = set(), set()
for node in ast.walk(tree):
if isinstance(node, ast.Import):
for a in node.names:
if a.name == "multiprocessing" or a.name.startswith("multiprocessing."):
modules.add(a.asname or a.name.split(".")[0])
elif a.name == "torch.multiprocessing":
modules.add(a.asname or "torch")
elif isinstance(node, ast.ImportFrom):
if node.module in (
"multiprocessing",
"multiprocessing.context",
"torch.multiprocessing",
):
constructors |= {
a.asname or a.name for a in node.names if a.name == "Process"
}
elif node.module == "concurrent.futures":
constructors |= {
a.asname or a.name
for a in node.names
if a.name == "ProcessPoolExecutor"
}
for node in ast.walk(tree):
if isinstance(node, ast.Assign) and isinstance(node.value, ast.Call):
func = node.value.func
if (
isinstance(func, ast.Attribute)
and func.attr == "get_context"
and isinstance(func.value, ast.Name)
and func.value.id in modules
):
modules |= {t.id for t in node.targets if isinstance(t, ast.Name)}
return modules, constructors
def _configured_accessors() -> frozenset:
"""The `configured_*_size()` names `runtime_context` exports.
Derived from that module, so a new accessor keeps its launcher watched
without a second list here.
"""
tree = ast.parse(
(_PACKAGE_ROOT / "srt/runtime_context.py").read_text(encoding="utf-8-sig")
)
names = frozenset(
node.name
for node in tree.body
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef))
and node.name.startswith("configured_")
and node.name.endswith("_size")
)
assert names, (
"no configured_*_size accessors found in runtime_context; the "
"derivation is broken, not the tree"
)
return names
def _spawns_from_a_size(tree) -> bool:
"""Does any function here construct a child process *and* read one of the
five sizes -- live off the parallel bag, or through its `configured_*_size()`
answer? That is a spawn count decided from the topology.
Counting the configured read too is what keeps a launcher watched after it
is converted. Deriving on the live read alone means the file drops out of
the scan the moment it stops offending, so the guard would only ever watch
the launchers that already fail it.
"""
configured = _configured_accessors()
modules, constructors = _multiprocessing_names(tree)
names, qualified = _parallel_bag_names(tree)
aliases = _bag_aliases(tree, names, qualified)
for fn in ast.walk(tree):
if not isinstance(fn, (ast.FunctionDef, ast.AsyncFunctionDef)):
continue
spawns = reads = False
for node in ast.walk(fn):
if isinstance(node, ast.Call):
func = node.func
if isinstance(func, ast.Attribute) and func.attr in (
"Process",
"ProcessPoolExecutor",
"Popen",
"spawn",
):
# `mp.Process(`, `mp.get_context("spawn").Process(` and
# `subprocess.Popen(` all reach a child process; the
# receiver of a chained call is itself a call, so this
# cannot require a bare Name.
spawns = True
elif isinstance(func, ast.Name) and func.id in constructors:
spawns = True
if (isinstance(func, ast.Name) and func.id in configured) or (
isinstance(func, ast.Attribute) and func.attr in configured
):
reads = True
elif (
isinstance(node, ast.Attribute)
and node.attr in _LIVE_SHADOWED
and (
_is_parallel_bag_call(node.value, names, qualified)
or (isinstance(node.value, ast.Name) and node.value.id in aliases)
)
):
# A record read (`server_args.tp_size`) sizes a spawn too, but
# it cannot raise pre-dist; only the bag read is this guard's
# subject, so only it forces a module into _PRE_DIST.
reads = True
if spawns and reads:
return True
return False
def _parallel_bag_names(tree):
"""What this module calls `get_parallel`, plus any runtime_context alias.
A literal-name match reads only one spelling; an aliased import or a
module-qualified call is the same read with a different surface.
"""
names, modules = set(), set()
for node in ast.walk(tree):
if (
isinstance(node, ast.ImportFrom)
and node.module
and node.module.endswith("runtime_context")
):
names |= {
a.asname or a.name for a in node.names if a.name == "get_parallel"
}
elif isinstance(node, ast.Import):
for a in node.names:
if a.name.endswith("runtime_context"):
modules.add(a.asname or a.name.split(".")[0])
return names, modules
def _is_parallel_bag_call(node, names, modules) -> bool:
if not isinstance(node, ast.Call):
return False
if isinstance(node.func, ast.Name):
return node.func.id in names
return (
isinstance(node.func, ast.Attribute)
and node.func.attr == "get_parallel"
and isinstance(node.func.value, ast.Name)
and node.func.value.id in modules
)
def _bag_aliases(tree, names, qualified):
"""Locals bound to the parallel bag: `p = get_parallel()` then `p.pp_size`
is the same read one line later."""
return {
target.id
for node in ast.walk(tree)
if isinstance(node, ast.Assign)
and _is_parallel_bag_call(node.value, names, qualified)
for target in node.targets
if isinstance(target, ast.Name)
}
def _launch_paths():
"""(relative path, tree) per module that runs before its process groups.
A module that sizes a spawn loop from a parallel-bag size is derived from
the spawn itself; `_HAND_CARRIED` holds the launch entries that spawn
nothing, which no derivation can reach.
"""
seen = {}
for path in sorted(_PACKAGE_ROOT.rglob("*.py")):
source = path.read_text()
# Every spawn shape below names Process, ProcessPoolExecutor or Popen.
if not any(name in source for name in ("Process", "Popen", "spawn")):
continue
try:
tree = ast.parse(source)
except SyntaxError:
continue
if _spawns_from_a_size(tree):
seen[str(path.relative_to(_PACKAGE_ROOT))] = tree
for rel in _HAND_CARRIED:
seen.setdefault(rel, ast.parse((_PACKAGE_ROOT / rel).read_text()))
return sorted(seen.items())
class TestLaunchPathsReadConfiguredSizes(CustomTestCase):
def test_no_live_topology_read_before_distributed_init(self):
offenders = []
for rel, tree in _launch_paths():
names, modules = _parallel_bag_names(tree)
aliases = _bag_aliases(tree, names, modules)
for node in ast.walk(tree):
if isinstance(node, ast.Attribute) and node.attr in _LIVE_SHADOWED:
base = node.value
if _is_parallel_bag_call(base, names, modules) or (
isinstance(base, ast.Name) and base.id in aliases
):
offenders.append(
f"{rel}:{node.lineno} reads the live {node.attr}; "
f"use {_LIVE_SHADOWED[node.attr]}"
)
elif (
isinstance(node, ast.Call)
and isinstance(node.func, ast.Name)
and node.func.id == "getattr"
and len(node.args) >= 2
and isinstance(node.args[1], ast.Constant)
and node.args[1].value in _LIVE_SHADOWED
and (
_is_parallel_bag_call(node.args[0], names, modules)
or (
isinstance(node.args[0], ast.Name)
and node.args[0].id in aliases
)
)
):
offenders.append(
f"{rel}:{node.lineno} reads the live "
f"{node.args[1].value} through getattr; "
f"use {_LIVE_SHADOWED[node.args[1].value]}"
)
self.assertEqual(
offenders,
[],
"launch paths run before distributed init:\n " + "\n ".join(offenders),
)
if __name__ == "__main__":
unittest.main()
@@ -104,9 +104,6 @@ _UNREAD_ENTRIES: dict = {
("multimodal_gen/test/unit/test_disagg_trace.py", "_srt_trace_server_args"): (
"a trace fixture publishing its own context"
),
("srt/managers/detokenizer_manager.py", "run_detokenizer_process"): (
"DetokenizerManager reads the handed instance at this revision"
),
}
# `publish` itself and its named wrappers live here; a call inside them is the
@@ -49,6 +49,7 @@ _PAIR_READERS = {
"batch_overlap/two_batch_overlap.py": "prefill (extend positions)",
"managers/scheduler.py": "prefill (truncation align knobs)",
"entrypoints/engine.py": "either half (flashinfer version floor)",
"models/sarvam_moe.py": "the half serving the forward (attn dispatch)",
}
@@ -129,6 +129,7 @@ _ENV_MATRIX = (({}, {"SGLANG_IS_IN_CI": "true"}),)
_PASSED = frozenset({"model_path", "device", "random_seed"})
_EXPOSED = {
("entrypoints/sidecar.py", "grpc_port"),
("configs/embedding_model_spec.py", "chunked_prefill_size"),
("configs/embedding_model_spec.py", "cuda_graph_config"),
("configs/embedding_model_spec.py", "disable_radix_cache"),
@@ -143,24 +144,10 @@ _EXPOSED = {
("configs/model_config.py", "quantization"),
("configs/model_config.py", "speculative_algorithm"),
("configs/model_config.py", "speculative_draft_model_quantization"),
("constrained/base_grammar_backend.py", "grammar_backend"),
("constrained/base_grammar_backend.py", "reasoning_parser"),
("disaggregation/common/conn.py", "disaggregation_bootstrap_port"),
("disaggregation/common/conn.py", "pp_size"),
("disaggregation/decode_kvcache_offload_manager.py", "hicache_io_backend"),
("disaggregation/decode_kvcache_offload_manager.py", "served_model_name"),
("disaggregation/encode_receiver.py", "disaggregation_ib_device"),
("disaggregation/encode_receiver.py", "encoder_transfer_backend"),
("disaggregation/encode_receiver.py", "mooncake_ib_device"),
("disaggregation/encode_receiver.py", "tokenizer_path"),
("disaggregation/encode_server.py", "allowed_media_domains"),
("disaggregation/encode_server.py", "device"),
("disaggregation/encode_server.py", "encoder_transfer_backend"),
("disaggregation/encode_server.py", "load_format"),
("disaggregation/encode_server.py", "mm_process_config"),
("disaggregation/encode_server.py", "model_path"),
("disaggregation/encode_server.py", "served_model_name"),
("disaggregation/encode_server.py", "tokenizer_path"),
("disaggregation/utils.py", "disaggregation_transfer_backend"),
("distributed/bootstrap.py", "disable_custom_all_reduce"),
("distributed/bootstrap.py", "enable_symm_mem"),
@@ -201,28 +188,14 @@ _EXPOSED = {
("elastic_ep/expert_backup_manager.py", "load_format"),
("elastic_ep/expert_backup_manager.py", "mooncake_ib_device"),
("entrypoints/engine.py", "attn_cp_size"),
("entrypoints/engine.py", "detokenizer_worker_num"),
("entrypoints/engine.py", "dtype"),
("entrypoints/engine.py", "enable_symm_mem"),
("entrypoints/engine.py", "ep_join_mode"),
("entrypoints/engine.py", "load_format"),
("entrypoints/engine.py", "model_path"),
("entrypoints/engine.py", "moe_dp_size"),
("entrypoints/engine.py", "pp_size"),
("entrypoints/engine.py", "quantization"),
("entrypoints/engine.py", "reasoning_parser"),
(
"entrypoints/engine.py",
"remote_instance_weight_loader_start_seed_via_transfer_engine",
),
("entrypoints/engine.py", "tool_call_parser"),
("entrypoints/http_server.py", "disaggregation_mode"),
("entrypoints/http_server.py", "ep_join_mode"),
("entrypoints/http_server.py", "grpc_port"),
("entrypoints/http_server.py", "model_path"),
("entrypoints/http_server.py", "served_model_name"),
("entrypoints/http_server.py", "skip_server_warmup"),
("entrypoints/sidecar.py", "grpc_port"),
("eplb/eplb_manager.py", "ep_dispatch_algorithm"),
("eplb/eplb_manager.py", "expert_distribution_recorder_buffer_size"),
("eplb/expert_distribution.py", "deepep_mode"),
@@ -236,7 +209,6 @@ _EXPOSED = {
("kv_canary/capacities.py", "chunked_prefill_size"),
("kv_canary/capacities.py", "cuda_graph_config"),
("kv_canary/capacities.py", "speculative_num_draft_tokens"),
("kv_canary/token_oracle/install.py", "sampling_backend"),
("layers/cp/base.py", "attn_cp_size"),
("layers/cp/base.py", "cp_strategy"),
("layers/cp/base.py", "enable_prefill_cp"),
@@ -259,9 +231,6 @@ _EXPOSED = {
("managers/data_parallel_controller.py", "moe_dp_size"),
("managers/data_parallel_controller.py", "pp_size"),
("managers/data_parallel_controller.py", "soft_watchdog_timeout"),
("managers/detokenizer_manager.py", "soft_watchdog_timeout"),
("managers/detokenizer_manager.py", "tokenizer_path"),
("managers/detokenizer_manager.py", "tool_call_parser"),
("managers/disagg_service.py", "disaggregation_bootstrap_port"),
("managers/disagg_service.py", "disaggregation_mode"),
("managers/disagg_service.py", "disaggregation_transfer_backend"),
@@ -279,27 +248,7 @@ _EXPOSED = {
("managers/scheduler.py", "pp_size"),
("managers/scheduler.py", "soft_watchdog_timeout"),
("managers/scheduler.py", "speculative_algorithm"),
(
"managers/scheduler_components/new_token_ratio_tracker.py",
"schedule_conservativeness",
),
("managers/tokenizer_manager.py", "disable_radix_cache"),
("managers/tokenizer_manager.py", "disaggregation_mode"),
("managers/tokenizer_manager.py", "disaggregation_transfer_backend"),
("managers/tokenizer_manager.py", "enable_lora"),
("managers/tokenizer_manager.py", "enable_tokenizer_batch_encode"),
("managers/tokenizer_manager.py", "encoder_transfer_backend"),
("managers/tokenizer_manager.py", "limit_mm_data_per_request"),
("managers/tokenizer_manager.py", "lora_paths"),
("managers/tokenizer_manager.py", "mm_feature_transport"),
("managers/tokenizer_manager.py", "model_path"),
("managers/tokenizer_manager.py", "preferred_sampling_params"),
("managers/tokenizer_manager.py", "return_hidden_states_mode"),
("managers/tokenizer_manager.py", "served_model_name"),
("managers/tokenizer_manager.py", "soft_watchdog_timeout"),
("managers/tokenizer_manager.py", "speculative_algorithm"),
("managers/tokenizer_manager.py", "speculative_num_draft_tokens"),
("managers/tokenizer_manager.py", "tokenizer_path"),
("managers/tp_worker.py", "disable_overlap_schedule"),
("managers/tp_worker.py", "model_path"),
("managers/tp_worker.py", "random_seed"),
@@ -313,8 +262,6 @@ _EXPOSED = {
("mem_cache/hybrid_cache/hybrid_pool_assembler.py", "served_model_name"),
("mem_cache/kv_cache_builder.py", "hicache_mem_layout"),
("mem_cache/radix_cache_cpp.py", "enable_hierarchical_cache"),
("model_executor/forward_batch_info.py", "enable_return_hidden_states"),
("model_executor/forward_batch_info.py", "return_hidden_states_mode"),
("model_executor/model_runner.py", "device"),
("model_executor/model_runner.py", "speculative_algorithm"),
("model_executor/model_runner.py", "speculative_draft_attention_backend"),
@@ -353,22 +300,11 @@ _EXPOSED = {
"model_executor/runner_backend/tc_piecewise_cuda_graph_backend.py",
"cuda_graph_config",
),
("models/sarvam_moe.py", "attention_backend"),
("models/sarvam_moe.py", "decode_attention_backend"),
("models/sarvam_moe.py", "prefill_attention_backend"),
("multimodal/cache/identity.py", "mm_process_config"),
("multimodal/processors/base_processor.py", "allowed_media_domains"),
("multimodal/processors/base_processor.py", "image_processor_backend"),
("multimodal/processors/base_processor.py", "mm_feature_transport"),
("multimodal/processors/base_processor.py", "mm_process_config"),
("multimodal/processors/mimo_v2.py", "device"),
("observability/metrics_collector.py", "disaggregation_mode"),
("observability/metrics_collector.py", "prefill_delayer_max_delay_passes"),
("observability/metrics_collector.py", "served_model_name"),
("parser/template_detection.py", "model_path"),
("ray/data_parallel_controller.py", "attn_cp_size"),
("ray/data_parallel_controller.py", "pp_size"),
("ray/engine.py", "pp_size"),
("speculative/adaptive_spec_params.py", "speculative_algorithm"),
("speculative/adaptive_spec_params.py", "speculative_eagle_topk"),
("speculative/dflash_worker_v2.py", "speculative_draft_window_size"),
@@ -435,29 +371,20 @@ _EXPOSED_CUDA_ONLY: frozenset = frozenset()
_OVERRIDDEN_AND_READ = {
("configs/model_config.py", "dtype"),
("configs/model_config.py", "model_path"),
("constrained/base_grammar_backend.py", "grammar_backend"),
("disaggregation/decode_kvcache_offload_manager.py", "hicache_storage_backend"),
(
"disaggregation/decode_kvcache_offload_manager.py",
"hicache_storage_backend_extra_config",
),
("disaggregation/encode_server.py", "load_format"),
("disaggregation/encode_server.py", "model_path"),
(
"distributed/device_communicators/mooncake_transfer_engine.py",
"hicache_storage_backend",
),
("dllm/config.py", "model_path"),
("elastic_ep/expert_backup_manager.py", "load_format"),
("entrypoints/engine.py", "dtype"),
("entrypoints/engine.py", "load_format"),
("entrypoints/engine.py", "model_path"),
("entrypoints/http_server.py", "model_path"),
("kv_canary/api.py", "speculative_num_steps"),
("kv_canary/capacities.py", "speculative_num_draft_tokens"),
("managers/scheduler.py", "hicache_storage_backend"),
("managers/tokenizer_manager.py", "model_path"),
("managers/tokenizer_manager.py", "speculative_num_draft_tokens"),
("managers/tp_worker.py", "model_path"),
("mem_cache/hiradix_cache.py", "hicache_storage_backend"),
("mem_cache/hiradix_cache.py", "hicache_storage_backend_extra_config"),