[Spec] Resolve shared-read ends from the backend declaration alone (#35059)
This commit is contained in:
@@ -1701,12 +1701,8 @@ class Scheduler(
|
|||||||
dispatch_event_loop(self)
|
dispatch_event_loop(self)
|
||||||
|
|
||||||
def _apply_war_barrier(self):
|
def _apply_war_barrier(self):
|
||||||
# Called right after each launch: order later schedule_stream work
|
# WAR: keep later schedule_stream writes behind this forward's shared reads.
|
||||||
# (result processing, next iteration's writes) behind the forward's
|
# Clearing matters: a phase that skips the publish then falls back to coarse.
|
||||||
# shared-buffer reads. Fast path: wait on the read-done event the
|
|
||||||
# forward published after its snapshot (non-spec: decode graph; spec:
|
|
||||||
# draft_extend), then clear it. Else whole-forward wait_stream
|
|
||||||
# (forceable via SGLANG_FORCE_COARSE_WAR_BARRIER).
|
|
||||||
if not self._war_barrier_enabled:
|
if not self._war_barrier_enabled:
|
||||||
return
|
return
|
||||||
runner = self.model_worker.last_shared_read_runner
|
runner = self.model_worker.last_shared_read_runner
|
||||||
|
|||||||
@@ -403,9 +403,8 @@ class ModelRunner:
|
|||||||
# Init forward stream for overlap schedule
|
# Init forward stream for overlap schedule
|
||||||
self.forward_stream = torch.get_device_module(self.device).Stream()
|
self.forward_stream = torch.get_device_module(self.device).Stream()
|
||||||
|
|
||||||
# Published by the step's last shared-buffer-reading phase (decode graph,
|
# Read-done mailbox: the scheduler's WAR barrier reads it from the runner
|
||||||
# eagle draft extend, or prefill); the scheduler's WAR barrier waits on it
|
# its worker names, and treats None as the coarse whole-forward fence.
|
||||||
# then clears it. None -> coarse whole-forward wait_stream.
|
|
||||||
self.shared_read_done_event: Optional[torch.cuda.Event] = None
|
self.shared_read_done_event: Optional[torch.cuda.Event] = None
|
||||||
|
|
||||||
# CPU offload
|
# CPU offload
|
||||||
|
|||||||
@@ -486,17 +486,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
return self.attn_backend
|
return self.attn_backend
|
||||||
|
|
||||||
def _resolve_shared_read_ends(self, attn_backend, forward_mode) -> SharedReadEnds:
|
def _resolve_shared_read_ends(self, attn_backend, forward_mode) -> SharedReadEnds:
|
||||||
"""The backend's declaration, demoted when this runner cannot record
|
|
||||||
there. UNKNOWN records nothing (scheduler keeps the coarse fence)."""
|
|
||||||
if forward_mode.is_target_verify():
|
|
||||||
if not self.model_runner.spec_algorithm.is_last_shared_read_phase(
|
|
||||||
forward_mode
|
|
||||||
):
|
|
||||||
return SharedReadEnds.UNKNOWN
|
|
||||||
elif not forward_mode.is_decode():
|
|
||||||
return SharedReadEnds.UNKNOWN
|
|
||||||
declared = attn_backend.shared_read_ends(forward_mode)
|
declared = attn_backend.shared_read_ends(forward_mode)
|
||||||
|
|
||||||
if (
|
if (
|
||||||
declared is SharedReadEnds.IN_REPLAY
|
declared is SharedReadEnds.IN_REPLAY
|
||||||
and self.in_graph_metadata_prep_done is None
|
and self.in_graph_metadata_prep_done is None
|
||||||
|
|||||||
@@ -419,6 +419,8 @@ class NGRAMWorker(BaseSpecWorker):
|
|||||||
batch_result = self.target_worker.forward_batch_generation(
|
batch_result = self.target_worker.forward_batch_generation(
|
||||||
batch, is_verify=True
|
batch, is_verify=True
|
||||||
)
|
)
|
||||||
|
# Verify reads shared state past the in-graph marker; keep it coarse.
|
||||||
|
self.target_worker.model_runner.shared_read_done_event = None
|
||||||
|
|
||||||
logits_output, can_run_cuda_graph = (
|
logits_output, can_run_cuda_graph = (
|
||||||
batch_result.logits_output,
|
batch_result.logits_output,
|
||||||
|
|||||||
@@ -127,14 +127,6 @@ 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 is_last_shared_read_phase(self, forward_mode) -> bool:
|
|
||||||
# The step's last shared-buffer-reading phase owns the shared-read-done
|
|
||||||
# publish. DSPARK has no draft_extend: its draft samples inside the
|
|
||||||
# verify graph, so verify is that last phase.
|
|
||||||
if self.is_dflash_family() or self.is_dspark():
|
|
||||||
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
|
||||||
(per-request verify lengths); gates the token-bucket-keyed verify
|
(per-request verify lengths); gates the token-bucket-keyed verify
|
||||||
|
|||||||
@@ -92,10 +92,6 @@ class CustomSpecAlgo:
|
|||||||
def supports_target_verify_for_draft(self) -> bool:
|
def supports_target_verify_for_draft(self) -> bool:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
def is_last_shared_read_phase(self, forward_mode) -> bool:
|
|
||||||
# The step's last shared-buffer-reading phase owns the shared-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
|
||||||
|
|
||||||
|
|||||||
+14
-23
@@ -16,18 +16,11 @@ 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")
|
||||||
|
|
||||||
DECODE = ForwardMode.DECODE
|
DECODE = ForwardMode.DECODE
|
||||||
VERIFY = ForwardMode.TARGET_VERIFY
|
|
||||||
EXTEND = ForwardMode.EXTEND
|
|
||||||
|
|
||||||
|
|
||||||
def _runner(*, owns_verify: bool = False, has_marker: bool = False):
|
def _runner(*, has_marker: bool = False):
|
||||||
runner = DecodeCudaGraphRunner.__new__(DecodeCudaGraphRunner)
|
runner = DecodeCudaGraphRunner.__new__(DecodeCudaGraphRunner)
|
||||||
runner.model_runner = SimpleNamespace(
|
runner.model_runner = SimpleNamespace(shared_read_done_event=None)
|
||||||
spec_algorithm=SimpleNamespace(
|
|
||||||
is_last_shared_read_phase=lambda fm: owns_verify and fm.is_target_verify()
|
|
||||||
),
|
|
||||||
shared_read_done_event=None,
|
|
||||||
)
|
|
||||||
runner.in_graph_metadata_prep_done = object() if has_marker else None
|
runner.in_graph_metadata_prep_done = object() if has_marker else None
|
||||||
return runner
|
return runner
|
||||||
|
|
||||||
@@ -40,24 +33,22 @@ def _backend(declared: SharedReadEnds):
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
"mode, owns_verify, declared, has_marker, expected",
|
"declared, has_marker, expected",
|
||||||
[
|
[
|
||||||
# Only decode / target verify publish; anything else keeps the coarse fence.
|
# The backend's declaration decides where the record lands.
|
||||||
(EXTEND, False, SharedReadEnds.IN_REPLAY, True, SharedReadEnds.UNKNOWN),
|
(SharedReadEnds.IN_REPLAY, True, SharedReadEnds.IN_REPLAY),
|
||||||
# Target verify publishes only when it is the step's last reading phase.
|
|
||||||
(VERIFY, False, SharedReadEnds.IN_REPLAY, True, SharedReadEnds.UNKNOWN),
|
|
||||||
(VERIFY, True, SharedReadEnds.IN_REPLAY, True, SharedReadEnds.IN_REPLAY),
|
|
||||||
# A backend that keeps reading through the graph is never advanced.
|
|
||||||
(VERIFY, True, SharedReadEnds.POST_REPLAY, True, SharedReadEnds.POST_REPLAY),
|
|
||||||
# Nothing to demote: the declaration is honored as-is.
|
|
||||||
(DECODE, False, SharedReadEnds.IN_REPLAY, True, SharedReadEnds.IN_REPLAY),
|
|
||||||
# Nowhere to record in-graph -> fall back to the pre-replay record.
|
# Nowhere to record in-graph -> fall back to the pre-replay record.
|
||||||
(DECODE, False, SharedReadEnds.IN_REPLAY, False, SharedReadEnds.PRE_REPLAY),
|
(SharedReadEnds.IN_REPLAY, False, SharedReadEnds.PRE_REPLAY),
|
||||||
|
# Only an in-graph declaration is demoted; the rest pass through.
|
||||||
|
(SharedReadEnds.POST_REPLAY, False, SharedReadEnds.POST_REPLAY),
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
def test_resolve_shared_read_ends(mode, owns_verify, declared, has_marker, expected):
|
def test_resolve_shared_read_ends(declared, has_marker, expected):
|
||||||
runner = _runner(owns_verify=owns_verify, has_marker=has_marker)
|
runner = _runner(has_marker=has_marker)
|
||||||
assert runner._resolve_shared_read_ends(_backend(declared), mode) is expected
|
backend = _backend(declared)
|
||||||
|
|
||||||
|
assert runner._resolve_shared_read_ends(backend, DECODE) is expected
|
||||||
|
backend.shared_read_ends.assert_called_once_with(DECODE)
|
||||||
|
|
||||||
|
|
||||||
def test_publish_read_done():
|
def test_publish_read_done():
|
||||||
|
|||||||
Reference in New Issue
Block a user