runtime_context: record the publishing process role (#33012)
This commit is contained in:
@@ -63,10 +63,10 @@ from sglang.srt.observability.trace import (
|
|||||||
process_tracing_init,
|
process_tracing_init,
|
||||||
trace_set_thread_info,
|
trace_set_thread_info,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import publish
|
||||||
from sglang.srt.server_args import (
|
from sglang.srt.server_args import (
|
||||||
PortArgs,
|
PortArgs,
|
||||||
ServerArgs,
|
ServerArgs,
|
||||||
set_global_server_args_for_scheduler,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
CLIENT_MEDIA_EXCEPTIONS,
|
CLIENT_MEDIA_EXCEPTIONS,
|
||||||
@@ -265,7 +265,7 @@ class MMEncoder:
|
|||||||
):
|
):
|
||||||
logger.info(f"init MMEncoder {rank}/{server_args.tp_size}")
|
logger.info(f"init MMEncoder {rank}/{server_args.tp_size}")
|
||||||
self.server_args = server_args
|
self.server_args = server_args
|
||||||
set_global_server_args_for_scheduler(server_args)
|
publish(server_args, role="encoder")
|
||||||
self.rank = rank
|
self.rank = rank
|
||||||
# DP rank for metric labels; overridden by run_dp_worker in DP mode.
|
# DP rank for metric labels; overridden by run_dp_worker in DP mode.
|
||||||
# 0 in the single-instance (non-DP) path.
|
# 0 in the single-instance (non-DP) path.
|
||||||
|
|||||||
@@ -17,10 +17,10 @@ from sglang.srt.managers.io_struct import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader
|
from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader
|
||||||
from sglang.srt.model_loader.utils import set_default_torch_dtype
|
from sglang.srt.model_loader.utils import set_default_torch_dtype
|
||||||
|
from sglang.srt.runtime_context import publish
|
||||||
from sglang.srt.server_args import (
|
from sglang.srt.server_args import (
|
||||||
PortArgs,
|
PortArgs,
|
||||||
ServerArgs,
|
ServerArgs,
|
||||||
set_global_server_args_for_scheduler,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.network import get_local_ip_auto
|
from sglang.srt.utils.network import get_local_ip_auto
|
||||||
|
|
||||||
@@ -159,7 +159,7 @@ def run_expert_backup_manager_process(
|
|||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
port_args: PortArgs,
|
port_args: PortArgs,
|
||||||
):
|
):
|
||||||
set_global_server_args_for_scheduler(server_args)
|
publish(server_args, role="expert_backup")
|
||||||
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
|
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
|
||||||
init_mooncake_transfer_engine,
|
init_mooncake_transfer_engine,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -48,6 +48,7 @@ from sglang.srt.managers.scheduler import run_scheduler_process
|
|||||||
from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
|
from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
|
||||||
from sglang.srt.observability.req_time_stats import DPControllerReqTimeStats
|
from sglang.srt.observability.req_time_stats import DPControllerReqTimeStats
|
||||||
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
|
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
|
||||||
|
from sglang.srt.runtime_context import publish
|
||||||
from sglang.srt.server_args import (
|
from sglang.srt.server_args import (
|
||||||
DP_ATTENTION_HANDSHAKE_PORT_DELTA,
|
DP_ATTENTION_HANDSHAKE_PORT_DELTA,
|
||||||
PortArgs,
|
PortArgs,
|
||||||
@@ -815,6 +816,8 @@ def run_data_parallel_controller_process(
|
|||||||
kill_itself_when_parent_died()
|
kill_itself_when_parent_died()
|
||||||
parent_process = psutil.Process().parent()
|
parent_process = psutil.Process().parent()
|
||||||
|
|
||||||
|
# This process reads the config namespaces before spawning schedulers.
|
||||||
|
publish(server_args, role="dp_controller")
|
||||||
configure_logger(server_args)
|
configure_logger(server_args)
|
||||||
if server_args.enable_trace:
|
if server_args.enable_trace:
|
||||||
process_tracing_init(
|
process_tracing_init(
|
||||||
|
|||||||
@@ -250,7 +250,7 @@ from sglang.srt.observability.trace import process_tracing_init, trace_set_threa
|
|||||||
from sglang.srt.parser.reasoning_parser import ReasoningParser
|
from sglang.srt.parser.reasoning_parser import ReasoningParser
|
||||||
from sglang.srt.platforms import current_platform
|
from sglang.srt.platforms import current_platform
|
||||||
from sglang.srt.plugins import load_plugins
|
from sglang.srt.plugins import load_plugins
|
||||||
from sglang.srt.runtime_context import get_context, get_parallel
|
from sglang.srt.runtime_context import get_context, get_parallel, publish
|
||||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||||
from sglang.srt.sampling.sampling_params import TOP_K_ALL
|
from sglang.srt.sampling.sampling_params import TOP_K_ALL
|
||||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||||
@@ -4765,6 +4765,9 @@ def run_scheduler_process(
|
|||||||
display_dp_rank=display_dp_rank,
|
display_dp_rank=display_dp_rank,
|
||||||
display_moe_ep_rank=display_moe_ep_rank,
|
display_moe_ep_rank=display_moe_ep_rank,
|
||||||
)
|
)
|
||||||
|
# Scheduler.__init__ reads the config namespaces before the model
|
||||||
|
# worker's own publish.
|
||||||
|
publish(server_args, role="scheduler")
|
||||||
parent_process = psutil.Process().parent()
|
parent_process = psutil.Process().parent()
|
||||||
|
|
||||||
# Set up tracing
|
# Set up tracing
|
||||||
|
|||||||
@@ -16,13 +16,12 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
from typing import TYPE_CHECKING, Any, Dict, Optional
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
import ray
|
import ray
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
from sglang.srt.runtime_context import publish
|
||||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||||
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -100,6 +99,10 @@ class SchedulerActor:
|
|||||||
f"[TP{tp_rank}] Bound to NUMA node {numa_node} for GPU {actual_gpu_id}"
|
f"[TP{tp_rank}] Bound to NUMA node {numa_node} for GPU {actual_gpu_id}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# This actor constructs Scheduler directly (no run_scheduler_process),
|
||||||
|
# which reads the config namespaces before the model worker's publish.
|
||||||
|
publish(server_args, role="scheduler")
|
||||||
|
|
||||||
# Create scheduler (loads model into GPU, initializes NCCL)
|
# Create scheduler (loads model into GPU, initializes NCCL)
|
||||||
self.scheduler = Scheduler(
|
self.scheduler = Scheduler(
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
|
|||||||
@@ -574,9 +574,17 @@ class _ConfigBag:
|
|||||||
writers are ``get_context().override(source, ...)`` (permanent) and
|
writers are ``get_context().override(source, ...)`` (permanent) and
|
||||||
the scoped ``.override(**kw)`` context manager (tests). Sub-namespaces
|
the scoped ``.override(**kw)`` context manager (tests). Sub-namespaces
|
||||||
(e.g. ``exec.moe``) are nested ``_ConfigBag`` instances reached by attribute.
|
(e.g. ``exec.moe``) are nested ``_ConfigBag`` instances reached by attribute.
|
||||||
"""
|
|
||||||
|
|
||||||
__slots__ = ("_path", "_fields", "_subs")
|
Leaves and sub-bags are stored as **real instance attributes** (in
|
||||||
|
``__dict__``), so ``bag.leaf`` / ``bag.sub`` is a plain attribute load that
|
||||||
|
``torch.compile`` / dynamo can trace — config reads inside a compiled model
|
||||||
|
forward (e.g. ``get_exec().comm.enable_symm_mem`` in the embedding layer)
|
||||||
|
must not graph-break. ``_fields`` / ``_subs`` keep the authoritative
|
||||||
|
name→value maps used for override routing, membership, and scoped restore;
|
||||||
|
``__getattr__`` is only a fallback for genuinely absent names. (Deliberately
|
||||||
|
no ``__slots__``: leaves are dynamic, and the ``__dict__`` is what makes the
|
||||||
|
reads traceable.)
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(self, path: str):
|
def __init__(self, path: str):
|
||||||
object.__setattr__(self, "_path", path)
|
object.__setattr__(self, "_path", path)
|
||||||
@@ -584,7 +592,9 @@ class _ConfigBag:
|
|||||||
object.__setattr__(self, "_subs", {}) # {subname: _ConfigBag}
|
object.__setattr__(self, "_subs", {}) # {subname: _ConfigBag}
|
||||||
|
|
||||||
def __getattr__(self, name: str) -> Any:
|
def __getattr__(self, name: str) -> Any:
|
||||||
# Reached only when ``name`` is not a real attribute (slot).
|
# Fallback only: real leaves/sub-bags resolve via __dict__ before this
|
||||||
|
# runs. Uses object.__getattribute__ (not self._fields) to stay safe if
|
||||||
|
# invoked before __init__ populates the bookkeeping dicts.
|
||||||
fields = object.__getattribute__(self, "_fields")
|
fields = object.__getattribute__(self, "_fields")
|
||||||
if name in fields:
|
if name in fields:
|
||||||
return fields[name]
|
return fields[name]
|
||||||
@@ -601,8 +611,16 @@ class _ConfigBag:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _set(self, name: str, value: Any) -> None:
|
def _set(self, name: str, value: Any) -> None:
|
||||||
"""Internal write (publish + override) that bypasses the read-only guard."""
|
"""Internal write (publish + override) that bypasses the read-only guard.
|
||||||
|
Updates both the bookkeeping map and the real attribute (traceable read)."""
|
||||||
object.__getattribute__(self, "_fields")[name] = value
|
object.__getattribute__(self, "_fields")[name] = value
|
||||||
|
object.__setattr__(self, name, value)
|
||||||
|
|
||||||
|
def _set_sub(self, name: str, sub: _ConfigBag) -> None:
|
||||||
|
"""Register a nested bag as both a bookkeeping entry and a real
|
||||||
|
attribute (so ``bag.sub`` is a plain, traceable attribute load)."""
|
||||||
|
object.__getattribute__(self, "_subs")[name] = sub
|
||||||
|
object.__setattr__(self, name, sub)
|
||||||
|
|
||||||
def __contains__(self, name: str) -> bool:
|
def __contains__(self, name: str) -> bool:
|
||||||
return name in object.__getattribute__(self, "_fields")
|
return name in object.__getattribute__(self, "_fields")
|
||||||
@@ -617,11 +635,13 @@ class _ConfigBag:
|
|||||||
path = object.__getattribute__(self, "_path")
|
path = object.__getattribute__(self, "_path")
|
||||||
raise ValueError(f"unknown config leaf for {path!r}: {sorted(unknown)}")
|
raise ValueError(f"unknown config leaf for {path!r}: {sorted(unknown)}")
|
||||||
saved = {name: fields[name] for name in kwargs}
|
saved = {name: fields[name] for name in kwargs}
|
||||||
fields.update(kwargs)
|
for name, value in kwargs.items():
|
||||||
|
self._set(name, value)
|
||||||
try:
|
try:
|
||||||
yield self
|
yield self
|
||||||
finally:
|
finally:
|
||||||
fields.update(saved)
|
for name, value in saved.items():
|
||||||
|
self._set(name, value)
|
||||||
|
|
||||||
|
|
||||||
def _build_config_bags(server_args: Any) -> dict:
|
def _build_config_bags(server_args: Any) -> dict:
|
||||||
@@ -660,7 +680,8 @@ def _build_config_bags(server_args: Any) -> dict:
|
|||||||
subs = object.__getattribute__(bag, "_subs")
|
subs = object.__getattribute__(bag, "_subs")
|
||||||
child = subs.get(name)
|
child = subs.get(name)
|
||||||
if child is None:
|
if child is None:
|
||||||
child = subs[name] = _ConfigBag(".".join(parts[: depth + 1]))
|
child = _ConfigBag(".".join(parts[: depth + 1]))
|
||||||
|
bag._set_sub(name, child)
|
||||||
bag = child
|
bag = child
|
||||||
if field in object.__getattribute__(bag, "_subs"):
|
if field in object.__getattribute__(bag, "_subs"):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -709,6 +730,7 @@ class RuntimeContext:
|
|||||||
"_server_args",
|
"_server_args",
|
||||||
"_config_bags",
|
"_config_bags",
|
||||||
"_overrides_log",
|
"_overrides_log",
|
||||||
|
"_publish_role",
|
||||||
"flags",
|
"flags",
|
||||||
"resources",
|
"resources",
|
||||||
"forward",
|
"forward",
|
||||||
@@ -719,6 +741,7 @@ class RuntimeContext:
|
|||||||
self._server_args: ServerArgs | None = None
|
self._server_args: ServerArgs | None = None
|
||||||
self._config_bags: dict | None = None
|
self._config_bags: dict | None = None
|
||||||
self._overrides_log: list = []
|
self._overrides_log: list = []
|
||||||
|
self._publish_role: str | None = None
|
||||||
self.flags = Flags()
|
self.flags = Flags()
|
||||||
self.resources = Resources()
|
self.resources = Resources()
|
||||||
self.forward = ForwardFlags()
|
self.forward = ForwardFlags()
|
||||||
@@ -784,8 +807,9 @@ class RuntimeContext:
|
|||||||
# leaves like pp_max_micro_batch_size are served via ParallelContext
|
# leaves like pp_max_micro_batch_size are served via ParallelContext
|
||||||
# __getattr__; live topology properties still win by name).
|
# __getattr__; live topology properties still win by name).
|
||||||
self.parallel._config = self._config_bags.get("parallel")
|
self.parallel._config = self._config_bags.get("parallel")
|
||||||
# Fresh config lifecycle: prior override provenance no longer applies.
|
# A direct install is roleless; ``publish`` assigns the role afterwards.
|
||||||
self._overrides_log = []
|
self._overrides_log = []
|
||||||
|
self._publish_role = None
|
||||||
|
|
||||||
def config_bag(self, name: str) -> _ConfigBag:
|
def config_bag(self, name: str) -> _ConfigBag:
|
||||||
"""Return the top-level config namespace bag (``device`` / ``model`` /
|
"""Return the top-level config namespace bag (``device`` / ``model`` /
|
||||||
@@ -843,8 +867,11 @@ class RuntimeContext:
|
|||||||
self._overrides_log.append((source, dict(fields)))
|
self._overrides_log.append((source, dict(fields)))
|
||||||
|
|
||||||
def overrides_log(self) -> list:
|
def overrides_log(self) -> list:
|
||||||
"""Provenance of post-publish ``override`` calls: ``[(source, {field: value})]``."""
|
"""Provenance of post-publish ``override`` calls: ``[(source, {field: value})]``.
|
||||||
return list(self._overrides_log)
|
|
||||||
|
Returns deep-ish copies (source, dict(fields)) so callers inspecting the
|
||||||
|
log cannot mutate the recorded provenance in place."""
|
||||||
|
return [(source, dict(fields)) for source, fields in self._overrides_log]
|
||||||
|
|
||||||
def resolved_server_args_dict(self, base: dict | None = None) -> dict:
|
def resolved_server_args_dict(self, base: dict | None = None) -> dict:
|
||||||
"""Serialize the *resolved* config: the pristine ``server_args`` fields
|
"""Serialize the *resolved* config: the pristine ``server_args`` fields
|
||||||
@@ -898,6 +925,7 @@ class RuntimeContext:
|
|||||||
prev_bags = self._config_bags
|
prev_bags = self._config_bags
|
||||||
prev_bag_values = _snapshot_bag_values(prev_bags)
|
prev_bag_values = _snapshot_bag_values(prev_bags)
|
||||||
prev_overrides_log = list(self._overrides_log)
|
prev_overrides_log = list(self._overrides_log)
|
||||||
|
prev_publish_role = self._publish_role
|
||||||
prev_parallel_config = self.parallel._config
|
prev_parallel_config = self.parallel._config
|
||||||
prev_capture = self.flags.capture.enable_torch_compile
|
prev_capture = self.flags.capture.enable_torch_compile
|
||||||
try:
|
try:
|
||||||
@@ -908,6 +936,7 @@ class RuntimeContext:
|
|||||||
if prev_bags is not None:
|
if prev_bags is not None:
|
||||||
_restore_bag_values(prev_bags, prev_bag_values)
|
_restore_bag_values(prev_bags, prev_bag_values)
|
||||||
self._overrides_log = prev_overrides_log
|
self._overrides_log = prev_overrides_log
|
||||||
|
self._publish_role = prev_publish_role
|
||||||
self.parallel._config = prev_parallel_config
|
self.parallel._config = prev_parallel_config
|
||||||
self.flags.capture.enable_torch_compile = prev_capture
|
self.flags.capture.enable_torch_compile = prev_capture
|
||||||
|
|
||||||
@@ -922,13 +951,21 @@ class _ServerArgsOverride:
|
|||||||
nondeterministic point.
|
nondeterministic point.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
__slots__ = ("_context", "_fields", "_previous", "_previous_capture", "_installed")
|
__slots__ = (
|
||||||
|
"_context",
|
||||||
|
"_fields",
|
||||||
|
"_prev_server_args",
|
||||||
|
"_prev_bags",
|
||||||
|
"_prev_overrides_log",
|
||||||
|
"_prev_publish_role",
|
||||||
|
"_prev_parallel_config",
|
||||||
|
"_prev_capture",
|
||||||
|
"_installed",
|
||||||
|
)
|
||||||
|
|
||||||
def __init__(self, context: RuntimeContext, fields: dict):
|
def __init__(self, context: RuntimeContext, fields: dict):
|
||||||
self._context = context
|
self._context = context
|
||||||
self._fields = fields
|
self._fields = fields
|
||||||
self._previous: ServerArgs | None = None
|
|
||||||
self._previous_capture = False
|
|
||||||
self._installed = False
|
self._installed = False
|
||||||
|
|
||||||
def install(self) -> ServerArgs:
|
def install(self) -> ServerArgs:
|
||||||
@@ -938,8 +975,13 @@ class _ServerArgsOverride:
|
|||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
|
||||||
assert not self._installed, "override_server_args already installed"
|
assert not self._installed, "override_server_args already installed"
|
||||||
self._previous = self._context._server_args
|
ctx = self._context
|
||||||
self._previous_capture = self._context.flags.capture.enable_torch_compile
|
self._prev_server_args = ctx._server_args
|
||||||
|
self._prev_bags = ctx._config_bags
|
||||||
|
self._prev_overrides_log = ctx._overrides_log
|
||||||
|
self._prev_publish_role = ctx._publish_role
|
||||||
|
self._prev_parallel_config = ctx.parallel._config
|
||||||
|
self._prev_capture = ctx.flags.capture.enable_torch_compile
|
||||||
server_args = ServerArgs(model_path="dummy")
|
server_args = ServerArgs(model_path="dummy")
|
||||||
if self._fields:
|
if self._fields:
|
||||||
server_args.override(source="test-override", **self._fields)
|
server_args.override(source="test-override", **self._fields)
|
||||||
@@ -948,24 +990,26 @@ class _ServerArgsOverride:
|
|||||||
# materialized so bare post-publish writes raise like they do on a
|
# materialized so bare post-publish writes raise like they do on a
|
||||||
# fully resolved config.
|
# fully resolved config.
|
||||||
object.__setattr__(server_args, "_declarations_materialized", True)
|
object.__setattr__(server_args, "_declarations_materialized", True)
|
||||||
self._context.set_server_args(server_args)
|
ctx.set_server_args(server_args)
|
||||||
self._installed = True
|
self._installed = True
|
||||||
return server_args
|
return server_args
|
||||||
|
|
||||||
def restore(self) -> None:
|
def restore(self) -> None:
|
||||||
"""Reinstate the previously published config (or the empty slot)."""
|
"""Reinstate the exact pre-install lifecycle state (or the empty slot)."""
|
||||||
if not self._installed:
|
if not self._installed:
|
||||||
return
|
return
|
||||||
self._installed = False
|
self._installed = False
|
||||||
previous, self._previous = self._previous, None
|
ctx = self._context
|
||||||
if previous is None:
|
ctx._server_args = self._prev_server_args
|
||||||
self._context._server_args = None
|
ctx._config_bags = self._prev_bags
|
||||||
else:
|
ctx._overrides_log = self._prev_overrides_log
|
||||||
self._context.set_server_args(previous)
|
ctx._publish_role = self._prev_publish_role
|
||||||
# set_server_args reseeds the capture tier from the published object
|
ctx.parallel._config = self._prev_parallel_config
|
||||||
# (and the empty-slot path does not touch it at all); the snapshot
|
ctx.flags.capture.enable_torch_compile = self._prev_capture
|
||||||
# puts back the exact pre-install runtime state either way.
|
self._prev_server_args = None
|
||||||
self._context.flags.capture.enable_torch_compile = self._previous_capture
|
self._prev_bags = None
|
||||||
|
self._prev_overrides_log = None
|
||||||
|
self._prev_parallel_config = None
|
||||||
|
|
||||||
def __enter__(self) -> ServerArgs:
|
def __enter__(self) -> ServerArgs:
|
||||||
return self.install()
|
return self.install()
|
||||||
@@ -1051,6 +1095,36 @@ def get_observability() -> _ConfigBag:
|
|||||||
return _CONTEXT.config_bag("observability")
|
return _CONTEXT.config_bag("observability")
|
||||||
|
|
||||||
|
|
||||||
|
def publish(server_args, *, role: str, hf_config: Any = None) -> RuntimeContext:
|
||||||
|
"""Install process-wide config for this OS process.
|
||||||
|
|
||||||
|
Records the process ``role`` (``tokenizer`` / ``scheduler`` /
|
||||||
|
``dp_controller`` / ``encoder`` / ``expert_backup`` /
|
||||||
|
``weight_cache_daemon`` / ``launcher`` / ``test``) and
|
||||||
|
projects the config bags. Draft workers skip publish (they must not clobber
|
||||||
|
the target). ``role`` is provenance today — per-role namespace projection
|
||||||
|
and fail-closed enforcement is a later unit. ``hf_config`` is accepted for
|
||||||
|
forward-compat and currently unused.
|
||||||
|
|
||||||
|
Normally one call per process, but re-publish is allowed and is
|
||||||
|
**last-publish-wins** (bags re-projected, provenance reset, role
|
||||||
|
overwritten). Two sanctioned multi-publish shapes exist: the in-process
|
||||||
|
Engine builds its ``TokenizerManager`` inside the launcher process (the
|
||||||
|
process ends up with the tokenizer publish), and multiple Engines in one
|
||||||
|
process publish in sequence — which is exactly why per-instance managers
|
||||||
|
must read ``self.server_args`` for anything engine-specific rather than
|
||||||
|
the process-global bags.
|
||||||
|
"""
|
||||||
|
_CONTEXT.set_server_args(server_args)
|
||||||
|
_CONTEXT._publish_role = role
|
||||||
|
return _CONTEXT
|
||||||
|
|
||||||
|
|
||||||
|
def publish_role() -> str | None:
|
||||||
|
"""The role recorded by the last ``publish`` (None for a legacy set)."""
|
||||||
|
return _CONTEXT._publish_role
|
||||||
|
|
||||||
|
|
||||||
def get_stream(name: str) -> Any:
|
def get_stream(name: str) -> Any:
|
||||||
return _CONTEXT.get_stream(name)
|
return _CONTEXT.get_stream(name)
|
||||||
|
|
||||||
@@ -1084,6 +1158,7 @@ def reset_context() -> None:
|
|||||||
_CONTEXT._server_args = None
|
_CONTEXT._server_args = None
|
||||||
_CONTEXT._config_bags = None
|
_CONTEXT._config_bags = None
|
||||||
_CONTEXT._overrides_log = []
|
_CONTEXT._overrides_log = []
|
||||||
|
_CONTEXT._publish_role = None
|
||||||
_CONTEXT.parallel._config = None
|
_CONTEXT.parallel._config = None
|
||||||
_CONTEXT.flags = Flags()
|
_CONTEXT.flags = Flags()
|
||||||
_CONTEXT.resources = Resources()
|
_CONTEXT.resources = Resources()
|
||||||
|
|||||||
@@ -9007,14 +9007,19 @@ class ServerArgs:
|
|||||||
# (decrease-only) by test/registered/unit/test_legacy_global_ratchet.py.
|
# (decrease-only) by test/registered/unit/test_legacy_global_ratchet.py.
|
||||||
# Imports are in-function so the two modules stay cycle-free at import time.
|
# Imports are in-function so the two modules stay cycle-free at import time.
|
||||||
def set_global_server_args_for_scheduler(server_args: ServerArgs):
|
def set_global_server_args_for_scheduler(server_args: ServerArgs):
|
||||||
"""Legacy publish shim — prefer ``get_context().set_server_args()`` from
|
"""Legacy publish shim (role=scheduler) — prefer
|
||||||
``sglang.srt.runtime_context`` in new code."""
|
``runtime_context.publish(server_args, role=...)`` in new code."""
|
||||||
from sglang.srt.runtime_context import get_context
|
from sglang.srt.runtime_context import publish
|
||||||
|
|
||||||
get_context().set_server_args(server_args)
|
publish(server_args, role="scheduler")
|
||||||
|
|
||||||
|
|
||||||
set_global_server_args_for_tokenizer = set_global_server_args_for_scheduler
|
def set_global_server_args_for_tokenizer(server_args: ServerArgs):
|
||||||
|
"""Legacy publish shim (role=tokenizer). Not aliased to the scheduler shim:
|
||||||
|
the process role differs."""
|
||||||
|
from sglang.srt.runtime_context import publish
|
||||||
|
|
||||||
|
publish(server_args, role="tokenizer")
|
||||||
|
|
||||||
|
|
||||||
def get_global_server_args() -> ServerArgs:
|
def get_global_server_args() -> ServerArgs:
|
||||||
|
|||||||
@@ -97,7 +97,9 @@ def build_draft_tp_worker(
|
|||||||
context_length=target_model_config.context_len,
|
context_length=target_model_config.context_len,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# The draft's layers must resolve config from the draft's own bags.
|
||||||
with get_context().preserve_config():
|
with get_context().preserve_config():
|
||||||
|
get_context().set_server_args(draft_server_args)
|
||||||
draft_worker = TpModelWorker(
|
draft_worker = TpModelWorker(
|
||||||
server_args=draft_server_args,
|
server_args=draft_server_args,
|
||||||
gpu_id=gpu_id,
|
gpu_id=gpu_id,
|
||||||
|
|||||||
@@ -46,6 +46,7 @@ import torch.distributed as dist
|
|||||||
|
|
||||||
from sglang.srt.configs.load_config import LoadConfig
|
from sglang.srt.configs.load_config import LoadConfig
|
||||||
from sglang.srt.platforms import current_platform
|
from sglang.srt.platforms import current_platform
|
||||||
|
from sglang.srt.runtime_context import publish
|
||||||
from sglang.srt.utils import MultiprocessingSerializer
|
from sglang.srt.utils import MultiprocessingSerializer
|
||||||
|
|
||||||
from .protocol import (
|
from .protocol import (
|
||||||
@@ -208,7 +209,6 @@ class WeightCacheDaemon:
|
|||||||
from sglang.srt.configs.device_config import DeviceConfig
|
from sglang.srt.configs.device_config import DeviceConfig
|
||||||
from sglang.srt.configs.model_config import ModelConfig
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
from sglang.srt.model_loader.loader import get_model_loader
|
from sglang.srt.model_loader.loader import get_model_loader
|
||||||
from sglang.srt.runtime_context import get_context
|
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
|
||||||
server_args = ServerArgs(
|
server_args = ServerArgs(
|
||||||
@@ -223,7 +223,7 @@ class WeightCacheDaemon:
|
|||||||
load_format=self.load_format,
|
load_format=self.load_format,
|
||||||
model_loader_extra_config=self.model_loader_extra_config,
|
model_loader_extra_config=self.model_loader_extra_config,
|
||||||
)
|
)
|
||||||
get_context().set_server_args(server_args)
|
publish(server_args, role="weight_cache_daemon")
|
||||||
|
|
||||||
# Initialize distributed backend for model loading
|
# Initialize distributed backend for model loading
|
||||||
# (must be done after server_args and model_config are available)
|
# (must be done after server_args and model_config are available)
|
||||||
|
|||||||
@@ -31,7 +31,7 @@ _RATCHETS = [
|
|||||||
(
|
(
|
||||||
"set_global_server_args_for_*",
|
"set_global_server_args_for_*",
|
||||||
r"\bset_global_server_args_for_(?:scheduler|tokenizer)\s*\(",
|
r"\bset_global_server_args_for_(?:scheduler|tokenizer)\s*\(",
|
||||||
5,
|
4,
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|||||||
@@ -215,8 +215,10 @@ class TestServerArgsOwnership(_IsolatedServerArgs):
|
|||||||
self.assertIs(get_server_args(), sentinel)
|
self.assertIs(get_server_args(), sentinel)
|
||||||
self.assertIs(get_context().server_args, sentinel)
|
self.assertIs(get_context().server_args, sentinel)
|
||||||
|
|
||||||
def test_tokenizer_alias_is_same_function(self):
|
def test_tokenizer_alias_is_distinct_role_shim(self):
|
||||||
self.assertIs(
|
# Deliberately NOT an alias: the two legacy setters publish with
|
||||||
|
# different process roles (scheduler vs tokenizer).
|
||||||
|
self.assertIsNot(
|
||||||
server_args_module.set_global_server_args_for_tokenizer,
|
server_args_module.set_global_server_args_for_tokenizer,
|
||||||
server_args_module.set_global_server_args_for_scheduler,
|
server_args_module.set_global_server_args_for_scheduler,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -164,6 +164,34 @@ class TestContextOverride(CustomTestCase):
|
|||||||
self.assertIs(rc.get_context().server_args, target)
|
self.assertIs(rc.get_context().server_args, target)
|
||||||
self.assertEqual(rc.get_schedule().page_size, 16)
|
self.assertEqual(rc.get_schedule().page_size, 16)
|
||||||
|
|
||||||
|
def test_publish_records_role(self):
|
||||||
|
rc.publish(ServerArgs(model_path="dummy"), role="scheduler")
|
||||||
|
self.assertEqual(rc.publish_role(), "scheduler")
|
||||||
|
|
||||||
|
def test_legacy_shims_record_roles(self):
|
||||||
|
# Unit 2a: the legacy setters publish with their process role.
|
||||||
|
from sglang.srt.server_args import (
|
||||||
|
set_global_server_args_for_scheduler,
|
||||||
|
set_global_server_args_for_tokenizer,
|
||||||
|
)
|
||||||
|
|
||||||
|
set_global_server_args_for_scheduler(ServerArgs(model_path="dummy"))
|
||||||
|
self.assertEqual(rc.publish_role(), "scheduler")
|
||||||
|
set_global_server_args_for_tokenizer(ServerArgs(model_path="dummy"))
|
||||||
|
self.assertEqual(rc.publish_role(), "tokenizer")
|
||||||
|
|
||||||
|
def test_reset_clears_role(self):
|
||||||
|
rc.publish(ServerArgs(model_path="dummy"), role="test")
|
||||||
|
rc.reset_context()
|
||||||
|
self.assertIsNone(rc.publish_role())
|
||||||
|
|
||||||
|
def test_direct_install_clears_role(self):
|
||||||
|
# A role-less set_server_args (test overrides, draft-worker builds)
|
||||||
|
# must not inherit the previous lifecycle's role.
|
||||||
|
rc.publish(ServerArgs(model_path="dummy"), role="scheduler")
|
||||||
|
rc.get_context().set_server_args(ServerArgs(model_path="dummy"))
|
||||||
|
self.assertIsNone(rc.publish_role())
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user