config: the runner and scheduler read resolved config from the bags (#34095)

This commit is contained in:
Cheng Wan
2026-08-09 14:45:11 -07:00
committed by GitHub
parent a2199c1dee
commit e216c2bc59
10 changed files with 94 additions and 40 deletions
+13 -7
View File
@@ -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)]
+7 -4
View File
@@ -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
@@ -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:
@@ -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():
@@ -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)
+15
View File
@@ -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"
@@ -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)
@@ -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: