config: the runner and scheduler read resolved config from the bags (#34095)
This commit is contained in:
@@ -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)]
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user