Publish gated DSV4 DFLASH-family target-prefill read completion (#35947)
Co-authored-by: weireweire <20922698+weireweire@users.noreply.github.com>
This commit is contained in:
@@ -391,6 +391,113 @@ class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase):
|
||||
DeepseekV4AttnBackend.use_captured_forward_metadata_for_breakable_cuda_graph
|
||||
)
|
||||
|
||||
def test_prefill_snapshot_declares_pre_replay_boundary(self):
|
||||
from sglang.srt.layers.attention.base_attn_backend import SharedReadEnds
|
||||
from sglang.srt.layers.attention.deepseek_v4_backend import (
|
||||
DeepseekV4AttnBackend,
|
||||
DSV4Metadata,
|
||||
)
|
||||
|
||||
backend = object.__new__(DeepseekV4AttnBackend)
|
||||
backend.forward_metadata = DSV4Metadata(
|
||||
self._make_core_metadata(0), indexer_metadata=None
|
||||
)
|
||||
self.assertIs(
|
||||
backend.shared_read_ends(ForwardMode.EXTEND),
|
||||
SharedReadEnds.UNKNOWN,
|
||||
)
|
||||
|
||||
backend.forward_metadata.prefill_shared_reads_snapshotted = True
|
||||
self.assertIs(
|
||||
backend.shared_read_ends(ForwardMode.EXTEND),
|
||||
SharedReadEnds.PRE_REPLAY,
|
||||
)
|
||||
|
||||
def test_snapshot_builds_cache_only_for_sparse_prefill(self):
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.deepseek_v4_backend import (
|
||||
_LARGE_INDEXER_QUERY_THRESHOLD,
|
||||
DeepseekV4AttnBackend,
|
||||
DSV4Metadata,
|
||||
)
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
|
||||
batch = SimpleNamespace(forward_mode=ForwardMode.EXTEND)
|
||||
cache = object()
|
||||
for num_qo_tokens, builds in (
|
||||
(_LARGE_INDEXER_QUERY_THRESHOLD, False),
|
||||
(_LARGE_INDEXER_QUERY_THRESHOLD + 1, True),
|
||||
):
|
||||
with self.subTest(num_qo_tokens=num_qo_tokens):
|
||||
backend = object.__new__(DeepseekV4AttnBackend)
|
||||
backend.model_runner = SimpleNamespace(
|
||||
spec_algorithm=SpeculativeAlgorithm.DFLASH
|
||||
)
|
||||
backend.forward_metadata = DSV4Metadata(
|
||||
self._make_core_metadata(0), indexer_metadata=None
|
||||
)
|
||||
backend._build_sparse_prefill_chunk_cache = mock.Mock(
|
||||
return_value=cache
|
||||
)
|
||||
|
||||
with (
|
||||
envs.SGLANG_ENABLE_PREFILL_WAR_READ_DONE.override(True),
|
||||
envs.SGLANG_OPT_FLASHMLA_SPARSE_PREFILL.override(False),
|
||||
mock.patch(
|
||||
"sglang.srt.layers.attention.deepseek_v4_backend._is_sm120",
|
||||
False,
|
||||
),
|
||||
):
|
||||
backend.prepare_prefill_shared_read_snapshot(
|
||||
batch, num_qo_tokens=num_qo_tokens
|
||||
)
|
||||
|
||||
metadata = backend.forward_metadata
|
||||
if builds:
|
||||
backend._build_sparse_prefill_chunk_cache.assert_called_once_with(
|
||||
batch, num_qo_tokens=num_qo_tokens
|
||||
)
|
||||
self.assertIs(metadata.sparse_prefill_cache, cache)
|
||||
else:
|
||||
backend._build_sparse_prefill_chunk_cache.assert_not_called()
|
||||
self.assertIsNone(metadata.sparse_prefill_cache)
|
||||
# Dense declares the boundary too; it reads only the metadata
|
||||
# that init_forward_metadata already snapshotted.
|
||||
self.assertTrue(metadata.prefill_shared_reads_snapshotted)
|
||||
|
||||
def test_sparse_prefill_snapshot_marks_success_only_after_build(self):
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.deepseek_v4_backend import (
|
||||
DeepseekV4AttnBackend,
|
||||
DSV4Metadata,
|
||||
)
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
|
||||
backend = object.__new__(DeepseekV4AttnBackend)
|
||||
backend.model_runner = SimpleNamespace(
|
||||
spec_algorithm=SpeculativeAlgorithm.DFLASH
|
||||
)
|
||||
backend.forward_metadata = DSV4Metadata(
|
||||
self._make_core_metadata(0), indexer_metadata=None
|
||||
)
|
||||
backend.forward_metadata.prefill_shared_reads_snapshotted = True
|
||||
backend._build_sparse_prefill_chunk_cache = mock.Mock(
|
||||
side_effect=RuntimeError("snapshot failed")
|
||||
)
|
||||
batch = SimpleNamespace(forward_mode=ForwardMode.EXTEND)
|
||||
|
||||
with (
|
||||
envs.SGLANG_ENABLE_PREFILL_WAR_READ_DONE.override(True),
|
||||
envs.SGLANG_OPT_FLASHMLA_SPARSE_PREFILL.override(True),
|
||||
mock.patch(
|
||||
"sglang.srt.layers.attention.deepseek_v4_backend._is_sm120", False
|
||||
),
|
||||
self.assertRaisesRegex(RuntimeError, "snapshot failed"),
|
||||
):
|
||||
backend.prepare_prefill_shared_read_snapshot(batch, num_qo_tokens=12288)
|
||||
|
||||
self.assertFalse(backend.forward_metadata.prefill_shared_reads_snapshotted)
|
||||
|
||||
def test_refresh_replay_metadata_preserves_captured_tensor_storage(self):
|
||||
capture_metadata = self._make_core_metadata(0)
|
||||
replay_metadata = self._make_core_metadata(1000)
|
||||
@@ -459,6 +566,7 @@ class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase):
|
||||
self._make_core_metadata(0), indexer_metadata=None
|
||||
)
|
||||
capture_metadata.sparse_prefill_cache = object()
|
||||
capture_metadata.prefill_shared_reads_snapshotted = True
|
||||
replay_metadata = DSV4Metadata(
|
||||
self._make_core_metadata(1000), indexer_metadata=None
|
||||
)
|
||||
@@ -487,6 +595,7 @@ class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase):
|
||||
self.assertTrue(calls[0][2])
|
||||
self.assertIs(backend.forward_metadata, capture_metadata)
|
||||
self.assertIsNone(capture_metadata.sparse_prefill_cache)
|
||||
self.assertFalse(capture_metadata.prefill_shared_reads_snapshotted)
|
||||
self.assertTrue(
|
||||
torch.equal(
|
||||
capture_metadata.core_attn_metadata.seq_lens_casual,
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest import mock
|
||||
|
||||
from sglang.srt.model_executor.cuda_graph_config import Backend
|
||||
from sglang.srt.model_executor.forward_batch_info import (
|
||||
@@ -52,6 +53,25 @@ class TestPrefillCudaGraphPadding(CustomTestCase):
|
||||
|
||||
self.assertTrue(runner.can_run_graph(self._make_forward_batch(8)))
|
||||
|
||||
def test_replay_snapshot_uses_padded_token_count(self):
|
||||
runner = self._make_runner()
|
||||
runner.use_captured_attn_metadata = False
|
||||
attn_backend = mock.Mock()
|
||||
runner.model_runner = SimpleNamespace(attn_backend=attn_backend)
|
||||
forward_batch = self._make_forward_batch(8)
|
||||
static_forward_batch = self._make_forward_batch(16)
|
||||
|
||||
runner._prepare_forward_metadata_for_replay(
|
||||
forward_batch,
|
||||
static_forward_batch,
|
||||
num_tokens=16,
|
||||
)
|
||||
|
||||
attn_backend.init_forward_metadata.assert_called_once_with(forward_batch)
|
||||
attn_backend.prepare_prefill_shared_read_snapshot.assert_called_once_with(
|
||||
forward_batch, num_qo_tokens=16
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -60,6 +60,17 @@ def test_disabled_when_flag_is_false():
|
||||
assert runner.shared_read_done_event is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"algorithm", (SpeculativeAlgorithm.DFLASH, SpeculativeAlgorithm.DSPARK)
|
||||
)
|
||||
def test_dflash_family_target_prefill_publishes(algorithm):
|
||||
runner = _model_runner(spec_algorithm=algorithm)
|
||||
with envs.SGLANG_ENABLE_PREFILL_WAR_READ_DONE.override(True):
|
||||
maybe_publish_prefill_shared_read_done(runner, _batch(), _DEVICE_MODULE)
|
||||
published = runner.shared_read_done_event
|
||||
assert isinstance(published, _Event) and published.recorded
|
||||
|
||||
|
||||
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 (
|
||||
|
||||
Reference in New Issue
Block a user