[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
|
||||
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.dsv4.quant_k_cache import (
|
||||
quant_to_nope_fp8_rope_bf16_pack_triton,
|
||||
@@ -594,6 +595,9 @@ class DSV4AttentionFixture:
|
||||
forward_batch: ForwardBatch
|
||||
prefix_hidden: list[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
|
||||
@@ -1096,14 +1100,20 @@ def prepare_dsv4_runner_inputs(
|
||||
_populate_extra_kv_cache(fixture, layer_id=0, num_entries=_DSV4_EXTRA_ENTRIES)
|
||||
|
||||
|
||||
def _seed_c4_if_needed(fixture: DSV4AttentionFixture) -> None:
|
||||
"""For compress_ratio=4, seed `c4_sparse_page_indices` to the entries the
|
||||
fixture wrote via `_populate_extra_kv_cache` (the C4Indexer would normally
|
||||
populate this; the smoke fixture skips the indexer). No-op for other
|
||||
compress_ratios.
|
||||
def _seed_c4_if_needed(
|
||||
fixture: DSV4AttentionFixture, *, num_entries: int = _DSV4_EXTRA_ENTRIES
|
||||
) -> None:
|
||||
"""For compress_ratio=4, seed the C4 metadata the exercised path consumes
|
||||
(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:
|
||||
_seed_c4_sparse_indices(fixture, num_entries=_DSV4_EXTRA_ENTRIES)
|
||||
if fixture.case.compress_ratio != 4:
|
||||
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:
|
||||
@@ -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(
|
||||
testcase,
|
||||
case: DSV4AttentionCase,
|
||||
@@ -1540,6 +1591,7 @@ def run_dsv4_compress_attention_case(
|
||||
case: DSV4AttentionCase,
|
||||
*,
|
||||
extra_entries: int = 32,
|
||||
sparse_prefill: bool = False,
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
device: str = "cuda",
|
||||
) -> None:
|
||||
@@ -1549,17 +1601,24 @@ def run_dsv4_compress_attention_case(
|
||||
Pre-writes random packed K into both the SWA cache and the extra
|
||||
(C4/C128) cache via the production pack+set paths, lets
|
||||
`init_forward_metadata` populate the compression metadata, manually seeds
|
||||
`c4_sparse_page_indices` for the C4 case (so the flash_mla `extra_k_cache`
|
||||
path actually attends to entries we wrote rather than the all-`-1` initial
|
||||
value that the un-run indexer would leave), then dispatches `forward(
|
||||
compress_ratio=case.compress_ratio)` and compares against an independent
|
||||
pure-PyTorch SWA + extra reference that reads the SAME cache bytes and
|
||||
metadata indices.
|
||||
the C4 metadata the exercised path consumes (see `_seed_c4_if_needed`; the
|
||||
un-run indexer would otherwise leave it at `-1` / uninitialized), then
|
||||
dispatches `forward(compress_ratio=case.compress_ratio)` and compares
|
||||
against an independent pure-PyTorch SWA + extra reference that reads the
|
||||
SAME cache bytes and 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 (
|
||||
4,
|
||||
128,
|
||||
), 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(
|
||||
testcase,
|
||||
case,
|
||||
@@ -1567,6 +1626,7 @@ def run_dsv4_compress_attention_case(
|
||||
device=device,
|
||||
compression_ratios=[case.compress_ratio],
|
||||
)
|
||||
fixture.seed_c4_for_sparse_prefill = sparse_prefill
|
||||
runner = fixture.runner
|
||||
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)
|
||||
|
||||
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)
|
||||
if case.compress_ratio == 4:
|
||||
_seed_c4_sparse_indices(fixture, num_entries=extra_entries)
|
||||
_seed_c4_if_needed(fixture, num_entries=extra_entries)
|
||||
actual = fixture.backend.forward(
|
||||
q=q_input,
|
||||
k=q_input,
|
||||
@@ -1588,6 +1651,17 @@ def run_dsv4_compress_attention_case(
|
||||
save_kv_cache=False,
|
||||
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)
|
||||
|
||||
torch.testing.assert_close(
|
||||
|
||||
Reference in New Issue
Block a user