[DeepSeek V4] Cover both dense and sparse prefill paths in the compress attention unittest (#29885)

Co-authored-by: Yuwei An <ayw.sirius19@gmail.com>
This commit is contained in:
YAMY
2026-07-01 21:20:12 -07:00
committed by GitHub
co-authored by Yuwei An
parent 9ba4b8f8ba
commit 307094dc7d
2 changed files with 105 additions and 17 deletions
@@ -19,6 +19,7 @@ from typing import Any
import torch import torch
from torch import nn from torch import nn
from sglang.srt.environ import envs
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.dsv4.quant_k_cache import ( from sglang.srt.layers.attention.dsv4.quant_k_cache import (
quant_to_nope_fp8_rope_bf16_pack_triton, quant_to_nope_fp8_rope_bf16_pack_triton,
@@ -594,6 +595,9 @@ class DSV4AttentionFixture:
forward_batch: ForwardBatch forward_batch: ForwardBatch
prefix_hidden: list[torch.Tensor] prefix_hidden: list[torch.Tensor]
input_hidden: torch.Tensor input_hidden: torch.Tensor
# Selects dense vs sparse-prefill C4 seeding; lives on the fixture because
# the reference re-seeds after rebuilding metadata (`_seed_c4_if_needed`).
seed_c4_for_sparse_prefill: bool = False
@dataclass @dataclass
@@ -1096,14 +1100,20 @@ def prepare_dsv4_runner_inputs(
_populate_extra_kv_cache(fixture, layer_id=0, num_entries=_DSV4_EXTRA_ENTRIES) _populate_extra_kv_cache(fixture, layer_id=0, num_entries=_DSV4_EXTRA_ENTRIES)
def _seed_c4_if_needed(fixture: DSV4AttentionFixture) -> None: def _seed_c4_if_needed(
"""For compress_ratio=4, seed `c4_sparse_page_indices` to the entries the fixture: DSV4AttentionFixture, *, num_entries: int = _DSV4_EXTRA_ENTRIES
fixture wrote via `_populate_extra_kv_cache` (the C4Indexer would normally ) -> None:
populate this; the smoke fixture skips the indexer). No-op for other """For compress_ratio=4, seed the C4 metadata the exercised path consumes
compress_ratios. (the C4Indexer would normally populate it; the smoke fixture skips the
indexer): `c4_sparse_page_indices` for the dense extend path,
`c4_sparse_raw_indices` for sparse prefill. No-op for other compress_ratios.
""" """
if fixture.case.compress_ratio == 4: if fixture.case.compress_ratio != 4:
_seed_c4_sparse_indices(fixture, num_entries=_DSV4_EXTRA_ENTRIES) return
if fixture.seed_c4_for_sparse_prefill:
_seed_c4_sparse_prefill_indices(fixture, num_entries=num_entries)
else:
_seed_c4_sparse_indices(fixture, num_entries=num_entries)
def run_dsv4_fixture_eager(fixture: DSV4AttentionFixture) -> torch.Tensor: def run_dsv4_fixture_eager(fixture: DSV4AttentionFixture) -> torch.Tensor:
@@ -1401,6 +1411,47 @@ def _seed_c4_sparse_indices(
) )
def _seed_c4_sparse_prefill_indices(
fixture: DSV4AttentionFixture,
*,
num_entries: int,
) -> None:
"""Seed C4 metadata for the sparse prefill extend path.
`_forward_prefill_sparse` reads `c4_sparse_raw_indices` (request-local
compressed positions, normally the indexer's output) and derives per-query
lengths as `(pos + 1) // 4`. Seed the sequential positions the indexer
emits for short sequences and mirror the same causal set into
`c4_sparse_page_indices` / `c4_sparse_topk_lengths` so the reference
attends identical entries. The mirror relies on raw position `k` mapping
to physical extra-cache id `k` (page 0 of a fresh single-request layout);
asserted below.
"""
md = fixture.backend.forward_metadata.core_metadata
raw_indices = md.c4_sparse_raw_indices
assert raw_indices is not None, "requires init_flashmla_related(is_prefill=True)"
num_q, width = raw_indices.shape
lens = (md.positions_casual + 1) // 4
max_len = int(lens.max().item())
pool = fixture.runner.token_to_kv_pool
c4_page_size = pool.get_extra_key_page_size(layer_id=0)
assert max_len <= min(
num_entries, c4_page_size
), f"case attends {max_len} c4 entries; only {min(num_entries, c4_page_size)} populated"
assert (
md.page_table[:, 0] == 0
).all(), "sparse seeding requires the raw==physical identity (first page 0)"
seq = (
torch.arange(width, dtype=raw_indices.dtype, device=raw_indices.device)
.unsqueeze(0)
.expand(num_q, -1)
)
seeded = torch.where(seq < lens.unsqueeze(1), seq, seq.new_full((), -1))
md.c4_sparse_raw_indices = seeded
md.c4_sparse_page_indices = seeded.clone()
md.c4_sparse_topk_lengths = lens.to(md.c4_sparse_topk_lengths.dtype)
def run_dsv4_target_verify_attention_case( def run_dsv4_target_verify_attention_case(
testcase, testcase,
case: DSV4AttentionCase, case: DSV4AttentionCase,
@@ -1540,6 +1591,7 @@ def run_dsv4_compress_attention_case(
case: DSV4AttentionCase, case: DSV4AttentionCase,
*, *,
extra_entries: int = 32, extra_entries: int = 32,
sparse_prefill: bool = False,
dtype: torch.dtype = torch.bfloat16, dtype: torch.dtype = torch.bfloat16,
device: str = "cuda", device: str = "cuda",
) -> None: ) -> None:
@@ -1549,17 +1601,24 @@ def run_dsv4_compress_attention_case(
Pre-writes random packed K into both the SWA cache and the extra Pre-writes random packed K into both the SWA cache and the extra
(C4/C128) cache via the production pack+set paths, lets (C4/C128) cache via the production pack+set paths, lets
`init_forward_metadata` populate the compression metadata, manually seeds `init_forward_metadata` populate the compression metadata, manually seeds
`c4_sparse_page_indices` for the C4 case (so the flash_mla `extra_k_cache` the C4 metadata the exercised path consumes (see `_seed_c4_if_needed`; the
path actually attends to entries we wrote rather than the all-`-1` initial un-run indexer would otherwise leave it at `-1` / uninitialized), then
value that the un-run indexer would leave), then dispatches `forward( dispatches `forward(compress_ratio=case.compress_ratio)` and compares
compress_ratio=case.compress_ratio)` and compares against an independent against an independent pure-PyTorch SWA + extra reference that reads the
pure-PyTorch SWA + extra reference that reads the SAME cache bytes and SAME cache bytes and metadata indices.
metadata indices.
`sparse_prefill` pins `SGLANG_OPT_FLASHMLA_SPARSE_PREFILL`, selecting the
dense `flash_mla_with_kvcache` extend path or `_forward_prefill_sparse`;
the C4 seeding dispatches on the same flag.
""" """
assert case.compress_ratio in ( assert case.compress_ratio in (
4, 4,
128, 128,
), f"smoke runner requires compress_ratio in (4, 128); got {case.compress_ratio}" ), f"smoke runner requires compress_ratio in (4, 128); got {case.compress_ratio}"
if sparse_prefill:
assert (
case.forward_mode.is_extend_without_speculative()
), f"sparse prefill only serves extend; got {case.forward_mode}"
fixture = build_dsv4_attention_fixture( fixture = build_dsv4_attention_fixture(
testcase, testcase,
case, case,
@@ -1567,6 +1626,7 @@ def run_dsv4_compress_attention_case(
device=device, device=device,
compression_ratios=[case.compress_ratio], compression_ratios=[case.compress_ratio],
) )
fixture.seed_c4_for_sparse_prefill = sparse_prefill
runner = fixture.runner runner = fixture.runner
max_context_len = runner.req_to_token_pool.req_to_token.shape[1] max_context_len = runner.req_to_token_pool.req_to_token.shape[1]
@@ -1574,10 +1634,13 @@ def run_dsv4_compress_attention_case(
_populate_extra_kv_cache(fixture, layer_id=0, num_entries=extra_entries) _populate_extra_kv_cache(fixture, layer_id=0, num_entries=extra_entries)
q_input, _ = fixture.actual_module.project(fixture.input_hidden) q_input, _ = fixture.actual_module.project(fixture.input_hidden)
with torch.no_grad(), forward_context(ForwardContext(attn_backend=fixture.backend)): with (
torch.no_grad(),
forward_context(ForwardContext(attn_backend=fixture.backend)),
envs.SGLANG_OPT_FLASHMLA_SPARSE_PREFILL.override(sparse_prefill),
):
fixture.backend.init_forward_metadata(fixture.forward_batch) fixture.backend.init_forward_metadata(fixture.forward_batch)
if case.compress_ratio == 4: _seed_c4_if_needed(fixture, num_entries=extra_entries)
_seed_c4_sparse_indices(fixture, num_entries=extra_entries)
actual = fixture.backend.forward( actual = fixture.backend.forward(
q=q_input, q=q_input,
k=q_input, k=q_input,
@@ -1588,6 +1651,17 @@ def run_dsv4_compress_attention_case(
save_kv_cache=False, save_kv_cache=False,
attn_sink=fixture.actual_module.attn_sink, attn_sink=fixture.actual_module.attn_sink,
) )
# Only `_forward_prefill_sparse` populates `sparse_prefill_cache`;
# verify the intended path ran before the reference rebuilds metadata.
sparse_cache = fixture.backend.forward_metadata.sparse_prefill_cache
if sparse_prefill:
testcase.assertIsNotNone(
sparse_cache, f"{case.name} did not take _forward_prefill_sparse"
)
else:
testcase.assertIsNone(
sparse_cache, f"{case.name} did not take the dense extend path"
)
expected = _pure_torch_dsv4_combined_reference(fixture, q_input) expected = _pure_torch_dsv4_combined_reference(fixture, q_input)
torch.testing.assert_close( torch.testing.assert_close(
@@ -183,13 +183,27 @@ class TestDSV4AttentionBackendCorrectness(CustomTestCase):
) )
def test_compress_attention_cases(self): def test_compress_attention_cases(self):
# Pinned to the dense extend path; the sparse prefill path is covered
# by test_compress_attention_cases_sparse_prefill below.
for case in self.COMPRESS_CASES: for case in self.COMPRESS_CASES:
with self.subTest( with self.subTest(
case=case.name, case=case.name,
backend=case.backend, backend=case.backend,
compress_ratio=case.compress_ratio, compress_ratio=case.compress_ratio,
): ):
run_dsv4_compress_attention_case(self, case) run_dsv4_compress_attention_case(self, case, sparse_prefill=False)
def test_compress_attention_cases_sparse_prefill(self):
# `_forward_prefill_sparse` extend path; decode never reaches it.
for case in self.COMPRESS_CASES:
if not case.forward_mode.is_extend_without_speculative():
continue
with self.subTest(
case=case.name,
backend=case.backend,
compress_ratio=case.compress_ratio,
):
run_dsv4_compress_attention_case(self, case, sparse_prefill=True)
def test_eagle_target_verify_chain_cases(self): def test_eagle_target_verify_chain_cases(self):
for case in self.TARGET_VERIFY_CASES: for case in self.TARGET_VERIFY_CASES: