From fef2128e190953e2d27ce5cd1d17e8576aae8f01 Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Thu, 9 Jul 2026 02:09:03 -0700 Subject: [PATCH] [refactor] Add the per-forward flags tier: ctx.forward (#30490) --- python/sglang/srt/layers/communicator.py | 33 ++-- python/sglang/srt/layers/dp_attention.py | 20 +-- .../sglang/srt/layers/moe/moe_runner/base.py | 19 +-- .../moe/moe_runner/flashinfer_trtllm.py | 5 +- python/sglang/srt/runtime_context.py | 126 +++++++++++++- python/sglang/srt/utils/multi_stream_utils.py | 23 +-- test/registered/unit/test_runtime_context.py | 158 ++++++++++++++++++ 7 files changed, 321 insertions(+), 63 deletions(-) diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index 65e5441eb..97f3986f3 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -73,7 +73,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( check_cuda_graph_backend, ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_forward, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.utils import ( @@ -262,8 +262,6 @@ class AttentionInputs: class AttnTpContext: def __init__(self): self.allow_input_scattered = False - self.input_scattered_ = False - self.attn_inputs_: Optional[AttentionInputs] = None self.is_dsa = False def init_context(self, q_lora_rank, is_dsa): @@ -299,30 +297,35 @@ class AttnTpContext: @property def input_scattered(self): - return self.input_scattered_ + return get_forward().attn_input_scattered def set_attn_inputs(self, attn_inputs: AttentionInputs): - self.attn_inputs_ = attn_inputs + get_forward().set("attn_inputs", attn_inputs) def fetch_qkv_latent(self): - assert self.attn_inputs_ is not None - return self.attn_inputs_.fetch_qkv_latent() + attn_inputs = get_forward().attn_inputs + assert attn_inputs is not None + return attn_inputs.fetch_qkv_latent() def fetch_hidden_states(self): - assert self.attn_inputs_ is not None - return self.attn_inputs_.fetch_hidden_states() + attn_inputs = get_forward().attn_inputs + assert attn_inputs is not None + return attn_inputs.fetch_hidden_states() def clear_attn_inputs(self) -> None: - self.attn_inputs_ = None + get_forward().set("attn_inputs", None) @contextmanager def maybe_input_scattered(self, forward_batch: ForwardBatch): flag = self.use_input_scattered(forward_batch) - old_flag = self.input_scattered - self.input_scattered_ = flag - yield - self.input_scattered_ = old_flag - self.attn_inputs_ = None + forward = get_forward() + # scoped() also restores when the forward raises — the old in-place + # swap leaked the flag on exceptions. + with forward.scoped(attn_input_scattered=flag): + try: + yield + finally: + forward.set("attn_inputs", None) ATTN_TP_CONTEXT = AttnTpContext() diff --git a/python/sglang/srt/layers/dp_attention.py b/python/sglang/srt/layers/dp_attention.py index 77b7992a1..c513cbdc4 100644 --- a/python/sglang/srt/layers/dp_attention.py +++ b/python/sglang/srt/layers/dp_attention.py @@ -108,7 +108,6 @@ class _DpGatheredBufferWrapper: _local_dp_buffer_len: int _dp_max_padding: bool _global_num_tokens: Optional[List[int]] - _is_extend_in_batch: bool @classmethod def set_metadata(cls, hidden_size: int, dtype: torch.dtype, device: torch.device): @@ -173,14 +172,6 @@ class _DpGatheredBufferWrapper: def get_dp_device(cls) -> torch.device: return cls._device - @classmethod - def set_is_extend_in_batch(cls, is_extend_in_batch: bool): - cls._is_extend_in_batch = is_extend_in_batch - - @classmethod - def get_is_extend_in_batch(cls) -> bool: - return cls._is_extend_in_batch - @classmethod def is_dp_max_padding(cls) -> bool: return cls._dp_max_padding @@ -230,11 +221,18 @@ def get_dp_device() -> torch.device: def set_is_extend_in_batch(is_extend_in_batch: bool): - _DpGatheredBufferWrapper.set_is_extend_in_batch(is_extend_in_batch) + # Sticky within the thread: every ForwardBatch construction writes it, + # graph runners force False around capture; readers are the EP + # dispatchers on the same (single) forward thread. + from sglang.srt.runtime_context import get_forward + + get_forward().set("is_extend_in_batch", is_extend_in_batch) def get_is_extend_in_batch() -> bool: - return _DpGatheredBufferWrapper.get_is_extend_in_batch() + from sglang.srt.runtime_context import get_forward + + return get_forward().is_extend_in_batch def is_dp_max_padding() -> bool: diff --git a/python/sglang/srt/layers/moe/moe_runner/base.py b/python/sglang/srt/layers/moe/moe_runner/base.py index 9a4c2aa0d..ef913a75f 100644 --- a/python/sglang/srt/layers/moe/moe_runner/base.py +++ b/python/sglang/srt/layers/moe/moe_runner/base.py @@ -1,10 +1,8 @@ from __future__ import annotations -import contextvars from abc import ABC, abstractmethod -from contextlib import contextmanager from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, Callable, Generator, Optional, Tuple, TypeGuard +from typing import TYPE_CHECKING, Any, Callable, Optional, Tuple, TypeGuard import torch @@ -28,18 +26,11 @@ if TYPE_CHECKING: ) -_moe_output_buf: contextvars.ContextVar[Optional[torch.Tensor]] = ( - contextvars.ContextVar("moe_output_buf", default=None) -) +def moe_output_buffer_ctx(buf: torch.Tensor): + """Provide the MoE output buffer for the current forward scope.""" + from sglang.srt.runtime_context import get_forward - -@contextmanager -def moe_output_buffer_ctx(buf: torch.Tensor) -> Generator[None, None, None]: - token = _moe_output_buf.set(buf) - try: - yield - finally: - _moe_output_buf.reset(token) + return get_forward().scoped(moe_output_buffer=buf) @dataclass diff --git a/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py b/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py index 19a1df20f..a9d9c15e8 100644 --- a/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py +++ b/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py @@ -26,7 +26,6 @@ from sglang.srt.layers.moe.flashinfer_trtllm_moe import ( from sglang.srt.layers.moe.moe_runner.base import ( MoeQuantInfo, MoeRunnerConfig, - _moe_output_buf, register_fused_func, ) from sglang.srt.layers.quantization.fp8_kernel import ( @@ -1001,7 +1000,9 @@ def fused_experts_none_to_flashinfer_trtllm_fp4( output_dtype = ( hidden_states.dtype if hidden_states_scale is None else torch.bfloat16 ) - _provided = _moe_output_buf.get() + from sglang.srt.runtime_context import get_forward + + _provided = get_forward().moe_output_buffer _symm_required = is_allocation_symmetric() if ( _provided is not None diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index cc8a112a2..bf75292d8 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -348,17 +348,134 @@ class Resources(_FlagGroupBase): tbo_event_pool: dict = dataclasses.field(default_factory=dict) +class ForwardFlags: + """Per-forward runtime flags with one API and two backings. + + Flags read only from eager Python are backed by context variables, so + nested scopes and threads stay isolated (a new thread sees the defaults). + Flags that are read or written *inside torch.compile-traced model code* + (``_GRAPH_VISIBLE``) are backed by plain dict slots instead: dynamo + cannot trace ``ContextVar.get``/``set``, while plain reads it guards on + — the storage form these flags had before joining the tier. Their + writers and readers are single-threaded per process (TBO interleaves + ubatches on one thread; attention-TP input scattering excludes TBO), so + context isolation is not needed for correctness. + + ``scoped(**kw)`` — the one regular write path — restores on exit for + both backings. ``set()`` exists for the legacy unscoped setters' shims. + """ + + _DEFAULTS = { + "multi_stream": False, + "moe_output_buffer": None, + # Attention-TP input-scattering (set per forward by + # AttnTpContext.maybe_input_scattered / set_attn_inputs). + "attn_input_scattered": False, + "attn_inputs": None, + # Sticky across forwards: every ForwardBatch construction writes it; + # graph runners force False around capture. + "is_extend_in_batch": False, + } + + # Read/written inside compiled graphs (vocab embedding, communicator, + # EP dispatch, DP gather/scatter): plain-slot backed. Before moving a + # flag out of this set, prove no read/write site sits under + # torch.compile. + _GRAPH_VISIBLE = frozenset( + { + "attn_input_scattered", + "attn_inputs", + "is_extend_in_batch", + } + ) + + __slots__ = ("_vars", "_plain") + + def __init__(self): + import contextvars + + object.__setattr__( + self, + "_plain", + { + name: default + for name, default in self._DEFAULTS.items() + if name in self._GRAPH_VISIBLE + }, + ) + object.__setattr__( + self, + "_vars", + { + name: contextvars.ContextVar(f"forward.{name}", default=default) + for name, default in self._DEFAULTS.items() + if name not in self._GRAPH_VISIBLE + }, + ) + + def __getattr__(self, name: str) -> Any: + plain = self._plain + if name in plain: + return plain[name] + try: + return self._vars[name].get() + except KeyError: + raise AttributeError( + f"ForwardFlags has no flag '{name}' (flags are declared in " + "ForwardFlags._DEFAULTS; check for typos)" + ) from None + + def __setattr__(self, name: str, value: Any) -> None: + raise AttributeError( + "ForwardFlags is written through scoped(**kw) (or the legacy " + "set() shim), never by attribute assignment" + ) + + def set(self, name: str, value: Any) -> None: + """Unscoped write for legacy setter shims; persists until the next + write (current context only, for contextvar-backed flags).""" + if name in self._plain: + self._plain[name] = value + else: + self._vars[name].set(value) + + @contextmanager + def scoped(self, **kwargs): + """Set flags for the current scope, restoring on exit. Transactional + (keys validated before any write) and exception-safe.""" + unknown = set(kwargs) - set(self._DEFAULTS) + if unknown: + raise ValueError(f"unknown forward flag(s): {sorted(unknown)}") + plain_saved = [ + (name, self._plain[name]) for name in kwargs if name in self._plain + ] + tokens = [] + for name, value in kwargs.items(): + if name in self._plain: + self._plain[name] = value + else: + tokens.append((self._vars[name], self._vars[name].set(value))) + try: + yield self + finally: + for var, token in reversed(tokens): + var.reset(token) + for name, value in reversed(plain_saved): + self._plain[name] = value + + class RuntimeContext: """Container for the structured runtime accessors; exposes ``parallel``, - ``server_args``, ``flags``, and ``resources``.""" + ``server_args``, ``flags``, ``resources``, and ``forward``.""" - __slots__ = ("parallel", "_server_args", "flags", "resources") + __slots__ = ("parallel", "_server_args", "flags", "resources", "forward") def __init__(self, parallel: ParallelContext): self.parallel = parallel self._server_args: ServerArgs | None = None self.flags = Flags() self.resources = Resources() + self.forward = ForwardFlags() def get_stream(self, name: str) -> Any: """Named process-level CUDA side stream: get-or-create, shared by @@ -439,6 +556,10 @@ def get_resources() -> Resources: return _CONTEXT.resources +def get_forward() -> ForwardFlags: + return _CONTEXT.forward + + def get_stream(name: str) -> Any: return _CONTEXT.get_stream(name) @@ -460,3 +581,4 @@ def reset_context() -> None: _CONTEXT._server_args = None _CONTEXT.flags = Flags() _CONTEXT.resources = Resources() + _CONTEXT.forward = ForwardFlags() diff --git a/python/sglang/srt/utils/multi_stream_utils.py b/python/sglang/srt/utils/multi_stream_utils.py index fb3d2cf7e..5bc174399 100644 --- a/python/sglang/srt/utils/multi_stream_utils.py +++ b/python/sglang/srt/utils/multi_stream_utils.py @@ -1,37 +1,22 @@ # Adapted from trtllm. -import threading -from contextlib import contextmanager from typing import Any, Callable, Optional import torch - -class do_multi_stream_local(threading.local): - - def __init__(self): - self.do_multi_stream = False - - -_local = do_multi_stream_local() +from sglang.srt.runtime_context import get_forward def set_do_multi_stream(enable: bool): - _local.do_multi_stream = enable + get_forward().set("multi_stream", enable) def do_multi_stream() -> bool: - return _local.do_multi_stream + return get_forward().multi_stream -@contextmanager def with_multi_stream(enable: bool): - prev_do_multi_stream = _local.do_multi_stream - set_do_multi_stream(enable) - try: - yield - finally: - set_do_multi_stream(prev_do_multi_stream) + return get_forward().scoped(multi_stream=enable) def maybe_execute_in_parallel( diff --git a/test/registered/unit/test_runtime_context.py b/test/registered/unit/test_runtime_context.py index f847597e6..e52cdbced 100644 --- a/test/registered/unit/test_runtime_context.py +++ b/test/registered/unit/test_runtime_context.py @@ -500,6 +500,164 @@ class TestEpBufferState(_IsolatedServerArgs): self.assertIsNone(DeepEPBuffer._state().buffer) +class TestForwardFlags(_IsolatedServerArgs): + """ctx.forward: contextvar-backed per-forward flags; scoped() restores, + threads see defaults.""" + + def test_scoped_set_restore_and_nesting(self): + from sglang.srt.runtime_context import get_forward + + reset_context() + fwd = get_forward() + self.assertFalse(fwd.multi_stream) + with fwd.scoped(multi_stream=True): + self.assertTrue(fwd.multi_stream) + with fwd.scoped(multi_stream=False): + self.assertFalse(fwd.multi_stream) + self.assertTrue(fwd.multi_stream) + self.assertFalse(fwd.multi_stream) + + def test_scoped_restores_on_exception_and_validates_keys(self): + from sglang.srt.runtime_context import get_forward + + reset_context() + fwd = get_forward() + with self.assertRaises(RuntimeError): + with fwd.scoped(moe_output_buffer="buf"): + raise RuntimeError("boom") + self.assertIsNone(fwd.moe_output_buffer) + with self.assertRaises(ValueError): + with fwd.scoped(nope=1): + pass + with self.assertRaises(AttributeError): + fwd.multi_stream = True # attribute writes are rejected + + def test_threads_see_defaults(self): + import threading + + from sglang.srt.runtime_context import get_forward + + reset_context() + fwd = get_forward() + seen = {} + with fwd.scoped(multi_stream=True): + + def probe(): + seen["value"] = get_forward().multi_stream + + worker = threading.Thread(target=probe) + worker.start() + worker.join() + self.assertFalse(seen["value"]) # a new thread sees the default + + def test_graph_visible_flags_trace_under_torch_compile(self): + # Regression: dynamo cannot trace ContextVar.get, and these flags are + # read inside compiled model code (vocab embedding, communicator, DP + # gather) — they must stay plain-slot backed. fullgraph=True turns + # any graph break back into a failure. + import torch + + from sglang.srt.runtime_context import get_forward + + reset_context() + + @torch.compile(fullgraph=True, backend="eager", dynamic=False) + def probe(x): + fwd = get_forward() + if fwd.attn_input_scattered: + x = x + 1 + if fwd.is_extend_in_batch: + x = x + 2 + return x + + self.assertEqual(probe(torch.zeros(())).item(), 0) + with get_forward().scoped(attn_input_scattered=True): + self.assertEqual(probe(torch.zeros(())).item(), 1) + get_forward().set("is_extend_in_batch", True) + self.assertEqual(probe(torch.zeros(())).item(), 2) + get_forward().set("is_extend_in_batch", False) + + def test_graph_visible_flags_are_process_visible_across_threads(self): + # Documented divergence from the contextvar-backed flags: plain slots + # are process-global (the storage form these flags had before the + # tier), so another thread sees the current value, not the default. + import threading + + from sglang.srt.runtime_context import get_forward + + reset_context() + seen = {} + with get_forward().scoped(attn_input_scattered=True): + + def probe(): + seen["value"] = get_forward().attn_input_scattered + + worker = threading.Thread(target=probe) + worker.start() + worker.join() + self.assertTrue(seen["value"]) + self.assertFalse(get_forward().attn_input_scattered) + + def test_multi_stream_shims(self): + from sglang.srt.utils.multi_stream_utils import ( + do_multi_stream, + with_multi_stream, + ) + + reset_context() + self.assertFalse(do_multi_stream()) + with with_multi_stream(True): + self.assertTrue(do_multi_stream()) + self.assertFalse(do_multi_stream()) + + def test_attn_tp_context_per_forward_slots(self): + from types import SimpleNamespace + + from sglang.srt.layers.communicator import get_attn_tp_context + from sglang.srt.runtime_context import get_forward + + reset_context() + ctx = get_attn_tp_context() + self.assertFalse(ctx.input_scattered) + fb = SimpleNamespace( + forward_mode=SimpleNamespace( + is_extend=lambda: False, is_target_verify=lambda: False + ), + input_ids=None, + can_run_tbo=False, + ) + sentinel = SimpleNamespace(fetch_qkv_latent=lambda: "qkv") + with ctx.maybe_input_scattered(fb): + ctx.set_attn_inputs(sentinel) + self.assertEqual(ctx.fetch_qkv_latent(), "qkv") + # attn inputs are cleared at scope exit, flag restored + self.assertIsNone(get_forward().attn_inputs) + self.assertFalse(ctx.input_scattered) + + def test_is_extend_in_batch_sticky_within_thread(self): + from sglang.srt.layers.dp_attention import ( + get_is_extend_in_batch, + set_is_extend_in_batch, + ) + + reset_context() + self.assertFalse(get_is_extend_in_batch()) + set_is_extend_in_batch(True) + self.assertTrue(get_is_extend_in_batch()) # sticky until next write + set_is_extend_in_batch(False) + self.assertFalse(get_is_extend_in_batch()) + + def test_moe_output_buffer_ctx(self): + from sglang.srt.layers.moe.moe_runner.base import moe_output_buffer_ctx + from sglang.srt.runtime_context import get_forward + + reset_context() + sentinel = object() + with moe_output_buffer_ctx(sentinel): + self.assertIs(get_forward().moe_output_buffer, sentinel) + self.assertIsNone(get_forward().moe_output_buffer) + + class TestPublishLifecycle(_IsolatedServerArgs): """Publish installs the resolved server_args and seeds the capture tier."""