config: the DP/EP topology reads come from the parallel bag (#35025)

This commit is contained in:
Cheng Wan
2026-08-17 16:17:26 -07:00
committed by GitHub
parent d2bc697396
commit a97bc8db32
33 changed files with 332 additions and 304 deletions
@@ -13,6 +13,7 @@ import torch
from sglang.benchmark.serving import run_benchmark
from sglang.srt.managers.prefill_delayer import PrefillDelayer
from sglang.srt.runtime_context import get_context
from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.run_eval import run_eval
@@ -76,6 +77,9 @@ def _run_negotiate_test(rank, test_cases):
cpu_group = torch.distributed.new_group(backend="gloo")
for case in test_cases:
# The DP-attention gate is a published config leaf.
override = get_context().override_server_args(enable_dp_attention=True)
override.install()
delayer = PrefillDelayer(
dp_size=world_size,
attn_tp_size=1,
@@ -127,6 +131,8 @@ def _run_negotiate_test(rank, test_cases):
result.wait_seconds > 0.0
), f"Case {case.name} rank {rank}: wait_seconds not surfaced"
override.restore()
_NEGOTIATE_TEST_CASES = [
NegotiateTestCase(
@@ -14,6 +14,12 @@ register_cpu_ci(est_time=5, suite="base-a-test-cpu")
class TestDisaggregationServerWarmup(unittest.IsolatedAsyncioTestCase):
async def test_sends_concurrent_scalar_request_to_each_dp_rank(self):
from sglang.srt.runtime_context import get_context
# The warmup fan-out width comes from the published topology.
override = get_context().override_server_args(dp_size=4)
override.install()
self.addCleanup(override.restore)
server_args = SimpleNamespace(dp_size=4)
all_started = asyncio.Event()
calls = []
@@ -5,6 +5,7 @@ import tempfile
import time
import unittest
from types import SimpleNamespace
from unittest import mock
from sglang.srt.managers.load_snapshot import (
LoadSnapshot,
@@ -18,6 +19,7 @@ from sglang.srt.managers.load_snapshot import (
should_use_zmq,
zmq_reader_owner,
)
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, maybe_stub_sgl_kernel
@@ -228,23 +230,17 @@ class TestZmqRoundTrip(CustomTestCase):
class TestFactoryFunctions(CustomTestCase):
def _publish(self, **fields):
override = get_context().override_server_args(**fields)
override.install()
self.addCleanup(override.restore)
def test_shm_mode(self):
server_args = SimpleNamespace(
enable_dp_attention=False,
nnodes=1,
dp_size=1,
load_balance_method="round_robin",
node_rank=0,
tokenizer_worker_num=1,
)
self._publish(enable_dp_attention=False, nnodes=1, dp_size=1)
port_args = SimpleNamespace(instance_id="test_shm_factory")
writer = create_load_snapshot_writer(
server_args, port_args, dp_size=1, dp_rank=0
)
writer = create_load_snapshot_writer(port_args, dp_size=1, dp_rank=0)
self.assertIsInstance(writer, ShmLoadSnapshotWriter)
reader = create_load_snapshot_reader(
server_args, port_args, caller="TokenizerManager"
)
reader = create_load_snapshot_reader(port_args, caller="TokenizerManager")
self.assertIsInstance(reader, ShmLoadSnapshotReader)
reader.close()
writer.close()
@@ -255,24 +251,13 @@ class TestFactoryFunctions(CustomTestCase):
os.unlink(path)
def test_zmq_mode_via_env(self):
server_args = SimpleNamespace(
enable_dp_attention=False,
nnodes=1,
dp_size=1,
load_balance_method="round_robin",
node_rank=0,
tokenizer_worker_num=1,
)
self._publish(enable_dp_attention=False, nnodes=1, dp_size=1)
port_args = SimpleNamespace(instance_id="test_zmq_factory")
os.environ["SGLANG_LOAD_SNAPSHOT_USE_ZMQ"] = "1"
try:
writer = create_load_snapshot_writer(
server_args, port_args, dp_size=1, dp_rank=0
)
writer = create_load_snapshot_writer(port_args, dp_size=1, dp_rank=0)
self.assertIsInstance(writer, ZmqLoadSnapshotWriter)
reader = create_load_snapshot_reader(
server_args, port_args, caller="TokenizerManager"
)
reader = create_load_snapshot_reader(port_args, caller="TokenizerManager")
self.assertIsInstance(reader, ZmqShmLoadSnapshotReader)
reader.close()
writer.close()
@@ -280,8 +265,8 @@ class TestFactoryFunctions(CustomTestCase):
del os.environ["SGLANG_LOAD_SNAPSHOT_USE_ZMQ"]
def test_should_use_zmq_multinode_dp_attention(self):
args = SimpleNamespace(enable_dp_attention=True, nnodes=2)
self.assertTrue(should_use_zmq(args))
self._publish(enable_dp_attention=True, nnodes=2, dp_size=2)
self.assertTrue(should_use_zmq())
class TestZmqReaderOwner(CustomTestCase):
@@ -289,9 +274,9 @@ class TestZmqReaderOwner(CustomTestCase):
CALLERS = ("TokenizerManager", "MultiTokenizerRouter", "DataParallelController")
@staticmethod
def _args(**overrides):
base = dict(
def _owners(self, **overrides):
"""Publish a config and return the callers that claim the socket."""
fields = dict(
enable_dp_attention=True,
nnodes=2,
node_rank=0,
@@ -299,54 +284,87 @@ class TestZmqReaderOwner(CustomTestCase):
load_balance_method="round_robin",
tokenizer_worker_num=1,
)
base.update(overrides)
return SimpleNamespace(**base)
def _owners(self, args):
return {c for c in self.CALLERS if zmq_reader_owner(args, c)}
fields.update(overrides)
override = get_context().override_server_args(**fields)
override.install()
try:
return {c for c in self.CALLERS if zmq_reader_owner(c)}
finally:
override.restore()
def test_zmq_disabled_no_owner(self):
args = self._args(enable_dp_attention=False, nnodes=1)
self.assertEqual(self._owners(args), set())
self.assertEqual(self._owners(enable_dp_attention=False, nnodes=1), set())
def test_non_zero_node_rank_no_owner(self):
args = self._args(node_rank=1, dp_size=4, tokenizer_worker_num=8)
self.assertEqual(self._owners(args), set())
self.assertEqual(
self._owners(node_rank=1, dp_size=4, tokenizer_worker_num=8), set()
)
def test_tokenizer_manager_owns_when_dp1(self):
self.assertEqual(self._owners(self._args(dp_size=1)), {"TokenizerManager"})
self.assertEqual(self._owners(dp_size=1), {"TokenizerManager"})
def test_multi_tokenizer_router_owns_in_multi_tokenizer_dp1(self):
args = self._args(dp_size=1, tokenizer_worker_num=8)
self.assertEqual(self._owners(args), {"MultiTokenizerRouter"})
self.assertEqual(
self._owners(dp_size=1, tokenizer_worker_num=8), {"MultiTokenizerRouter"}
)
def test_multi_tokenizer_router_owns_in_multi_tokenizer_round_robin(self):
args = self._args(dp_size=4, tokenizer_worker_num=8)
self.assertEqual(self._owners(args), {"MultiTokenizerRouter"})
self.assertEqual(
self._owners(dp_size=4, tokenizer_worker_num=8), {"MultiTokenizerRouter"}
)
def test_data_parallel_controller_owns_load_aware(self):
for method in ("total_tokens", "total_requests"):
args = self._args(
dp_size=4, tokenizer_worker_num=8, load_balance_method=method
self.assertEqual(
self._owners(
dp_size=4, tokenizer_worker_num=8, load_balance_method=method
),
{"DataParallelController"},
)
self.assertEqual(self._owners(args), {"DataParallelController"})
def test_tokenizer_manager_owns_dp4_round_robin(self):
args = self._args(dp_size=4, tokenizer_worker_num=1)
self.assertEqual(self._owners(args), {"TokenizerManager"})
self.assertEqual(
self._owners(dp_size=4, tokenizer_worker_num=1), {"TokenizerManager"}
)
def test_the_controller_answers_within_its_audited_namespaces(self):
"""The DP controller publishes with a narrowed namespace set.
Under `SGLANG_ROLE_NAMESPACES=enforce` a read outside that set raises,
and this decision runs during its startup -- so asking which tokenizer
process owns the socket would abort the controller before it has a
reader.
"""
import sglang.srt.runtime_context as rc
fields = dict(
enable_dp_attention=True,
nnodes=2,
node_rank=0,
dp_size=4,
load_balance_method="total_tokens",
tokenizer_worker_num=8,
)
override = get_context().override_server_args(**fields)
override.install()
self.addCleanup(override.restore)
with mock.patch.object(rc, "_ROLE_NS_MODE", "enforce"), mock.patch.object(
rc._CONTEXT, "_publish_role", "dp_controller"
):
self.assertTrue(zmq_reader_owner("DataParallelController"))
def test_at_most_one_owner_across_configs(self):
for dp_size in (1, 4):
for tw in (1, 8):
for method in ("round_robin", "total_tokens", "total_requests"):
for node_rank in (0, 1):
args = self._args(
owners = self._owners(
dp_size=dp_size,
tokenizer_worker_num=tw,
load_balance_method=method,
node_rank=node_rank,
)
self.assertLessEqual(len(self._owners(args)), 1, args)
self.assertLessEqual(len(owners), 1, owners)
class TestZmqAddr(CustomTestCase):
@@ -1,6 +1,7 @@
import unittest
from types import SimpleNamespace
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, maybe_stub_sgl_kernel
@@ -14,11 +15,15 @@ from sglang.srt.model_executor.forward_batch_info import ForwardMode # noqa: E4
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
def _server_args(interval, enable_dp_attention=False):
return SimpleNamespace(
def _publish(case, interval, enable_dp_attention=False):
"""Publish the config the skipper reads; the double is a published config,
not an injected object."""
override = get_context().override_server_args(
scheduler_recv_interval=interval,
enable_dp_attention=enable_dp_attention,
)
override.install()
case.addCleanup(override.restore)
def _batch(forward_mode, recv_skipper_forward_mode=None):
@@ -31,22 +36,24 @@ def _batch(forward_mode, recv_skipper_forward_mode=None):
class TestSchedulerRecvSkipper(CustomTestCase):
def test_disabled_at_default_interval(self):
# interval <= 1 disables the skipper entirely.
self.assertIsNone(SchedulerRecvSkipper.maybe_create(_server_args(1)))
_publish(self, 1)
self.assertIsNone(SchedulerRecvSkipper.maybe_create())
def test_enabled_under_dp_attention(self):
# Regression: the constructor used to assert `not enable_dp_attention`.
skipper = SchedulerRecvSkipper.maybe_create(
_server_args(50, enable_dp_attention=True)
)
_publish(self, 50, enable_dp_attention=True)
skipper = SchedulerRecvSkipper.maybe_create()
self.assertIsNotNone(skipper)
def test_no_last_batch_accumulates_slowly(self):
skipper = SchedulerRecvSkipper.maybe_create(_server_args(50))
_publish(self, 50)
skipper = SchedulerRecvSkipper.maybe_create()
self.assertFalse(skipper.handle(None))
def test_decode_accumulates_until_threshold(self):
# DECODE weight = 1: recv only every `interval` decode steps.
skipper = SchedulerRecvSkipper.maybe_create(_server_args(3))
_publish(self, 3)
skipper = SchedulerRecvSkipper.maybe_create()
decode = _batch(ForwardMode.DECODE)
self.assertFalse(skipper.handle(decode)) # counter 1
self.assertFalse(skipper.handle(decode)) # counter 2
@@ -55,21 +62,20 @@ class TestSchedulerRecvSkipper(CustomTestCase):
def test_prefill_triggers_recv_immediately(self):
# Non-decode passes use the large default weight: recv right away.
skipper = SchedulerRecvSkipper.maybe_create(_server_args(50))
_publish(self, 50)
skipper = SchedulerRecvSkipper.maybe_create()
self.assertTrue(skipper.handle(_batch(ForwardMode.EXTEND)))
def test_dp_uses_synced_mode_not_local(self):
# Local EXTEND (weight 1000) must be ignored in favor of the synced
# DECODE (weight 1); a recv here would mean the local mode leaked in.
skipper = SchedulerRecvSkipper.maybe_create(
_server_args(50, enable_dp_attention=True)
)
_publish(self, 50, enable_dp_attention=True)
skipper = SchedulerRecvSkipper.maybe_create()
self.assertFalse(skipper.handle(_batch(ForwardMode.EXTEND, ForwardMode.DECODE)))
def test_dp_synced_extend_triggers_recv(self):
skipper = SchedulerRecvSkipper.maybe_create(
_server_args(50, enable_dp_attention=True)
)
_publish(self, 50, enable_dp_attention=True)
skipper = SchedulerRecvSkipper.maybe_create()
self.assertTrue(skipper.handle(_batch(ForwardMode.IDLE, ForwardMode.EXTEND)))
def test_derive_forward_mode(self):
@@ -28,7 +28,7 @@ from sglang.srt.model_executor.model_runner_components.startup_weight_load impor
)
from sglang.srt.model_loader.loader import DefaultModelLoader
from sglang.srt.model_loader.weight_utils import initialize_capture_safe_weights
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 ServerArgs
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
@@ -203,10 +203,14 @@ class TestStartupWeightLoadSelector(CustomTestCase):
def test_options_accept_current_server_args_schema(self):
"""Removed server options must not break overlap startup initialization."""
server_args = ServerArgs(
model_path="dummy", cuda_graph_config=CudaGraphConfig()
)
# The parallel sizes come from the bags, so the config has to be published.
publish(server_args, role="test")
self.addCleanup(reset_context)
options = StartupWeightLoadOptions.from_server_args(
server_args=ServerArgs(
model_path="dummy", cuda_graph_config=CudaGraphConfig()
),
server_args=server_args,
is_draft_worker=False,
)
@@ -83,16 +83,22 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
get_model_architecture.assert_not_called()
def test_vmm_transport_initializes_pool(self):
from sglang.srt.runtime_context import get_context
from sglang.srt.utils import cuda_vmm_transport_utils as vmm
server_args = SimpleNamespace(
mm_feature_transport="cuda_vmm",
tokenizer_worker_num=2,
base_gpu_id=3,
enable_dp_attention=False,
tp_size=4,
nnodes=1,
)
# The consumer count comes from the published topology.
override = get_context().override_server_args(
enable_dp_attention=False, tp_size=4
)
override.install()
self.addCleanup(override.restore)
pool = object()
with (
patch.object(vmm, "get_mm_feature_pool_size_per_worker", return_value=123),
@@ -76,6 +76,21 @@ _CONFIGURED_SIZE_CALL_SITES = {
"the one where initialize_model_parallel aliases _MOE_DP to _ATTN_CP, so "
"the live sizes are equal there and a live comparison is always false"
),
("srt/managers/scheduler.py", "configured_tp_size"): (
"configure_scheduler_process runs before the scheduler's own process "
"groups exist -- configuring the process is what it is for -- so there "
"is nothing live to ask yet"
),
("srt/managers/scheduler.py", "configured_moe_dp_size"): (
"same pre-distributed-init arithmetic in configure_scheduler_process"
),
("srt/managers/scheduler.py", "configured_attn_cp_size"): (
"same pre-distributed-init arithmetic in configure_scheduler_process"
),
("srt/utils/cuda_vmm_transport_utils.py", "configured_tp_size"): (
"the consumer count is configured fan-out arithmetic (tp_size // "
"dp_size), which is what the record answered before"
),
("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"
@@ -104,10 +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/entrypoints/engine.py", "_launch_subprocesses"): (
"its subprocess targets and its tokenizer-manager factory arrive as "
"parameters, so the walk resolves none of them"
),
("srt/managers/detokenizer_manager.py", "run_detokenizer_process"): (
"DetokenizerManager reads the handed instance at this revision"
),
@@ -155,7 +155,6 @@ _EXPOSED = {
("disaggregation/encode_receiver.py", "tokenizer_path"),
("disaggregation/encode_server.py", "allowed_media_domains"),
("disaggregation/encode_server.py", "device"),
("disaggregation/encode_server.py", "dp_size"),
("disaggregation/encode_server.py", "encoder_transfer_backend"),
("disaggregation/encode_server.py", "load_format"),
("disaggregation/encode_server.py", "mm_process_config"),
@@ -203,12 +202,9 @@ _EXPOSED = {
("elastic_ep/expert_backup_manager.py", "mooncake_ib_device"),
("entrypoints/engine.py", "attn_cp_size"),
("entrypoints/engine.py", "detokenizer_worker_num"),
("entrypoints/engine.py", "dp_size"),
("entrypoints/engine.py", "dtype"),
("entrypoints/engine.py", "enable_dp_attention"),
("entrypoints/engine.py", "enable_symm_mem"),
("entrypoints/engine.py", "ep_join_mode"),
("entrypoints/engine.py", "ep_size"),
("entrypoints/engine.py", "load_format"),
("entrypoints/engine.py", "model_path"),
("entrypoints/engine.py", "moe_dp_size"),
@@ -221,7 +217,6 @@ _EXPOSED = {
),
("entrypoints/engine.py", "tool_call_parser"),
("entrypoints/http_server.py", "disaggregation_mode"),
("entrypoints/http_server.py", "dp_size"),
("entrypoints/http_server.py", "ep_join_mode"),
("entrypoints/http_server.py", "grpc_port"),
("entrypoints/http_server.py", "model_path"),
@@ -247,10 +242,6 @@ _EXPOSED = {
("layers/cp/base.py", "enable_prefill_cp"),
("layers/cp/bcg.py", "cp_strategy"),
("layers/cp/bcg.py", "enable_prefill_cp"),
("layers/dp_attention.py", "attn_cp_size"),
("layers/dp_attention.py", "device"),
("layers/dp_attention.py", "dp_size"),
("layers/dp_attention.py", "enable_dp_attention"),
("layers/flashinfer_comm_fusion.py", "flashinfer_allreduce_fusion_backend"),
("layers/moe/kt_ep_wrapper.py", "chunked_prefill_size"),
("layers/moe/utils.py", "deepep_mode"),
@@ -259,19 +250,11 @@ _EXPOSED = {
("layers/moe/utils.py", "quantization"),
("layers/moe/utils.py", "speculative_moe_runner_backend"),
("layers/quantization/unquant.py", "enable_deterministic_inference"),
("lora/lora_manager.py", "enable_dp_attention"),
("lora/lora_manager.py", "enable_lora_overlap_loading"),
("lora/marlin_lora_temp/policy.py", "enable_lora"),
("lora/marlin_lora_temp/policy.py", "lora_paths"),
("managers/data_parallel_controller.py", "attn_cp_size"),
("managers/data_parallel_controller.py", "disaggregation_mode"),
("managers/data_parallel_controller.py", "dp_size"),
("managers/data_parallel_controller.py", "enable_dp_attention"),
(
"managers/data_parallel_controller.py",
"enable_dp_attention_local_control_broadcast",
),
("managers/data_parallel_controller.py", "ep_size"),
("managers/data_parallel_controller.py", "load_balance_method"),
("managers/data_parallel_controller.py", "moe_dp_size"),
("managers/data_parallel_controller.py", "pp_size"),
@@ -282,23 +265,16 @@ _EXPOSED = {
("managers/disagg_service.py", "disaggregation_bootstrap_port"),
("managers/disagg_service.py", "disaggregation_mode"),
("managers/disagg_service.py", "disaggregation_transfer_backend"),
("managers/load_snapshot.py", "dp_size"),
("managers/load_snapshot.py", "enable_dp_attention"),
("managers/load_snapshot.py", "load_balance_method"),
("managers/overlap_utils.py", "speculative_algorithm"),
("managers/prefill_delayer.py", "disable_overlap_schedule"),
("managers/prefill_delayer.py", "enable_dp_attention"),
("managers/rust_server.py", "mm_process_config"),
("managers/schedule_batch.py", "disaggregation_mode"),
("managers/scheduler.py", "attn_cp_size"),
("managers/scheduler.py", "disable_overlap_schedule"),
("managers/scheduler.py", "disaggregation_mode"),
("managers/scheduler.py", "dp_size"),
("managers/scheduler.py", "enable_dp_attention"),
("managers/scheduler.py", "enable_hierarchical_cache"),
("managers/scheduler.py", "enable_lora"),
("managers/scheduler.py", "enable_lora_overlap_loading"),
("managers/scheduler.py", "ep_size"),
("managers/scheduler.py", "moe_dp_size"),
("managers/scheduler.py", "pp_size"),
("managers/scheduler.py", "soft_watchdog_timeout"),
@@ -307,13 +283,9 @@ _EXPOSED = {
"managers/scheduler_components/new_token_ratio_tracker.py",
"schedule_conservativeness",
),
("managers/scheduler_components/recv_skipper.py", "enable_dp_attention"),
("managers/tokenizer_control_mixin.py", "dp_size"),
("managers/tokenizer_manager.py", "disable_radix_cache"),
("managers/tokenizer_manager.py", "disaggregation_mode"),
("managers/tokenizer_manager.py", "disaggregation_transfer_backend"),
("managers/tokenizer_manager.py", "dp_size"),
("managers/tokenizer_manager.py", "enable_dp_attention"),
("managers/tokenizer_manager.py", "enable_lora"),
("managers/tokenizer_manager.py", "enable_tokenizer_batch_encode"),
("managers/tokenizer_manager.py", "encoder_transfer_backend"),
@@ -339,9 +311,6 @@ _EXPOSED = {
("mem_cache/hybrid_cache/hybrid_pool_assembler.py", "hicache_io_backend"),
("mem_cache/hybrid_cache/hybrid_pool_assembler.py", "hicache_mem_layout"),
("mem_cache/hybrid_cache/hybrid_pool_assembler.py", "served_model_name"),
("mem_cache/kv_cache_builder.py", "disable_radix_cache"),
("mem_cache/kv_cache_builder.py", "disaggregation_mode"),
("mem_cache/kv_cache_builder.py", "enable_dp_attention"),
("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"),
@@ -373,9 +342,7 @@ _EXPOSED = {
"custom_weight_loader",
),
("model_executor/model_runner_components/startup_weight_load.py", "device"),
("model_executor/model_runner_components/startup_weight_load.py", "dp_size"),
("model_executor/model_runner_components/startup_weight_load.py", "enable_lora"),
("model_executor/model_runner_components/startup_weight_load.py", "ep_size"),
("model_executor/model_runner_components/startup_weight_load.py", "lora_paths"),
("model_executor/model_runner_components/startup_weight_load.py", "pp_size"),
(
@@ -400,11 +367,7 @@ _EXPOSED = {
("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", "dp_size"),
("ray/data_parallel_controller.py", "enable_dp_attention"),
("ray/data_parallel_controller.py", "pp_size"),
("ray/engine.py", "dp_size"),
("ray/engine.py", "enable_dp_attention"),
("ray/engine.py", "pp_size"),
("speculative/adaptive_spec_params.py", "speculative_algorithm"),
("speculative/adaptive_spec_params.py", "speculative_eagle_topk"),
@@ -416,29 +379,24 @@ _EXPOSED = {
"speculative/dspark_components/dspark_config.py",
"speculative_draft_model_revision",
),
("speculative/dspark_components/dspark_worker_v2.py", "disable_cuda_graph"),
("speculative/dspark_components/dspark_worker_v2.py", "disaggregation_mode"),
("speculative/dspark_components/dspark_worker_v2.py", "enable_dp_attention"),
(
"speculative/dspark_components/dspark_worker_v2.py",
"speculative_num_draft_tokens",
),
("speculative/eagle_worker_v2.py", "device"),
("speculative/eagle_worker_v2.py", "enable_dp_attention"),
("speculative/eagle_worker_v2.py", "speculative_adaptive"),
("speculative/eagle_worker_v2.py", "speculative_algorithm"),
("speculative/eagle_worker_v2.py", "speculative_eagle_topk"),
("speculative/eagle_worker_v2.py", "speculative_num_draft_tokens"),
("speculative/eagle_worker_v2.py", "speculative_num_steps"),
("speculative/frozen_kv_mtp_worker_v2.py", "device"),
("speculative/frozen_kv_mtp_worker_v2.py", "enable_dp_attention"),
("speculative/frozen_kv_mtp_worker_v2.py", "speculative_adaptive"),
("speculative/frozen_kv_mtp_worker_v2.py", "speculative_algorithm"),
("speculative/frozen_kv_mtp_worker_v2.py", "speculative_eagle_topk"),
("speculative/frozen_kv_mtp_worker_v2.py", "speculative_num_draft_tokens"),
("speculative/frozen_kv_mtp_worker_v2.py", "speculative_num_steps"),
("speculative/multi_layer_eagle_worker_v2.py", "device"),
("speculative/multi_layer_eagle_worker_v2.py", "enable_dp_attention"),
("speculative/multi_layer_eagle_worker_v2.py", "speculative_algorithm"),
("speculative/multi_layer_eagle_worker_v2.py", "speculative_eagle_topk"),
("speculative/multi_layer_eagle_worker_v2.py", "speculative_num_draft_tokens"),
@@ -451,7 +409,6 @@ _EXPOSED = {
("speculative/spec_info.py", "enable_multi_layer_eagle"),
("speculative/spec_registry.py", "disable_overlap_schedule"),
("speculative/standalone_worker_v2.py", "device"),
("speculative/standalone_worker_v2.py", "enable_dp_attention"),
("speculative/standalone_worker_v2.py", "speculative_algorithm"),
("speculative/standalone_worker_v2.py", "speculative_eagle_topk"),
("speculative/standalone_worker_v2.py", "speculative_num_draft_tokens"),
@@ -460,11 +417,8 @@ _EXPOSED = {
("utils/common.py", "speculative_eagle_topk"),
("utils/common.py", "speculative_num_draft_tokens"),
("utils/common.py", "speculative_num_steps"),
("utils/cuda_vmm_transport_utils.py", "dp_size"),
("utils/cuda_vmm_transport_utils.py", "enable_dp_attention"),
("utils/cuda_vmm_transport_utils.py", "mm_feature_transport"),
("utils/hf_transformers/processor.py", "image_processor_backend"),
("utils/offloader.py", "dp_size"),
}
# Pairs whose resolution write only happens on a CUDA host (capability or
@@ -487,7 +441,6 @@ _OVERRIDDEN_AND_READ = {
"disaggregation/decode_kvcache_offload_manager.py",
"hicache_storage_backend_extra_config",
),
("disaggregation/encode_server.py", "dp_size"),
("disaggregation/encode_server.py", "load_format"),
("disaggregation/encode_server.py", "model_path"),
(
@@ -496,24 +449,13 @@ _OVERRIDDEN_AND_READ = {
),
("dllm/config.py", "model_path"),
("elastic_ep/expert_backup_manager.py", "load_format"),
("entrypoints/engine.py", "dp_size"),
("entrypoints/engine.py", "dtype"),
("entrypoints/engine.py", "ep_size"),
("entrypoints/engine.py", "load_format"),
("entrypoints/engine.py", "model_path"),
("entrypoints/http_server.py", "dp_size"),
("entrypoints/http_server.py", "model_path"),
("kv_canary/api.py", "speculative_num_steps"),
("kv_canary/capacities.py", "speculative_num_draft_tokens"),
("layers/dp_attention.py", "dp_size"),
("managers/data_parallel_controller.py", "dp_size"),
("managers/data_parallel_controller.py", "ep_size"),
("managers/load_snapshot.py", "dp_size"),
("managers/scheduler.py", "dp_size"),
("managers/scheduler.py", "ep_size"),
("managers/scheduler.py", "hicache_storage_backend"),
("managers/tokenizer_control_mixin.py", "dp_size"),
("managers/tokenizer_manager.py", "dp_size"),
("managers/tokenizer_manager.py", "model_path"),
("managers/tokenizer_manager.py", "speculative_num_draft_tokens"),
("managers/tp_worker.py", "model_path"),
@@ -532,13 +474,8 @@ _OVERRIDDEN_AND_READ = {
("mem_cache/unified_radix_cache.py", "hicache_storage_prefetch_policy"),
("mem_cache/unified_radix_cache.py", "hicache_write_policy"),
("model_executor/model_runner_components/load_model_utils.py", "load_format"),
("model_executor/model_runner_components/startup_weight_load.py", "dp_size"),
("model_executor/model_runner_components/startup_weight_load.py", "ep_size"),
("parser/template_detection.py", "model_path"),
("ray/data_parallel_controller.py", "dp_size"),
("ray/engine.py", "dp_size"),
("speculative/dflash_worker_v2.py", "speculative_num_draft_tokens"),
("speculative/dspark_components/dspark_worker_v2.py", "disable_cuda_graph"),
(
"speculative/dspark_components/dspark_worker_v2.py",
"speculative_num_draft_tokens",
@@ -555,8 +492,6 @@ _OVERRIDDEN_AND_READ = {
("speculative/standalone_worker_v2.py", "speculative_num_steps"),
("utils/common.py", "speculative_num_draft_tokens"),
("utils/common.py", "speculative_num_steps"),
("utils/cuda_vmm_transport_utils.py", "dp_size"),
("utils/offloader.py", "dp_size"),
}