[Scheduler] Unify WAR read-done gating behind shared-read boundary declarations (#34052)

This commit is contained in:
Liangsheng Yin
2026-08-08 03:26:36 -07:00
committed by GitHub
parent 2c0188cc78
commit a1ca76b24b
11 changed files with 105 additions and 70 deletions
@@ -1,6 +1,7 @@
from __future__ import annotations from __future__ import annotations
from abc import ABC from abc import ABC
from enum import Enum, auto
from typing import TYPE_CHECKING, Optional from typing import TYPE_CHECKING, Optional
import torch import torch
@@ -14,10 +15,21 @@ if TYPE_CHECKING:
) )
from sglang.srt.layers.attention.verify_mask import VerifyMask from sglang.srt.layers.attention.verify_mask import VerifyMask
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.speculative.spec_info import SpecInput
class SharedReadBoundary(Enum):
"""Where a backend's scheduler-shared reads end, relative to the replay;
the WAR read-done record must land at or after this point. IN_REPLAY
means at the captured (in-graph) metadata init."""
PRE_REPLAY = auto()
IN_REPLAY = auto()
POST_REPLAY = auto()
UNKNOWN = auto() # not audited -> coarse whole-forward fence
class AttentionBackend(ABC): class AttentionBackend(ABC):
"""The base class of attention backends. """The base class of attention backends.
@@ -113,8 +125,16 @@ class AttentionBackend(ABC):
# object during capture, and refresh its dynamic fields before each replay. # object during capture, and refresh its dynamic fields before each replay.
use_captured_forward_metadata_for_breakable_cuda_graph: bool = False use_captured_forward_metadata_for_breakable_cuda_graph: bool = False
# Whether prefill metadata initialization finishes all scheduler-shared reads. def shared_read_boundary(self, forward_mode: ForwardMode) -> SharedReadBoundary:
prefill_shared_reads_end_at_metadata_init: bool = False """Declare where this backend's scheduler-shared reads end per mode.
Decode/verify default to IN_REPLAY: the out-graph/in-graph init
contract above makes it a safe upper bound for any backend honoring
the contract. Override for audited deviations.
"""
if forward_mode.is_decode() or forward_mode.is_target_verify():
return SharedReadBoundary.IN_REPLAY
return SharedReadBoundary.UNKNOWN
# Chunked-prefix FullCG capture has a second model topology and stable # Chunked-prefix FullCG capture has a second model topology and stable
# prefix buffers. Backends must opt in explicitly so the runner does not # prefix buffers. Backends must opt in explicitly so the runner does not
@@ -39,7 +39,10 @@ from sglang.kernels.ops.speculative.dspark.dspark_attn_metadata import (
ComputeDsparkWindowGather, ComputeDsparkWindowGather,
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import (
AttentionBackend,
SharedReadBoundary,
)
from sglang.srt.layers.attention.dsa.dsa_topk_backend import DSATopKBackend from sglang.srt.layers.attention.dsa.dsa_topk_backend import DSATopKBackend
from sglang.srt.layers.attention.dsv4.compressor_v2 import ( from sglang.srt.layers.attention.dsv4.compressor_v2 import (
CompressorBackendMixin, CompressorBackendMixin,
@@ -501,6 +504,12 @@ class DeepseekV4AttnBackend(
supports_ragged_verify_graph: bool = True supports_ragged_verify_graph: bool = True
needs_cpu_seq_lens: bool = False needs_cpu_seq_lens: bool = False
def shared_read_boundary(self, forward_mode: ForwardMode) -> SharedReadBoundary:
# Breakable-graph verify rereads shared state across segments.
if forward_mode.is_target_verify():
return SharedReadBoundary.POST_REPLAY
return super().shared_read_boundary(forward_mode)
def __init__( def __init__(
self, self,
model_runner: ModelRunner, model_runner: ModelRunner,
@@ -21,6 +21,7 @@ from sglang.kernels.ops.kvcache.trtllm_mha_page_table import (
build_trtllm_mha_page_table, build_trtllm_mha_page_table,
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.attention.base_attn_backend import SharedReadBoundary
from sglang.srt.layers.attention.flashinfer_backend import ( from sglang.srt.layers.attention.flashinfer_backend import (
FlashInferAttnBackend, FlashInferAttnBackend,
FlashInferMultiStepDraftBackend, FlashInferMultiStepDraftBackend,
@@ -102,8 +103,11 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
supports_ragged_verify_graph: bool = True supports_ragged_verify_graph: bool = True
# Prefill metadata initialization snapshots all scheduler-shared inputs. def shared_read_boundary(self, forward_mode: ForwardMode) -> SharedReadBoundary:
prefill_shared_reads_end_at_metadata_init: bool = True # Prefill metadata init snapshots all scheduler-shared inputs pre-replay.
if forward_mode == ForwardMode.EXTEND:
return SharedReadBoundary.PRE_REPLAY
return super().shared_read_boundary(forward_mode)
def __init__( def __init__(
self, self,
@@ -44,6 +44,7 @@ from sglang.srt.distributed.parallel_state import (
) )
from sglang.srt.dllm.config import DllmConfig from sglang.srt.dllm.config import DllmConfig
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.attention.base_attn_backend import SharedReadBoundary
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
DpPaddingMode, DpPaddingMode,
@@ -82,7 +83,6 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen
from sglang.srt.model_executor.runner_backend_utils import ( from sglang.srt.model_executor.runner_backend_utils import (
CUDA_GRAPH_CAPTURE_FAILED_MSG, CUDA_GRAPH_CAPTURE_FAILED_MSG,
) )
from sglang.srt.model_executor.runner_utils import WarReadDonePolicy
from sglang.srt.model_executor.runner_utils.buffers import ( from sglang.srt.model_executor.runner_utils.buffers import (
DecodeInputBuffers, DecodeInputBuffers,
) )
@@ -438,23 +438,22 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
self.model_runner.war_read_done_event.record() self.model_runner.war_read_done_event.record()
self._war_read_done_node_planted = True self._war_read_done_node_planted = True
def _war_read_done_policy(self, attn_backend, forward_mode) -> WarReadDonePolicy: def _war_read_done_record(self, attn_backend, forward_mode) -> SharedReadBoundary:
"""Whether and where this replay records its WAR read-done event.""" """Where this replay records its WAR read-done event; UNKNOWN records
nothing."""
if forward_mode.is_target_verify(): if forward_mode.is_target_verify():
if not self.model_runner.spec_algorithm.is_war_publish_phase(forward_mode):
return SharedReadBoundary.UNKNOWN
elif not forward_mode.is_decode():
return SharedReadBoundary.UNKNOWN
boundary = attn_backend.shared_read_boundary(forward_mode)
if ( if (
not self.model_runner.spec_algorithm.supports_target_verify_war_read_done() boundary is SharedReadBoundary.IN_REPLAY
and not self._war_read_done_node_planted
): ):
return WarReadDonePolicy.NONE # Non-capturing runs / no external-event support.
if attn_backend.use_captured_forward_metadata_for_breakable_cuda_graph: return SharedReadBoundary.PRE_REPLAY
return WarReadDonePolicy.POST_REPLAY return boundary
if self._war_read_done_node_planted:
return WarReadDonePolicy.IN_GRAPH
return WarReadDonePolicy.PRE_REPLAY
elif forward_mode.is_decode():
if self._war_read_done_node_planted:
return WarReadDonePolicy.IN_GRAPH
return WarReadDonePolicy.PRE_REPLAY
return WarReadDonePolicy.NONE
def _publish_war_read_done(self, in_graph: bool): def _publish_war_read_done(self, in_graph: bool):
"""Publish the read-done event the scheduler's WAR barrier waits on.""" """Publish the read-done event the scheduler's WAR barrier waits on."""
@@ -1305,7 +1304,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
timer_ctx = device_timer_ctx( timer_ctx = device_timer_ctx(
self.model_runner.device_timer, forward_batch.forward_mode.name.lower() self.model_runner.device_timer, forward_batch.forward_mode.name.lower()
) )
war_policy = self._war_read_done_policy( war_record = self._war_read_done_record(
self.attn_backend, forward_batch.forward_mode self.attn_backend, forward_batch.forward_mode
) )
with timer_ctx, self.backend.replay_session(): with timer_ctx, self.backend.replay_session():
@@ -1324,12 +1323,12 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
else "" else ""
), ),
) )
if war_policy is WarReadDonePolicy.PRE_REPLAY: if war_record is SharedReadBoundary.PRE_REPLAY:
self._publish_war_read_done(in_graph=False) self._publish_war_read_done(in_graph=False)
output = self.backend.replay(self._replay_graph_key, forward_batch) output = self.backend.replay(self._replay_graph_key, forward_batch)
if war_policy is WarReadDonePolicy.POST_REPLAY: if war_record is SharedReadBoundary.POST_REPLAY:
self._publish_war_read_done(in_graph=False) self._publish_war_read_done(in_graph=False)
elif war_policy is WarReadDonePolicy.IN_GRAPH: elif war_record is SharedReadBoundary.IN_REPLAY:
self._publish_war_read_done(in_graph=True) self._publish_war_read_done(in_graph=True)
if isinstance(output, LogitsProcessorOutput): if isinstance(output, LogitsProcessorOutput):
@@ -27,7 +27,6 @@ from sglang.srt.model_executor.runner_utils.pool import ( # noqa: F401
set_global_graph_memory_pool, set_global_graph_memory_pool,
) )
from sglang.srt.model_executor.runner_utils.war_event import ( # noqa: F401 from sglang.srt.model_executor.runner_utils.war_event import ( # noqa: F401
WarReadDonePolicy,
make_war_read_done_event, make_war_read_done_event,
maybe_publish_prefill_war_read_done, maybe_publish_prefill_war_read_done,
) )
@@ -1,30 +1,18 @@
"""WAR read-done event utilities for CUDA graph runners.""" """WAR read-done event utilities for CUDA graph runners."""
import logging import logging
from enum import Enum, auto
from typing import Optional from typing import Optional
import torch import torch
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.attention.base_attn_backend import SharedReadBoundary
from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.utils import is_cuda from sglang.srt.utils import is_cuda
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# Whether and where the WAR read-done record lands for a replay.
class WarReadDonePolicy(Enum):
# This forward mode or algorithm does not publish a read-done event.
NONE = auto()
# Snapshot backends finish all shared reads before launch.
PRE_REPLAY = auto()
# Captured metadata initialization finishes all shared reads.
IN_GRAPH = auto()
# Captured-metadata replays keep reading shared buffers throughout the graph.
POST_REPLAY = auto()
def make_war_read_done_event(device_module) -> Optional[torch.cuda.Event]: def make_war_read_done_event(device_module) -> Optional[torch.cuda.Event]:
"""Create a persistent external event for CUDA graph capture.""" """Create a persistent external event for CUDA graph capture."""
if not is_cuda(): if not is_cuda():
@@ -47,7 +35,11 @@ def maybe_publish_prefill_war_read_done(
# WAR boundaries are validated. # WAR boundaries are validated.
if not model_runner.spec_algorithm.is_none(): if not model_runner.spec_algorithm.is_none():
return return
if not model_runner.attn_backend.prefill_shared_reads_end_at_metadata_init: # The record lands right after replay prep, so PRE_REPLAY only.
boundary = model_runner.attn_backend.shared_read_boundary(
forward_batch.forward_mode
)
if boundary is not SharedReadBoundary.PRE_REPLAY:
return return
logger.info_once( logger.info_once(
"Prefill WAR read-done fastpath active (%s)", "Prefill WAR read-done fastpath active (%s)",
@@ -590,7 +590,8 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
self.draft_extend_attn_backend.init_forward_metadata_out_graph(fb_view) self.draft_extend_attn_backend.init_forward_metadata_out_graph(fb_view)
# Snapshot built -- the forward is done reading the shared pool. Publish # Snapshot built -- the forward is done reading the shared pool. Publish
# a read-done event the scheduler's WAR barrier waits on. # a read-done event the scheduler's WAR barrier waits on (draft extend
# is the EAGLE-family war-publish phase; last write wins the mailbox).
read_done = self.device_module.Event() read_done = self.device_module.Event()
read_done.record() read_done.record()
self.model_runner.war_fastpath_read_done_event = read_done self.model_runner.war_fastpath_read_done_event = read_done
+5 -2
View File
@@ -127,8 +127,11 @@ class SpeculativeAlgorithm(Enum):
def supports_target_verify_for_draft(self) -> bool: def supports_target_verify_for_draft(self) -> bool:
return self.is_dflash_family() return self.is_dflash_family()
def supports_target_verify_war_read_done(self) -> bool: def is_war_publish_phase(self, forward_mode) -> bool:
return self.is_dflash_family() # The step's last shared-buffer-reading phase owns the WAR read-done publish.
if self.is_dflash_family():
return forward_mode.is_target_verify()
return forward_mode.is_draft_extend_v2()
def supports_ragged_verify(self) -> bool: def supports_ragged_verify(self) -> bool:
"""Whether this algorithm's verify step may carry a RaggedVerifyLayout """Whether this algorithm's verify step may carry a RaggedVerifyLayout
@@ -92,8 +92,9 @@ class CustomSpecAlgo:
def supports_target_verify_for_draft(self) -> bool: def supports_target_verify_for_draft(self) -> bool:
return False return False
def supports_target_verify_war_read_done(self) -> bool: def is_war_publish_phase(self, forward_mode) -> bool:
return False # The step's last shared-buffer-reading phase owns the WAR read-done publish.
return forward_mode.is_draft_extend_v2()
def supports_ragged_verify(self) -> bool: def supports_ragged_verify(self) -> bool:
return False return False
@@ -4,12 +4,12 @@ from types import SimpleNamespace
import pytest import pytest
import torch import torch
from sglang.srt.layers.attention.base_attn_backend import SharedReadBoundary
from sglang.srt.model_executor.forward_batch_info import ForwardMode, PPProxyTensors from sglang.srt.model_executor.forward_batch_info import ForwardMode, PPProxyTensors
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import ( from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
DecodeCudaGraphRunner, DecodeCudaGraphRunner,
) )
from sglang.srt.model_executor.runner.shape_key import ShapeKey from sglang.srt.model_executor.runner.shape_key import ShapeKey
from sglang.srt.model_executor.runner_utils import WarReadDonePolicy
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=1, suite="base-a-test-cpu") register_cpu_ci(est_time=1, suite="base-a-test-cpu")
@@ -19,14 +19,19 @@ class _SpecAlgorithm:
def __init__(self, target_verify_war: bool = False): def __init__(self, target_verify_war: bool = False):
self._target_verify_war = target_verify_war self._target_verify_war = target_verify_war
def supports_target_verify_war_read_done(self) -> bool: def is_war_publish_phase(self, forward_mode) -> bool:
return self._target_verify_war return self._target_verify_war and forward_mode.is_target_verify()
def _attn_backend(*, breakable_metadata=False): def _attn_backend(*, breakable_metadata=False):
return SimpleNamespace( def shared_read_boundary(forward_mode):
use_captured_forward_metadata_for_breakable_cuda_graph=breakable_metadata, if breakable_metadata and forward_mode.is_target_verify():
) return SharedReadBoundary.POST_REPLAY
if forward_mode.is_decode() or forward_mode.is_target_verify():
return SharedReadBoundary.IN_REPLAY
return SharedReadBoundary.UNKNOWN
return SimpleNamespace(shared_read_boundary=shared_read_boundary)
def _runner(*, target_verify_war: bool = False, planted: bool = False): def _runner(*, target_verify_war: bool = False, planted: bool = False):
@@ -42,35 +47,35 @@ def _runner(*, target_verify_war: bool = False, planted: bool = False):
return runner return runner
def test_war_read_done_policy(): def test_war_read_done_record():
# Planted node: the graph re-arms it every replay. # Planted node: the graph re-arms it every replay.
assert ( assert (
_runner(planted=True)._war_read_done_policy(_attn_backend(), ForwardMode.DECODE) _runner(planted=True)._war_read_done_record(_attn_backend(), ForwardMode.DECODE)
is WarReadDonePolicy.IN_GRAPH is SharedReadBoundary.IN_REPLAY
) )
# No node, snapshot backend: all shared reads finish before launch. # No planted node: fall back to a pre-replay record.
assert ( assert (
_runner()._war_read_done_policy(_attn_backend(), ForwardMode.DECODE) _runner()._war_read_done_record(_attn_backend(), ForwardMode.DECODE)
is WarReadDonePolicy.PRE_REPLAY is SharedReadBoundary.PRE_REPLAY
) )
# Unrelated modes never publish from the decode graph runner. # Unrelated modes never publish from the decode graph runner.
assert ( assert (
_runner(planted=True)._war_read_done_policy(_attn_backend(), ForwardMode.EXTEND) _runner(planted=True)._war_read_done_record(_attn_backend(), ForwardMode.EXTEND)
is WarReadDonePolicy.NONE is SharedReadBoundary.UNKNOWN
) )
# Backend placement cannot opt an unsupported algorithm into publication. # The algorithm gate precedes the backend declaration.
assert ( assert (
_runner(planted=True)._war_read_done_policy( _runner(planted=True)._war_read_done_record(
_attn_backend(breakable_metadata=True), ForwardMode.TARGET_VERIFY _attn_backend(breakable_metadata=True), ForwardMode.TARGET_VERIFY
) )
is WarReadDonePolicy.NONE is SharedReadBoundary.UNKNOWN
) )
# Captured-metadata verify keeps reading throughout the graph, even planted. # Captured-metadata verify keeps reading throughout the graph, even planted.
assert ( assert (
_runner(target_verify_war=True, planted=True)._war_read_done_policy( _runner(target_verify_war=True, planted=True)._war_read_done_record(
_attn_backend(breakable_metadata=True), ForwardMode.TARGET_VERIFY _attn_backend(breakable_metadata=True), ForwardMode.TARGET_VERIFY
) )
is WarReadDonePolicy.POST_REPLAY is SharedReadBoundary.POST_REPLAY
) )
@@ -3,6 +3,7 @@ from types import SimpleNamespace
import pytest import pytest
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.attention.base_attn_backend import SharedReadBoundary
from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.model_executor.runner_utils import maybe_publish_prefill_war_read_done from sglang.srt.model_executor.runner_utils import maybe_publish_prefill_war_read_done
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
@@ -20,11 +21,12 @@ class _Event:
def _model_runner(*, spec_algorithm=SpeculativeAlgorithm.NONE, compliant=True): def _model_runner(*, spec_algorithm=SpeculativeAlgorithm.NONE, compliant=True):
boundary = (
SharedReadBoundary.PRE_REPLAY if compliant else SharedReadBoundary.UNKNOWN
)
return SimpleNamespace( return SimpleNamespace(
spec_algorithm=spec_algorithm, spec_algorithm=spec_algorithm,
attn_backend=SimpleNamespace( attn_backend=SimpleNamespace(shared_read_boundary=lambda mode: boundary),
prefill_shared_reads_end_at_metadata_init=compliant
),
war_fastpath_read_done_event=None, war_fastpath_read_done_event=None,
) )
@@ -60,7 +62,7 @@ def test_gates_exclude_non_prefill_unsupported_algorithm_and_noncompliant_backen
(_model_runner(), _batch(ForwardMode.DECODE)), (_model_runner(), _batch(ForwardMode.DECODE)),
# The algorithm has a later prefill reader or unverified ownership. # The algorithm has a later prefill reader or unverified ownership.
(_model_runner(spec_algorithm=SpeculativeAlgorithm.EAGLE), _batch()), (_model_runner(spec_algorithm=SpeculativeAlgorithm.EAGLE), _batch()),
# Backend has not declared metadata-init compliance. # Backend has not declared a pre-replay prefill read boundary.
(_model_runner(compliant=False), _batch()), (_model_runner(compliant=False), _batch()),
): ):
maybe_publish_prefill_war_read_done(runner, batch, _DEVICE_MODULE) maybe_publish_prefill_war_read_done(runner, batch, _DEVICE_MODULE)