config: preserve resolved config across nested publishes + mutation ratchets (#33011)
This commit is contained in:
@@ -671,6 +671,34 @@ def _build_config_bags(server_args: Any) -> dict:
|
||||
return tops
|
||||
|
||||
|
||||
def _snapshot_bag_values(bags: dict | None) -> dict | None:
|
||||
"""Per-leaf value snapshot of a config-bag tree (bags are mutated in
|
||||
place by ``override``, so reference snapshots alias live state)."""
|
||||
if bags is None:
|
||||
return None
|
||||
snap: dict = {}
|
||||
|
||||
def walk(prefix: str, bag) -> None:
|
||||
snap[prefix] = dict(object.__getattribute__(bag, "_fields"))
|
||||
for name, sub in object.__getattribute__(bag, "_subs").items():
|
||||
walk(f"{prefix}.{name}", sub)
|
||||
|
||||
for name, bag in bags.items():
|
||||
walk(name, bag)
|
||||
return snap
|
||||
|
||||
|
||||
def _restore_bag_values(bags: dict, snap: dict) -> None:
|
||||
def walk(prefix: str, bag) -> None:
|
||||
for key, value in snap[prefix].items():
|
||||
bag._set(key, value)
|
||||
for name, sub in object.__getattribute__(bag, "_subs").items():
|
||||
walk(f"{prefix}.{name}", sub)
|
||||
|
||||
for name, bag in bags.items():
|
||||
walk(name, bag)
|
||||
|
||||
|
||||
class RuntimeContext:
|
||||
"""Container for the structured runtime accessors; exposes ``parallel``,
|
||||
``server_args``, the resolved config namespace bags, ``flags``,
|
||||
@@ -858,6 +886,31 @@ class RuntimeContext:
|
||||
"""
|
||||
return _ServerArgsOverride(self, fields)
|
||||
|
||||
@contextmanager
|
||||
def preserve_config(self):
|
||||
"""Snapshot the full config lifecycle and reinstate it verbatim on exit.
|
||||
|
||||
For nested construction steps that publish a private ``ServerArgs``
|
||||
copy (e.g. a draft-worker build) and must leave the enclosing
|
||||
lifecycle — including its post-publish overrides — untouched.
|
||||
"""
|
||||
prev_server_args = self._server_args
|
||||
prev_bags = self._config_bags
|
||||
prev_bag_values = _snapshot_bag_values(prev_bags)
|
||||
prev_overrides_log = list(self._overrides_log)
|
||||
prev_parallel_config = self.parallel._config
|
||||
prev_capture = self.flags.capture.enable_torch_compile
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
self._server_args = prev_server_args
|
||||
self._config_bags = prev_bags
|
||||
if prev_bags is not None:
|
||||
_restore_bag_values(prev_bags, prev_bag_values)
|
||||
self._overrides_log = prev_overrides_log
|
||||
self.parallel._config = prev_parallel_config
|
||||
self.flags.capture.enable_torch_compile = prev_capture
|
||||
|
||||
|
||||
class _ServerArgsOverride:
|
||||
"""Scoped config override (see ``RuntimeContext.override_server_args``).
|
||||
|
||||
@@ -10,7 +10,7 @@ import torch
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
|
||||
from sglang.srt.runtime_context import get_context, get_server_args
|
||||
from sglang.srt.runtime_context import get_context
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.speculative.dflash_info import DFlashVerifyInput
|
||||
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
|
||||
@@ -97,8 +97,7 @@ def build_draft_tp_worker(
|
||||
context_length=target_model_config.context_len,
|
||||
)
|
||||
|
||||
saved_server_args = get_server_args()
|
||||
try:
|
||||
with get_context().preserve_config():
|
||||
draft_worker = TpModelWorker(
|
||||
server_args=draft_server_args,
|
||||
gpu_id=gpu_id,
|
||||
@@ -106,8 +105,6 @@ def build_draft_tp_worker(
|
||||
nccl_port=nccl_port,
|
||||
is_draft_worker=True,
|
||||
)
|
||||
finally:
|
||||
get_context().set_server_args(saved_server_args)
|
||||
|
||||
draft_model_runner = draft_worker.model_runner
|
||||
draft_worker.draft_runner = draft_model_runner
|
||||
|
||||
Reference in New Issue
Block a user