[Scheduler] Align WAR fences with CUDA graph metadata reads (#33587)
This commit is contained in:
@@ -207,6 +207,121 @@ def test_metadata_update_records_inside_cuda_graph():
|
||||
)
|
||||
|
||||
|
||||
def test_graph_read_done_event_fences_slot_mutation():
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA required")
|
||||
|
||||
backend = _make_backend_for_hook_test()
|
||||
backend.device = torch.device(DEVICE)
|
||||
backend.page_size = 2
|
||||
backend.max_num_pages = 2
|
||||
backend.use_sliding_window_kv_pool = True
|
||||
backend._swa_kv_pool = object()
|
||||
backend.req_to_token = torch.tensor(
|
||||
[[0, 1, 2, 3], [8, 9, 10, 11]], dtype=torch.int32, device=DEVICE
|
||||
)
|
||||
backend._swa_full_to_swa_mapping = (
|
||||
torch.arange(32, dtype=torch.int64, device=DEVICE) * 2
|
||||
)
|
||||
backend.init_cuda_graph_state(max_bs=1, max_num_tokens=1)
|
||||
forward_batch = SimpleNamespace(
|
||||
batch_size=1,
|
||||
req_pool_indices=torch.tensor([0], dtype=torch.int64, device=DEVICE),
|
||||
seq_lens=torch.tensor([4], dtype=torch.int32, device=DEVICE),
|
||||
forward_mode=ForwardMode.DECODE,
|
||||
spec_info=None,
|
||||
positions=torch.tensor([0], dtype=torch.int64, device=DEVICE),
|
||||
out_cache_loc=torch.tensor([3], dtype=torch.int64, device=DEVICE),
|
||||
)
|
||||
|
||||
backend.init_forward_metadata_out_graph(forward_batch, in_capture=True)
|
||||
backend.init_forward_metadata_in_graph(forward_batch)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
read_done = torch.cuda.Event(external=True)
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
backend.init_forward_metadata_in_graph(forward_batch)
|
||||
read_done.record()
|
||||
|
||||
fence_stream = torch.cuda.Stream()
|
||||
mutation_done = torch.cuda.Event()
|
||||
graph.replay()
|
||||
with torch.cuda.stream(fence_stream):
|
||||
fence_stream.wait_event(read_done)
|
||||
backend.req_to_token.copy_(
|
||||
torch.tensor(
|
||||
[[8, 9, 10, 11], [0, 1, 2, 3]],
|
||||
dtype=torch.int32,
|
||||
device=DEVICE,
|
||||
)
|
||||
)
|
||||
backend._swa_full_to_swa_mapping.add_(64)
|
||||
mutation_done.record()
|
||||
torch.cuda.current_stream().wait_event(mutation_done)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
torch.testing.assert_close(
|
||||
backend.forward_metadata.page_table,
|
||||
torch.tensor([[0, 1]], dtype=torch.int32, device=DEVICE),
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
backend.forward_metadata.swa_page_table,
|
||||
torch.tensor([[0, 2]], dtype=torch.int32, device=DEVICE),
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
backend.forward_metadata.swa_out_cache_loc,
|
||||
torch.tensor([6], dtype=torch.int64, device=DEVICE),
|
||||
)
|
||||
|
||||
graph.replay()
|
||||
with torch.cuda.stream(fence_stream):
|
||||
fence_stream.wait_event(read_done)
|
||||
backend.req_to_token.copy_(
|
||||
torch.tensor(
|
||||
[[0, 1, 2, 3], [8, 9, 10, 11]],
|
||||
dtype=torch.int32,
|
||||
device=DEVICE,
|
||||
)
|
||||
)
|
||||
backend._swa_full_to_swa_mapping.sub_(64)
|
||||
mutation_done.record()
|
||||
torch.cuda.current_stream().wait_event(mutation_done)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
torch.testing.assert_close(
|
||||
backend.forward_metadata.page_table,
|
||||
torch.tensor([[4, 5]], dtype=torch.int32, device=DEVICE),
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
backend.forward_metadata.swa_page_table,
|
||||
torch.tensor([[40, 42]], dtype=torch.int32, device=DEVICE),
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
backend.forward_metadata.swa_out_cache_loc,
|
||||
torch.tensor([70], dtype=torch.int64, device=DEVICE),
|
||||
)
|
||||
|
||||
|
||||
def test_swa_cache_write_uses_metadata_slot_snapshot():
|
||||
snapshot = torch.tensor([6], dtype=torch.int64)
|
||||
|
||||
def translate_live_mapping(_):
|
||||
raise AssertionError("cache writes must not read the live SWA mapping")
|
||||
|
||||
backend = TRTLLMHAAttnBackend.__new__(TRTLLMHAAttnBackend)
|
||||
backend._swa_kv_pool = SimpleNamespace(
|
||||
layers_mapping={1: (0, True)},
|
||||
translate_loc_from_full_to_swa=translate_live_mapping,
|
||||
)
|
||||
backend.forward_metadata = SimpleNamespace(swa_out_cache_loc=snapshot)
|
||||
forward_batch = SimpleNamespace(out_cache_loc=torch.tensor([3], dtype=torch.int64))
|
||||
|
||||
cache_loc = backend._get_layer_cache_loc(SimpleNamespace(layer_id=1), forward_batch)
|
||||
|
||||
torch.testing.assert_close(cache_loc, snapshot, rtol=0, atol=0)
|
||||
|
||||
|
||||
def _build_inputs(bs, pool_size, max_num_pages, max_seq_pages, seq_max, seed):
|
||||
"""Build random pool / indices / seq_lens consistent with backend buffers."""
|
||||
g = torch.Generator(device="cpu").manual_seed(seed)
|
||||
|
||||
@@ -0,0 +1,169 @@
|
||||
import contextlib
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode, PPProxyTensors
|
||||
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
|
||||
DecodeCudaGraphRunner,
|
||||
)
|
||||
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
|
||||
|
||||
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
class _SpecAlgorithm:
|
||||
def __init__(self, target_verify_war: bool = False):
|
||||
self._target_verify_war = target_verify_war
|
||||
|
||||
def supports_target_verify_war_read_done(self) -> bool:
|
||||
return self._target_verify_war
|
||||
|
||||
|
||||
def _attn_backend(*, breakable_metadata=False):
|
||||
return SimpleNamespace(
|
||||
use_captured_forward_metadata_for_breakable_cuda_graph=breakable_metadata,
|
||||
)
|
||||
|
||||
|
||||
def _runner(*, target_verify_war: bool = False, planted: bool = False):
|
||||
runner = DecodeCudaGraphRunner.__new__(DecodeCudaGraphRunner)
|
||||
runner.model_runner = SimpleNamespace(
|
||||
spec_algorithm=_SpecAlgorithm(target_verify_war),
|
||||
device_timer=None,
|
||||
is_draft_worker=False,
|
||||
war_read_done_event=None,
|
||||
war_fastpath_read_done_event=None,
|
||||
)
|
||||
runner._war_read_done_node_planted = planted
|
||||
return runner
|
||||
|
||||
|
||||
def test_war_read_done_policy():
|
||||
# Planted node: the graph re-arms it every replay.
|
||||
assert (
|
||||
_runner(planted=True)._war_read_done_policy(_attn_backend(), ForwardMode.DECODE)
|
||||
is WarReadDonePolicy.IN_GRAPH
|
||||
)
|
||||
# No node, snapshot backend: all shared reads finish before launch.
|
||||
assert (
|
||||
_runner()._war_read_done_policy(_attn_backend(), ForwardMode.DECODE)
|
||||
is WarReadDonePolicy.PRE_REPLAY
|
||||
)
|
||||
# Unrelated modes never publish from the decode graph runner.
|
||||
assert (
|
||||
_runner(planted=True)._war_read_done_policy(_attn_backend(), ForwardMode.EXTEND)
|
||||
is WarReadDonePolicy.NONE
|
||||
)
|
||||
# Backend placement cannot opt an unsupported algorithm into publication.
|
||||
assert (
|
||||
_runner(planted=True)._war_read_done_policy(
|
||||
_attn_backend(breakable_metadata=True), ForwardMode.TARGET_VERIFY
|
||||
)
|
||||
is WarReadDonePolicy.NONE
|
||||
)
|
||||
# Captured-metadata verify keeps reading throughout the graph, even planted.
|
||||
assert (
|
||||
_runner(target_verify_war=True, planted=True)._war_read_done_policy(
|
||||
_attn_backend(breakable_metadata=True), ForwardMode.TARGET_VERIFY
|
||||
)
|
||||
is WarReadDonePolicy.POST_REPLAY
|
||||
)
|
||||
|
||||
|
||||
def test_publish_war_read_done():
|
||||
runner = _runner()
|
||||
graph_event = object()
|
||||
runner.model_runner.war_read_done_event = graph_event
|
||||
runner._publish_war_read_done(in_graph=True)
|
||||
assert runner.model_runner.war_fastpath_read_done_event is graph_event
|
||||
|
||||
recorded = []
|
||||
|
||||
class Event:
|
||||
def record(self):
|
||||
recorded.append(self)
|
||||
|
||||
runner.device_module = SimpleNamespace(Event=Event)
|
||||
runner._publish_war_read_done(in_graph=False)
|
||||
published = runner.model_runner.war_fastpath_read_done_event
|
||||
assert isinstance(published, Event) and recorded == [published]
|
||||
|
||||
|
||||
def _execute_harness(runner, calls, mode=ForwardMode.DECODE):
|
||||
key = ShapeKey(size=1)
|
||||
output = PPProxyTensors({"hidden_states": torch.ones(1, 1)})
|
||||
runner.ragged_verify_mode = False
|
||||
runner.bs = 1
|
||||
runner.load_batch = lambda *_: setattr(runner, "_replay_graph_key", key)
|
||||
|
||||
class Backend:
|
||||
def replay_session(self):
|
||||
return contextlib.nullcontext()
|
||||
|
||||
def replay(self, replay_key, _forward_batch):
|
||||
assert replay_key == key
|
||||
calls.append("replay")
|
||||
return output
|
||||
|
||||
runner.backend = Backend()
|
||||
return SimpleNamespace(forward_mode=mode, batch_size=1)
|
||||
|
||||
|
||||
def test_execute_publishes_the_planted_graph_event():
|
||||
runner = _runner(planted=True)
|
||||
graph_event = object()
|
||||
runner.model_runner.war_read_done_event = graph_event
|
||||
runner.attn_backend = _attn_backend()
|
||||
runner.device_module = SimpleNamespace(
|
||||
Event=lambda: (_ for _ in ()).throw(
|
||||
AssertionError("execute must reuse the graph-recorded event")
|
||||
)
|
||||
)
|
||||
calls = []
|
||||
forward_batch = _execute_harness(runner, calls)
|
||||
|
||||
result = runner.execute(forward_batch)
|
||||
|
||||
assert result.tensors["hidden_states"].shape == (1, 1)
|
||||
assert runner.model_runner.war_fastpath_read_done_event is graph_event
|
||||
|
||||
|
||||
def test_execute_records_pre_replay_for_snapshot_backends():
|
||||
runner = _runner()
|
||||
runner.attn_backend = _attn_backend()
|
||||
calls = []
|
||||
|
||||
class Event:
|
||||
def record(self):
|
||||
calls.append("record")
|
||||
|
||||
runner.device_module = SimpleNamespace(Event=Event)
|
||||
forward_batch = _execute_harness(runner, calls)
|
||||
|
||||
runner.execute(forward_batch)
|
||||
|
||||
# The eager record lands before the replay so the fence stays truthful.
|
||||
assert calls == ["record", "replay"]
|
||||
assert isinstance(runner.model_runner.war_fastpath_read_done_event, Event)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("supported", [False, True])
|
||||
def test_target_verify_requires_war_capability(supported):
|
||||
runner = _runner(target_verify_war=supported, planted=True)
|
||||
graph_event = object()
|
||||
runner.model_runner.war_read_done_event = graph_event
|
||||
runner.attn_backend = _attn_backend()
|
||||
runner.device_module = SimpleNamespace(Event=lambda: None)
|
||||
|
||||
runner.execute(_execute_harness(runner, [], ForwardMode.TARGET_VERIFY))
|
||||
|
||||
expected = graph_event if supported else None
|
||||
assert runner.model_runner.war_fastpath_read_done_event is expected
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(pytest.main([__file__, "-v"]))
|
||||
@@ -0,0 +1,71 @@
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
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.speculative.spec_info import SpeculativeAlgorithm
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
class _Event:
|
||||
def __init__(self):
|
||||
self.recorded = False
|
||||
|
||||
def record(self):
|
||||
self.recorded = True
|
||||
|
||||
|
||||
def _model_runner(*, spec_algorithm=SpeculativeAlgorithm.NONE, compliant=True):
|
||||
return SimpleNamespace(
|
||||
spec_algorithm=spec_algorithm,
|
||||
attn_backend=SimpleNamespace(
|
||||
prefill_shared_reads_end_at_metadata_init=compliant
|
||||
),
|
||||
war_fastpath_read_done_event=None,
|
||||
)
|
||||
|
||||
|
||||
_DEVICE_MODULE = SimpleNamespace(Event=_Event)
|
||||
|
||||
|
||||
def _batch(mode=ForwardMode.EXTEND):
|
||||
return SimpleNamespace(forward_mode=mode)
|
||||
|
||||
|
||||
def test_publishes_recorded_event_when_enabled():
|
||||
runner = _model_runner()
|
||||
with envs.SGLANG_ENABLE_PREFILL_WAR_READ_DONE.override(True):
|
||||
maybe_publish_prefill_war_read_done(runner, _batch(), _DEVICE_MODULE)
|
||||
published = runner.war_fastpath_read_done_event
|
||||
assert isinstance(published, _Event) and published.recorded
|
||||
|
||||
|
||||
def test_disabled_when_flag_is_false():
|
||||
runner = _model_runner()
|
||||
with envs.SGLANG_ENABLE_PREFILL_WAR_READ_DONE.override(False):
|
||||
maybe_publish_prefill_war_read_done(runner, _batch(), _DEVICE_MODULE)
|
||||
assert runner.war_fastpath_read_done_event is None
|
||||
|
||||
|
||||
def test_gates_exclude_non_prefill_unsupported_algorithm_and_noncompliant_backend():
|
||||
with envs.SGLANG_ENABLE_PREFILL_WAR_READ_DONE.override(True):
|
||||
for runner, batch in (
|
||||
# Verify/mixed/decode publish through the decode graph runner.
|
||||
(_model_runner(), _batch(ForwardMode.TARGET_VERIFY)),
|
||||
(_model_runner(), _batch(ForwardMode.MIXED)),
|
||||
(_model_runner(), _batch(ForwardMode.DECODE)),
|
||||
# The algorithm has a later prefill reader or unverified ownership.
|
||||
(_model_runner(spec_algorithm=SpeculativeAlgorithm.EAGLE), _batch()),
|
||||
# Backend has not declared metadata-init compliance.
|
||||
(_model_runner(compliant=False), _batch()),
|
||||
):
|
||||
maybe_publish_prefill_war_read_done(runner, batch, _DEVICE_MODULE)
|
||||
assert runner.war_fastpath_read_done_event is None
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(pytest.main([__file__, "-v"]))
|
||||
Reference in New Issue
Block a user