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:
@@ -152,6 +152,16 @@ class AttentionBackend(ABC):
|
|||||||
return SharedReadEnds.IN_REPLAY
|
return SharedReadEnds.IN_REPLAY
|
||||||
return SharedReadEnds.UNKNOWN
|
return SharedReadEnds.UNKNOWN
|
||||||
|
|
||||||
|
def prepare_prefill_shared_read_snapshot(
|
||||||
|
self, forward_batch: ForwardBatch, *, num_qo_tokens: int
|
||||||
|
) -> None:
|
||||||
|
"""Snapshot late prefill reads before a PRE_REPLAY event is published.
|
||||||
|
|
||||||
|
Runners call this only after the actual eager/replay query geometry is
|
||||||
|
known. Backends that retain scheduler-shared reads into the model
|
||||||
|
forward keep the default no-op and must not declare PRE_REPLAY.
|
||||||
|
"""
|
||||||
|
|
||||||
# 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
|
||||||
# assume that generic ForwardBatch metadata is sufficient for every
|
# assume that generic ForwardBatch metadata is sufficient for every
|
||||||
|
|||||||
@@ -408,10 +408,10 @@ class DSV4Metadata:
|
|||||||
c4_compress_metadata: Optional[FusedCompressMetadata] = None
|
c4_compress_metadata: Optional[FusedCompressMetadata] = None
|
||||||
c128_compress_metadata: Optional[FusedCompressMetadata] = None
|
c128_compress_metadata: Optional[FusedCompressMetadata] = None
|
||||||
|
|
||||||
# Lazily populated on the first call to ``_forward_prefill_sparse`` and
|
# Built at the runner's prefill WAR boundary when the fast path is on,
|
||||||
# reused across every layer in the chunk. Reset to ``None`` when graph
|
# otherwise lazily by ``_forward_prefill_sparse``.
|
||||||
# metadata is refreshed so replay rebuilds it from the live batch.
|
|
||||||
sparse_prefill_cache: Optional[SparsePrefillChunkCache] = None
|
sparse_prefill_cache: Optional[SparsePrefillChunkCache] = None
|
||||||
|
prefill_shared_reads_snapshotted: bool = False
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def core_metadata(self) -> DSV4AttnMetadata:
|
def core_metadata(self) -> DSV4AttnMetadata:
|
||||||
@@ -425,6 +425,7 @@ class DSV4Metadata:
|
|||||||
self.c128_compress_metadata, src=other.c128_compress_metadata
|
self.c128_compress_metadata, src=other.c128_compress_metadata
|
||||||
)
|
)
|
||||||
self.sparse_prefill_cache = None
|
self.sparse_prefill_cache = None
|
||||||
|
self.prefill_shared_reads_snapshotted = False
|
||||||
|
|
||||||
def refresh_for_breakable_cuda_graph_replay_(self, static_metadata: DSV4Metadata):
|
def refresh_for_breakable_cuda_graph_replay_(self, static_metadata: DSV4Metadata):
|
||||||
self.core_attn_metadata.refresh_for_breakable_cuda_graph_replay_(
|
self.core_attn_metadata.refresh_for_breakable_cuda_graph_replay_(
|
||||||
@@ -444,6 +445,7 @@ class DSV4Metadata:
|
|||||||
src=static_metadata.c128_compress_metadata,
|
src=static_metadata.c128_compress_metadata,
|
||||||
)
|
)
|
||||||
self.sparse_prefill_cache = None
|
self.sparse_prefill_cache = None
|
||||||
|
self.prefill_shared_reads_snapshotted = False
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -519,6 +521,13 @@ class DeepseekV4AttnBackend(
|
|||||||
if self.model_runner.spec_algorithm.is_dspark():
|
if self.model_runner.spec_algorithm.is_dspark():
|
||||||
return SharedReadEnds.IN_REPLAY
|
return SharedReadEnds.IN_REPLAY
|
||||||
return SharedReadEnds.POST_REPLAY
|
return SharedReadEnds.POST_REPLAY
|
||||||
|
metadata = self.forward_metadata
|
||||||
|
if (
|
||||||
|
fm == ForwardMode.EXTEND
|
||||||
|
and isinstance(metadata, DSV4Metadata)
|
||||||
|
and metadata.prefill_shared_reads_snapshotted
|
||||||
|
):
|
||||||
|
return SharedReadEnds.PRE_REPLAY
|
||||||
return super().shared_read_ends(fm)
|
return super().shared_read_ends(fm)
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -1352,6 +1361,66 @@ class DeepseekV4AttnBackend(
|
|||||||
self.forward_metadata = self._build_forward_metadata(forward_batch)
|
self.forward_metadata = self._build_forward_metadata(forward_batch)
|
||||||
self.init_forward_metadata_in_graph(forward_batch)
|
self.init_forward_metadata_in_graph(forward_batch)
|
||||||
|
|
||||||
|
def prepare_prefill_shared_read_snapshot(
|
||||||
|
self, forward_batch: ForwardBatch, *, num_qo_tokens: int
|
||||||
|
) -> None:
|
||||||
|
# Sparse prefill otherwise reads req_to_token/full_to_swa lazily in its
|
||||||
|
# first layer. DFLASH/DSPARK have no later prefill draft-extend reader;
|
||||||
|
# CP-v2 shards the query layout that this global snapshot assumes.
|
||||||
|
metadata = self.forward_metadata
|
||||||
|
if isinstance(metadata, DSV4Metadata):
|
||||||
|
metadata.prefill_shared_reads_snapshotted = False
|
||||||
|
snapshot_shared_prefill_reads = (
|
||||||
|
envs.SGLANG_ENABLE_PREFILL_WAR_READ_DONE.get()
|
||||||
|
and forward_batch.forward_mode == ForwardMode.EXTEND
|
||||||
|
and self.model_runner.spec_algorithm.is_dflash_family()
|
||||||
|
and not is_cp_v2_active(forward_batch)
|
||||||
|
)
|
||||||
|
if not snapshot_shared_prefill_reads:
|
||||||
|
return
|
||||||
|
|
||||||
|
assert isinstance(metadata, DSV4Metadata)
|
||||||
|
use_sparse_prefill = not _is_sm120 and (
|
||||||
|
num_qo_tokens > _LARGE_INDEXER_QUERY_THRESHOLD
|
||||||
|
or envs.SGLANG_OPT_FLASHMLA_SPARSE_PREFILL.get()
|
||||||
|
)
|
||||||
|
if use_sparse_prefill:
|
||||||
|
metadata.sparse_prefill_cache = self._build_sparse_prefill_chunk_cache(
|
||||||
|
forward_batch, num_qo_tokens=num_qo_tokens
|
||||||
|
)
|
||||||
|
# Marked for dense prefill too: that path reads only core_attn_metadata,
|
||||||
|
# which init_forward_metadata already snapshotted.
|
||||||
|
metadata.prefill_shared_reads_snapshotted = True
|
||||||
|
|
||||||
|
def _build_sparse_prefill_chunk_cache(
|
||||||
|
self, forward_batch: ForwardBatch, *, num_qo_tokens: int
|
||||||
|
) -> SparsePrefillChunkCache:
|
||||||
|
seq_lens_cpu = forward_batch.seq_lens_cpu
|
||||||
|
assert seq_lens_cpu is not None
|
||||||
|
extend_seq_lens_cpu = forward_batch.extend_seq_lens_cpu
|
||||||
|
assert extend_seq_lens_cpu is not None
|
||||||
|
seq_lens_cpu_list = seq_lens_cpu.tolist()
|
||||||
|
total_swa = sum(
|
||||||
|
min(int(seq_len), int(extend_len) + SWA_WINDOW - 1)
|
||||||
|
for seq_len, extend_len in zip(
|
||||||
|
seq_lens_cpu_list, extend_seq_lens_cpu, strict=True
|
||||||
|
)
|
||||||
|
)
|
||||||
|
# ``swa_window_size`` on the pool is its storage page size, not the
|
||||||
|
# model's SWA window, so pass both explicitly.
|
||||||
|
return SparsePrefillChunkCache.build(
|
||||||
|
seq_lens=forward_batch.seq_lens.to(torch.int32),
|
||||||
|
extend_seq_lens=forward_batch.extend_seq_lens.to(torch.int32),
|
||||||
|
req_pool_indices=forward_batch.req_pool_indices.to(torch.int32),
|
||||||
|
req_to_token=self.req_to_token,
|
||||||
|
full_to_swa=self.token_to_kv_pool.full_to_swa_index_mapping,
|
||||||
|
swa_window_size=SWA_WINDOW,
|
||||||
|
swa_page_size=self.token_to_kv_pool.swa_window_size,
|
||||||
|
num_qo_tokens=num_qo_tokens,
|
||||||
|
max_seq_len=max(seq_lens_cpu_list),
|
||||||
|
total_swa=total_swa,
|
||||||
|
)
|
||||||
|
|
||||||
def _build_forward_metadata(
|
def _build_forward_metadata(
|
||||||
self,
|
self,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
@@ -1791,29 +1860,8 @@ class DeepseekV4AttnBackend(
|
|||||||
|
|
||||||
cache = self.forward_metadata.sparse_prefill_cache
|
cache = self.forward_metadata.sparse_prefill_cache
|
||||||
if cache is None:
|
if cache is None:
|
||||||
seq_lens_cpu = forward_batch.seq_lens_cpu
|
cache = self._build_sparse_prefill_chunk_cache(
|
||||||
assert seq_lens_cpu is not None
|
forward_batch, num_qo_tokens=q_flat.shape[0]
|
||||||
extend_seq_lens_cpu = forward_batch.extend_seq_lens_cpu
|
|
||||||
assert extend_seq_lens_cpu is not None
|
|
||||||
total_swa = sum(
|
|
||||||
min(int(seq_len), int(extend_len) + SWA_WINDOW - 1)
|
|
||||||
for seq_len, extend_len in zip(
|
|
||||||
seq_lens_cpu.tolist(), extend_seq_lens_cpu, strict=True
|
|
||||||
)
|
|
||||||
)
|
|
||||||
# ``swa_window_size`` on the pool is its storage page size, not
|
|
||||||
# the model's SWA window — pass both explicitly.
|
|
||||||
cache = SparsePrefillChunkCache.build(
|
|
||||||
seq_lens=forward_batch.seq_lens.to(torch.int32),
|
|
||||||
extend_seq_lens=forward_batch.extend_seq_lens.to(torch.int32),
|
|
||||||
req_pool_indices=forward_batch.req_pool_indices.to(torch.int32),
|
|
||||||
req_to_token=self.req_to_token,
|
|
||||||
full_to_swa=token_to_kv_pool.full_to_swa_index_mapping,
|
|
||||||
swa_window_size=SWA_WINDOW,
|
|
||||||
swa_page_size=token_to_kv_pool.swa_window_size,
|
|
||||||
num_qo_tokens=q_flat.shape[0],
|
|
||||||
max_seq_len=int(seq_lens_cpu.max().item()),
|
|
||||||
total_swa=total_swa,
|
|
||||||
)
|
)
|
||||||
self.forward_metadata.sparse_prefill_cache = cache
|
self.forward_metadata.sparse_prefill_cache = cache
|
||||||
|
|
||||||
@@ -1987,27 +2035,8 @@ class DeepseekV4AttnBackend(
|
|||||||
|
|
||||||
cache = self.forward_metadata.sparse_prefill_cache
|
cache = self.forward_metadata.sparse_prefill_cache
|
||||||
if cache is None:
|
if cache is None:
|
||||||
seq_lens_cpu = forward_batch.seq_lens_cpu
|
cache = self._build_sparse_prefill_chunk_cache(
|
||||||
assert seq_lens_cpu is not None
|
forward_batch, num_qo_tokens=q_flat.shape[0]
|
||||||
extend_seq_lens_cpu = forward_batch.extend_seq_lens_cpu
|
|
||||||
assert extend_seq_lens_cpu is not None
|
|
||||||
total_swa = sum(
|
|
||||||
min(int(seq_len), int(extend_len) + SWA_WINDOW - 1)
|
|
||||||
for seq_len, extend_len in zip(
|
|
||||||
seq_lens_cpu.tolist(), extend_seq_lens_cpu, strict=True
|
|
||||||
)
|
|
||||||
)
|
|
||||||
cache = SparsePrefillChunkCache.build(
|
|
||||||
seq_lens=forward_batch.seq_lens.to(torch.int32),
|
|
||||||
extend_seq_lens=forward_batch.extend_seq_lens.to(torch.int32),
|
|
||||||
req_pool_indices=forward_batch.req_pool_indices.to(torch.int32),
|
|
||||||
req_to_token=self.req_to_token,
|
|
||||||
full_to_swa=token_to_kv_pool.full_to_swa_index_mapping,
|
|
||||||
swa_window_size=SWA_WINDOW,
|
|
||||||
swa_page_size=token_to_kv_pool.swa_window_size,
|
|
||||||
num_qo_tokens=q_flat.shape[0],
|
|
||||||
max_seq_len=int(seq_lens_cpu.max().item()),
|
|
||||||
total_swa=total_swa,
|
|
||||||
)
|
)
|
||||||
self.forward_metadata.sparse_prefill_cache = cache
|
self.forward_metadata.sparse_prefill_cache = cache
|
||||||
|
|
||||||
|
|||||||
@@ -49,6 +49,9 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
|
|||||||
enable_tc_piecewise_cuda_graph,
|
enable_tc_piecewise_cuda_graph,
|
||||||
set_tc_piecewise_forward_context,
|
set_tc_piecewise_forward_context,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.model_executor.runner_utils import (
|
||||||
|
maybe_publish_prefill_shared_read_done,
|
||||||
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_spec,
|
get_spec,
|
||||||
@@ -305,6 +308,15 @@ class EagerRunner(BaseRunner):
|
|||||||
# e.g. Moss-VL's prefill cross-attention custom mask.
|
# e.g. Moss-VL's prefill cross-attention custom mask.
|
||||||
model_runner.model.prepare_forward_batch(forward_batch)
|
model_runner.model.prepare_forward_batch(forward_batch)
|
||||||
model_runner.attn_backend.init_forward_metadata(forward_batch)
|
model_runner.attn_backend.init_forward_metadata(forward_batch)
|
||||||
|
model_runner.attn_backend.prepare_prefill_shared_read_snapshot(
|
||||||
|
forward_batch,
|
||||||
|
num_qo_tokens=len(forward_batch.input_ids),
|
||||||
|
)
|
||||||
|
maybe_publish_prefill_shared_read_done(
|
||||||
|
model_runner,
|
||||||
|
forward_batch,
|
||||||
|
torch.get_device_module(model_runner.device),
|
||||||
|
)
|
||||||
|
|
||||||
if not cp_v2_active:
|
if not cp_v2_active:
|
||||||
forward_batch.attn_cp_metadata = None
|
forward_batch.attn_cp_metadata = None
|
||||||
|
|||||||
@@ -1044,6 +1044,9 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
return
|
return
|
||||||
if not self.use_captured_attn_metadata:
|
if not self.use_captured_attn_metadata:
|
||||||
attn_backend.init_forward_metadata(forward_batch)
|
attn_backend.init_forward_metadata(forward_batch)
|
||||||
|
attn_backend.prepare_prefill_shared_read_snapshot(
|
||||||
|
forward_batch, num_qo_tokens=num_tokens
|
||||||
|
)
|
||||||
return
|
return
|
||||||
assert self.attn_metadata_buffers is not None
|
assert self.attn_metadata_buffers is not None
|
||||||
metadata = self.attn_metadata_buffers[num_tokens]
|
metadata = self.attn_metadata_buffers[num_tokens]
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
"""Shared-read-done event utilities for CUDA graph runners."""
|
"""Shared-read-done event utilities for graph and eager runners."""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
@@ -31,9 +31,13 @@ def maybe_publish_prefill_shared_read_done(
|
|||||||
return
|
return
|
||||||
if forward_batch.forward_mode != ForwardMode.EXTEND:
|
if forward_batch.forward_mode != ForwardMode.EXTEND:
|
||||||
return
|
return
|
||||||
# TODO(Jialin): Relax this gate for speculative decoding after its prefill
|
# TODO(Jialin): Relax for EAGLE/MTP after validating the later
|
||||||
# WAR boundaries are validated.
|
# draft-extend reader's WAR boundary.
|
||||||
if not model_runner.spec_algorithm.is_none():
|
if (
|
||||||
|
not model_runner.spec_algorithm.is_none()
|
||||||
|
and not model_runner.spec_algorithm.is_dflash_family()
|
||||||
|
):
|
||||||
|
# Other speculative algorithms may have a later draft-extend reader.
|
||||||
return
|
return
|
||||||
# The record lands right after replay prep, so PRE_REPLAY only.
|
# The record lands right after replay prep, so PRE_REPLAY only.
|
||||||
declared = model_runner.attn_backend.shared_read_ends(forward_batch.forward_mode)
|
declared = model_runner.attn_backend.shared_read_ends(forward_batch.forward_mode)
|
||||||
|
|||||||
@@ -391,6 +391,113 @@ class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase):
|
|||||||
DeepseekV4AttnBackend.use_captured_forward_metadata_for_breakable_cuda_graph
|
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):
|
def test_refresh_replay_metadata_preserves_captured_tensor_storage(self):
|
||||||
capture_metadata = self._make_core_metadata(0)
|
capture_metadata = self._make_core_metadata(0)
|
||||||
replay_metadata = self._make_core_metadata(1000)
|
replay_metadata = self._make_core_metadata(1000)
|
||||||
@@ -459,6 +566,7 @@ class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase):
|
|||||||
self._make_core_metadata(0), indexer_metadata=None
|
self._make_core_metadata(0), indexer_metadata=None
|
||||||
)
|
)
|
||||||
capture_metadata.sparse_prefill_cache = object()
|
capture_metadata.sparse_prefill_cache = object()
|
||||||
|
capture_metadata.prefill_shared_reads_snapshotted = True
|
||||||
replay_metadata = DSV4Metadata(
|
replay_metadata = DSV4Metadata(
|
||||||
self._make_core_metadata(1000), indexer_metadata=None
|
self._make_core_metadata(1000), indexer_metadata=None
|
||||||
)
|
)
|
||||||
@@ -487,6 +595,7 @@ class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase):
|
|||||||
self.assertTrue(calls[0][2])
|
self.assertTrue(calls[0][2])
|
||||||
self.assertIs(backend.forward_metadata, capture_metadata)
|
self.assertIs(backend.forward_metadata, capture_metadata)
|
||||||
self.assertIsNone(capture_metadata.sparse_prefill_cache)
|
self.assertIsNone(capture_metadata.sparse_prefill_cache)
|
||||||
|
self.assertFalse(capture_metadata.prefill_shared_reads_snapshotted)
|
||||||
self.assertTrue(
|
self.assertTrue(
|
||||||
torch.equal(
|
torch.equal(
|
||||||
capture_metadata.core_attn_metadata.seq_lens_casual,
|
capture_metadata.core_attn_metadata.seq_lens_casual,
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
from unittest import mock
|
||||||
|
|
||||||
from sglang.srt.model_executor.cuda_graph_config import Backend
|
from sglang.srt.model_executor.cuda_graph_config import Backend
|
||||||
from sglang.srt.model_executor.forward_batch_info import (
|
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)))
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -60,6 +60,17 @@ def test_disabled_when_flag_is_false():
|
|||||||
assert runner.shared_read_done_event is None
|
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():
|
def test_gates_exclude_non_prefill_unsupported_algorithm_and_noncompliant_backend():
|
||||||
with envs.SGLANG_ENABLE_PREFILL_WAR_READ_DONE.override(True):
|
with envs.SGLANG_ENABLE_PREFILL_WAR_READ_DONE.override(True):
|
||||||
for runner, batch in (
|
for runner, batch in (
|
||||||
|
|||||||
Reference in New Issue
Block a user