[refactor] Add the per-forward flags tier: ctx.forward (#30490)

This commit is contained in:
Cheng Wan
2026-07-09 02:09:03 -07:00
committed by GitHub
parent 65b14881c5
commit fef2128e19
7 changed files with 321 additions and 63 deletions
+18 -15
View File
@@ -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()
+9 -11
View File
@@ -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:
@@ -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
@@ -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
+124 -2
View File
@@ -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()
+4 -19
View File
@@ -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(
@@ -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."""