diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 02db16862..d50c635db 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -69,7 +69,11 @@ from sglang.srt.mem_cache.common import ( ) from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.observability.req_time_stats import set_schedule_time_batch -from sglang.srt.runtime_context import get_disagg +from sglang.srt.runtime_context import ( + get_disagg, + get_parallel, + get_schedule, +) from sglang.srt.utils import is_npu from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method @@ -158,23 +162,25 @@ class PrefillBootstrapQueue: "SGLANG_DISAGG_STAGING_BUFFER is designed for non-MLA models " "(e.g. GQA, MHA). MLA models should not set this flag." ) - server_args = self.scheduler.server_args page_size = self.scheduler.token_to_kv_pool_allocator.page_size - cps = server_args.chunked_prefill_size or 8192 + # Same source as send_kv_chunk's staging grid below, so validation + # and the grid cannot disagree after a post-publish override. + chunked_prefill_size = get_schedule().chunked_prefill_size + cps = chunked_prefill_size or 8192 # Staging slices each send into a fixed page-aligned grid, so an # unbounded (-1) or non-page-aligned chunk size has no valid grid. if cps <= 0 or cps % page_size != 0: raise RuntimeError( f"SGLANG_DISAGG_STAGING_BUFFER requires a positive " f"chunked_prefill_size that is a multiple of page_size " - f"({page_size}); got {server_args.chunked_prefill_size}." + f"({page_size}); got {chunked_prefill_size}." ) if self.pp_size > 1: # Staging writer accounting has no pp dimension. raise RuntimeError( "SGLANG_DISAGG_STAGING_BUFFER does not support pp_size > 1." ) - if server_args.enable_prefill_context_parallel: + if get_parallel().enable_prefill_context_parallel: # CP rewrites index_slice per rank, breaking the chunk grid. raise RuntimeError( "SGLANG_DISAGG_STAGING_BUFFER does not support " @@ -1145,7 +1151,7 @@ class SchedulerDisaggregationPrefillMixin: # prefetched grid, so non-last sends must end on a grid # boundary; the remainder rides with the next send. grid_tokens = staging_grid_tokens( - self.server_args.chunked_prefill_size, page_size + get_schedule().chunked_prefill_size, page_size ) base = req.disagg_decode_prefix_len end_idx = base + ((end_idx - base) // grid_tokens) * grid_tokens @@ -1271,7 +1277,7 @@ class SchedulerDisaggregationPrefillMixin: start_idx, end_idx, req.disagg_decode_prefix_len, - staging_grid_tokens(self.server_args.chunked_prefill_size, page_size), + staging_grid_tokens(get_schedule().chunked_prefill_size, page_size), ) else: segments = [(start_idx, end_idx)] diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 7954fd44e..3119c7728 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -420,7 +420,7 @@ class Scheduler( self.enable_overlap = not server_args.disable_overlap_schedule and not use_mlx() self.enable_overlap_mlx = not server_args.disable_overlap_schedule and use_mlx() self.enable_pdmux = server_args.enable_pdmux - self.skip_tokenizer_init = server_args.skip_tokenizer_init + self.skip_tokenizer_init = get_serving().skip_tokenizer_init self.stream_interval = server_args.stream_interval self.spec_algorithm = SpeculativeAlgorithm.from_string( server_args.speculative_algorithm @@ -731,7 +731,10 @@ class Scheduler( self.ipc_channels = SchedulerIpcChannels.create( port_args=port_args, is_rank_zero=is_rank_zero, - skip_tokenizer_init=self.server_args.skip_tokenizer_init, + # The snapshot taken at construction, not a second bag read: this + # scheduler gates its tokenizer init on the same value, and the two + # must not be able to disagree. + skip_tokenizer_init=self.skip_tokenizer_init, metrics_enabled=get_observability().enable_metrics and ( self.ps.attn_tp_rank == 0 @@ -793,7 +796,7 @@ class Scheduler( server_args = self.server_args self.is_generation = self.model_config.is_generation - if server_args.skip_tokenizer_init: + if self.skip_tokenizer_init: self.tokenizer = self.processor = None else: if self.model_config.is_multimodal: @@ -1484,7 +1487,7 @@ class Scheduler( "triton": ("SGLANG_TRITON_PREFILL_TRUNCATION_ALIGN_SIZE", 4096), } env_var, default_size = backend_sizes.get( - self.server_args.attention_backend, (None, None) + get_exec().kernel.attention_backend, (None, None) ) self.truncation_align_size = ( get_int_env_var(env_var, default_size) if env_var else None diff --git a/python/sglang/srt/managers/scheduler_components/request_receiver.py b/python/sglang/srt/managers/scheduler_components/request_receiver.py index e35bff05e..bcdfff5a3 100644 --- a/python/sglang/srt/managers/scheduler_components/request_receiver.py +++ b/python/sglang/srt/managers/scheduler_components/request_receiver.py @@ -27,7 +27,7 @@ from sglang.srt.managers.mm_utils import ( has_shm_features, unwrap_shm_features, ) -from sglang.srt.runtime_context import get_disagg, get_parallel +from sglang.srt.runtime_context import get_disagg, get_parallel, is_ep_scale_joiner from sglang.srt.utils import ( broadcast_pyobj, point_to_point_pyobj, @@ -181,7 +181,7 @@ class SchedulerRequestReceiver: # all-ranks gloo sync. _local_ctrl = ( get_parallel().enable_dp_attention_local_control_broadcast - or self.server_args.is_ep_scale_joiner + or is_ep_scale_joiner() ) if _local_ctrl: if self.ps.attn_tp_size != 1: diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 9788acf62..8aa3f6551 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -178,6 +178,8 @@ from sglang.srt.runtime_context import ( get_parallel, get_schedule, get_spec, + is_ep_joiner, + is_ep_scale_joiner, set_global_dwdp_manager, ) from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo @@ -458,10 +460,7 @@ class ModelRunner: self.graph_time_usage: dict[str, float] = {} def _initialize_elastic_ep_joiner(self) -> None: - if not ( - get_exec().moe.elastic_ep_backend is not None - and self.server_args.is_ep_scale_joiner - ): + if not (get_exec().moe.elastic_ep_backend is not None and is_ep_scale_joiner()): return join_effective_ep_size = get_parallel().ep_join_rank_offset + self.ps.tp_size @@ -677,9 +676,7 @@ class ModelRunner: if self.is_draft_worker: return expert_rank = self.ps.moe_ep_rank + ( - get_parallel().ep_join_rank_offset - if self.server_args.is_ep_scale_joiner - else 0 + get_parallel().ep_join_rank_offset if is_ep_scale_joiner() else 0 ) set_global_expert_location_metadata( compute_initial_expert_location_metadata( @@ -1120,7 +1117,15 @@ class ModelRunner: pyt_hooks = PytHooks() pyt_hooks.register_hooks(self.model, module_prefix="model") - load_kv_cache_scales(model=self.model, server_args=self.server_args) + # Same leaf `configure_kv_cache_dtype` reads: the bag, not the startup + # record, so the FP8 gate and the pool cannot disagree after an + # override. (The runner's own stamp is not set yet -- load_model runs + # before configure_kv_cache_dtype.) + load_kv_cache_scales( + model=self.model, + server_args=self.server_args, + kv_cache_dtype=get_model().kv_cache_dtype, + ) self.sliding_window_size = resolve_sliding_window_size( self.model, self.model_config @@ -1176,7 +1181,7 @@ class ModelRunner: dist_barrier_after_load( elastic_ep_backend=get_exec().moe.elastic_ep_backend, tp_rank=self.ps.tp_rank, - is_ep_joiner=self.server_args.is_ep_joiner, + is_ep_joiner=is_ep_joiner(), ) def maybe_precompile_model_kernels_after_loading(self) -> None: @@ -1268,7 +1273,7 @@ class ModelRunner: spec_algorithm = getattr(self, "spec_algorithm", None) resolved_kv_cache_dtype, self.kv_cache_dtype = ( kv_cache_dtype.configure_kv_cache_dtype( - server_args_kv_cache_dtype=self.server_args.kv_cache_dtype, + server_args_kv_cache_dtype=get_model().kv_cache_dtype, model=getattr(self, "model", None), model_dtype=getattr(self, "dtype", torch.bfloat16), is_draft_worker=getattr(self, "is_draft_worker", False), @@ -1287,7 +1292,7 @@ class ModelRunner: self.kv_cache_dtype_str = ( resolved_kv_cache_dtype if resolved_kv_cache_dtype is not None - else self.server_args.kv_cache_dtype + else get_model().kv_cache_dtype ) def _get_attention_backend(self, init_new_workspace: bool = False): @@ -1835,7 +1840,7 @@ class ModelRunner: self._rearm_eplb_after_elastic_scale() def _report_elastic_scale_failure(self, error: str, effective_size: int) -> None: - if self.ps.tp_rank != 0 or self.server_args.is_ep_scale_joiner: + if self.ps.tp_rank != 0 or is_ep_scale_joiner(): return from sglang.srt.managers.io_struct import ElasticScaleUpdateReq @@ -1915,12 +1920,12 @@ class ModelRunner: ElasticEPStateManager.mark_syncing_new_world() self._elastic_scale_ready_barrier( target_size=target_size, - log_tag="JOINER" if self.server_args.is_ep_scale_joiner else "PRIMARY", + log_tag="JOINER" if is_ep_scale_joiner() else "PRIMARY", ) ElasticEPStateManager.commit_scale() self._rearm_eplb_after_elastic_scale() - if self.ps.tp_rank == 0 and not self.server_args.is_ep_scale_joiner: + if self.ps.tp_rank == 0 and not is_ep_scale_joiner(): from sglang.srt.managers.io_struct import ElasticScaleUpdateReq self._pending_elastic_scale_update = ElasticScaleUpdateReq( @@ -1953,7 +1958,7 @@ class ModelRunner: ) ElasticEPStateManager.fail_recovery(error) self._report_elastic_scale_failure(error, effective_size) - if self.ps.tp_rank == 0 and not self.server_args.is_ep_scale_joiner: + if self.ps.tp_rank == 0 and not is_ep_scale_joiner(): logger.error("[Elastic EP] %s", error) return @@ -1979,7 +1984,7 @@ class ModelRunner: ElasticEPStateManager.fail_scale(error) self._reset_eplb_after_elastic_scale_failure() self._report_elastic_scale_failure(error, effective_size) - if self.ps.tp_rank == 0 and not self.server_args.is_ep_scale_joiner: + if self.ps.tp_rank == 0 and not is_ep_scale_joiner(): logger.error("[Elastic EP] %s", error) return @@ -1995,7 +2000,7 @@ class ModelRunner: ElasticEPStateManager.fail_scale(error) self._reset_eplb_after_elastic_scale_failure() self._report_elastic_scale_failure(error, effective_size) - if self.ps.tp_rank == 0 and not self.server_args.is_ep_scale_joiner: + if self.ps.tp_rank == 0 and not is_ep_scale_joiner(): logger.error("[Elastic EP] %s", error) return if not ElasticEPStateManager.begin_scale(): diff --git a/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py b/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py index 14f9755a5..6f2a23ea7 100644 --- a/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py +++ b/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py @@ -103,8 +103,13 @@ def maybe_trigger_remote_instance_nccl_send_group( t.start() -def load_kv_cache_scales(*, model, server_args: ServerArgs) -> None: - if server_args.kv_cache_dtype == "fp8_e4m3": +def load_kv_cache_scales( + *, model, server_args: ServerArgs, kv_cache_dtype: str +) -> None: + """``kv_cache_dtype`` is the caller's resolved value. Required rather than + defaulted: a fallback to ``server_args`` would be a hidden global read for + any future caller that forgets to pass one.""" + if kv_cache_dtype == "fp8_e4m3": if server_args.quantization_param_path is not None: if callable(getattr(model, "load_kv_cache_scales", None)): model.load_kv_cache_scales(server_args.quantization_param_path) diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index 386f7c5ac..f39a11f5c 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -1553,3 +1553,18 @@ def configured_moe_dp_size() -> int: def configured_attn_cp_size() -> int: return _configured_parallel("attn_cp_size") + + +def is_ep_joiner() -> bool: + """True in a process launched as an elastic-EP joiner (scale or recover). + + A predicate over the published ``exec.moe.ep_join_mode`` leaf, so it follows + a post-publish override; the same-named ``ServerArgs`` property is the + pre-publish equivalent. + """ + return get_exec().moe.ep_join_mode in ("scale", "recover") + + +def is_ep_scale_joiner() -> bool: + """True in a process launched as an elastic-EP scale-up joiner.""" + return get_exec().moe.ep_join_mode == "scale" diff --git a/python/sglang/srt/speculative/dspark_components/dspark_planner.py b/python/sglang/srt/speculative/dspark_components/dspark_planner.py index c67da766d..2d38ab0b8 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_planner.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_planner.py @@ -177,7 +177,7 @@ class DSparkVerifyPlanner: and not get_schedule().disable_overlap_schedule and not get_spec().speculative_skip_dp_mlp_sync and get_disagg().disaggregation_mode == "null" - and self.server_args.pp_size == 1 + and get_parallel().pp_size == 1 and not envs.SGLANG_SCHEDULER_SKIP_ALL_GATHER.get() ) if tp_rank == 0: @@ -609,7 +609,7 @@ class DSparkVerifyPlanner: ) broadcast_group, group_size = verify_lens_broadcast_group( - tp_size=self.server_args.tp_size + tp_size=get_parallel().tp_size ) if group_size > 1: broadcast_group.broadcast(verify_lens, src=0) diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py b/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py index f084f0596..e31798633 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py @@ -44,6 +44,7 @@ from sglang.srt.model_executor.forward_batch_info import ( ) from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig +from sglang.srt.runtime_context import attention_backends, get_spec from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase from sglang.srt.speculative.eagle_utils import ( @@ -228,11 +229,15 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker): pass def _resolve_draft_backend_type(self) -> str: - return ( - self.server_args.speculative_draft_attention_backend - or self.server_args.decode_attention_backend - or self.server_args.attention_backend - ) + # The same chain as before, off the bags: the speculative override if + # the operator set one, else the configured decode backend (which falls + # back to the base one). Deliberately NOT the runner's stamp: this + # worker does not hand its runner a draft backend, so the stamp is the + # ordinary pair and reading it would drop the speculative setting -- + # and forcing the runner onto one backend would collapse a hybrid + # prefill/decode config for the topk==1 path, which uses the runner's + # own backend. + return get_spec().speculative_draft_attention_backend or attention_backends()[1] def _init_draft_attn_backend(self): if self.topk == 1: diff --git a/test/registered/unit/layers/quantization/test_fp4_kv_cache_quant_method.py b/test/registered/unit/layers/quantization/test_fp4_kv_cache_quant_method.py index 118de4a62..0ebf64efa 100644 --- a/test/registered/unit/layers/quantization/test_fp4_kv_cache_quant_method.py +++ b/test/registered/unit/layers/quantization/test_fp4_kv_cache_quant_method.py @@ -72,10 +72,16 @@ class TestKVCacheQuantRegistry(CustomTestCase): from types import SimpleNamespace from sglang.srt.model_executor.model_runner import ModelRunner + from sglang.srt.runtime_context import get_context runner = object.__new__(ModelRunner) - runner.server_args = SimpleNamespace(kv_cache_dtype="fp4_e2m1") + runner.server_args = SimpleNamespace() runner.draft_attention_backend = None + # The runner reads the requested dtype off the model bag, so the double + # publishes it rather than carrying it on a stand-in config. + override = get_context().override_server_args(kv_cache_dtype="fp4_e2m1") + override.install() + self.addCleanup(override.restore) with self.assertRaisesRegex(ValueError, "fp4_mx_block16"): runner.configure_kv_cache_dtype() diff --git a/test/registered/unit/managers/test_scheduler_init_req_max_new_tokens.py b/test/registered/unit/managers/test_scheduler_init_req_max_new_tokens.py index 7e4921c27..93f41fb5c 100644 --- a/test/registered/unit/managers/test_scheduler_init_req_max_new_tokens.py +++ b/test/registered/unit/managers/test_scheduler_init_req_max_new_tokens.py @@ -4,6 +4,7 @@ from types import SimpleNamespace from sglang.srt.environ import envs from sglang.srt.managers.scheduler import Scheduler +from sglang.srt.runtime_context import get_parallel from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=1, suite="base-a-test-cpu") @@ -36,6 +37,14 @@ class TestSchedulerInitReqMaxNewTokens(unittest.TestCase): def tearDownClass(cls): cls._scheduler_logger.setLevel(cls._old_level) + def setUp(self): + # The scheduler scales the budget by the live DCP size + # (`get_parallel().attn_dcp_size`), so the double states a topology + # rather than publishing a config it does not otherwise need. + cm = get_parallel().override(attn_dcp_size=1) + cm.__enter__() + self.addCleanup(cm.__exit__, None, None, None) + def _new_scheduler( self, max_req_len: int = 128,