[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:
@@ -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:
|
||||||
|
|||||||
Reference in New Issue
Block a user