diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index d57fb45c301f..6ff590c3ba29 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -3273,6 +3273,7 @@ def _prefill_indexer_inputs( kv=kv, k_cache=k_cache.view(k_cache.shape[0], index_page_size, 1, 68), page_size=index_page_size, + # apply_cp_reindex keeps these rows aligned with the local queries. kv_page_table=self.forward_metadata.core_metadata.page_table[:num_tokens], kv_page_size=self.page_size, compress_ratio=ratio, @@ -3297,6 +3298,8 @@ def _publish_prefill(self, published: CandidateMetadata) -> None: if tail_lens is not None: self.tail_forward_metadata.candidate_metadata = ( self.candidate_indexer.prefill_tail(published, tail_lens) + if any(tail_lens) + else None ) def _publish_prefill_masks(self, masks: CandidateMasks) -> None: diff --git a/python/sglang/srt/layers/attention/dsv4/candidate_indexer.py b/python/sglang/srt/layers/attention/dsv4/candidate_indexer.py index f6df650b7df8..ac1e7e5d9e58 100644 --- a/python/sglang/srt/layers/attention/dsv4/candidate_indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/candidate_indexer.py @@ -9,7 +9,7 @@ import torch.nn.functional as F from sglang.srt.layers.attention.dsv4.metadata import PagedIndexerMetadata -from sglang.srt.runtime_context import get_parallel, get_platform +from sglang.srt.runtime_context import get_platform class CandidateMetadata: @@ -130,10 +130,8 @@ def prefill_tail( def make_candidate_indexer( topk_blocks: int, block_size: int ) -> Optional[CandidateIndexer]: - """DeepGEMM's paged sparse indexer on SM100, with the dense-score - implementation for prefill under context parallelism (a rank's local rows - are not the page table's rows); None on Hopper, whose decode and prefill - indexers select through masks inline.""" + """DeepGEMM's paged sparse indexer on SM100, including CP prefill. + Hopper's decode and prefill indexers still select through masks inline.""" if topk_blocks <= 0 or get_platform().device_sm < 100: return None from sglang.srt.layers.deep_gemm_wrapper.configurer import ( @@ -148,16 +146,8 @@ def make_candidate_indexer( from sglang.srt.layers.attention.dsv4.candidate_indexer_deep_gemm import ( DeepGemmCandidateIndexer, ) - from sglang.srt.layers.attention.dsv4.dense_prefill_indexer import ( - DenseCandidateIndexer, - ) - prefill_dense = None - if get_parallel().attn_cp_size > 1: - prefill_dense = DenseCandidateIndexer(topk_blocks, block_size) - return DeepGemmCandidateIndexer( - topk_blocks, block_size, prefill_dense=prefill_dense - ) + return DeepGemmCandidateIndexer(topk_blocks, block_size) # TODO(candidate): Hopper decode and the torch prefill path still select through diff --git a/python/sglang/srt/layers/attention/dsv4/candidate_indexer_deep_gemm.py b/python/sglang/srt/layers/attention/dsv4/candidate_indexer_deep_gemm.py index afa2d3d40450..1562b76f3528 100644 --- a/python/sglang/srt/layers/attention/dsv4/candidate_indexer_deep_gemm.py +++ b/python/sglang/srt/layers/attention/dsv4/candidate_indexer_deep_gemm.py @@ -24,7 +24,6 @@ expand_index_page_table, ) from sglang.srt.layers.attention.dsv4.dense_prefill_indexer import ( - DenseCandidateIndexer, score_tiles, ) from sglang.srt.layers.attention.dsv4.indexer import ( @@ -158,20 +157,13 @@ class PrefillSparseBlockTable(SparseBlockTable): # TODO(dark): support fusion of publish + topk of publish layer class DeepGemmCandidateIndexer(CandidateIndexer): - def __init__( - self, - topk_blocks: int, - block_size: int, - prefill_dense: Optional[DenseCandidateIndexer] = None, - ): + def __init__(self, topk_blocks: int, block_size: int): assert block_size == CANDIDATE_BLOCK_SIZE, block_size self.topk_blocks = topk_blocks self.block_size = block_size self.alt_stream = torch.cuda.Stream() self._row_ids: Optional[torch.Tensor] = None self._retired_row_ids: list = [] # captured graphs keep reading the buffers they saw - # set for the forwards whose prefill rows the page table does not describe - self._prefill_dense = prefill_dense def _request_ids( self, request_ids: Optional[torch.Tensor], rows: int, device: torch.device @@ -288,8 +280,6 @@ def publish_prefill( """The source layer's own top-k and the block table of the chunk from one pass over its dense scores, tile by tile: per tile the plain top-k into ``out_positions`` and one read for the block keys.""" - if self._prefill_dense is not None: - return self._prefill_dense.publish_prefill(inputs, out_positions) rows = inputs.num_rows device = inputs.q_fp4.device nblocks, valid_lens = candidate_row_lens(inputs.compress_lens, self.topk_blocks) @@ -380,8 +370,6 @@ def _prefill_table( def prefill_tail( self, published: CandidateMetadata, tail_lens: List[int] ) -> CandidateMetadata: - if self._prefill_dense is not None: - return self._prefill_dense.prefill_tail(published, tail_lens) table = published assert isinstance(table, PrefillSparseBlockTable), "prefill block table missing" rows, start = [], 0 @@ -408,8 +396,6 @@ def select_prefill( inputs: PrefillIndexerInputs, out_positions: torch.Tensor, ) -> None: - if self._prefill_dense is not None: - return self._prefill_dense.select_prefill(published, inputs, out_positions) table = published assert isinstance(table, PrefillSparseBlockTable), "prefill block table missing" rows, heads = inputs.q_sf.shape diff --git a/python/sglang/srt/layers/attention/dsv4/dense_prefill_indexer.py b/python/sglang/srt/layers/attention/dsv4/dense_prefill_indexer.py index d4a478a24de1..c207c7812635 100644 --- a/python/sglang/srt/layers/attention/dsv4/dense_prefill_indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/dense_prefill_indexer.py @@ -174,7 +174,7 @@ class DenseCandidateIndexer(CandidateIndexer): """Candidates as block ids per request (``PrefillCandidateBlocks``): the source keeps its best blocks from its dense scores, a consumer masks its own dense scores to -inf outside them and runs the plain top-k, tile by tile. - The prefill implementation for the CP layout; Hopper still runs the same + Used as a reference for sparse prefill; Hopper still runs the same selection inline in the backend.""" def __init__(self, topk_blocks: int, block_size: int): diff --git a/test/registered/kernels/ops/attention/test_dsv41_prefill_sparse_indexer.py b/test/registered/kernels/ops/attention/test_dsv41_prefill_sparse_indexer.py index 96d146a3cf18..14017948c3ca 100644 --- a/test/registered/kernels/ops/attention/test_dsv41_prefill_sparse_indexer.py +++ b/test/registered/kernels/ops/attention/test_dsv41_prefill_sparse_indexer.py @@ -5,16 +5,19 @@ from typing import NamedTuple import msgspec +import pytest import torch from sglang.kernels.ops.attention.dsv4.fp4_indexer import quantize_fp4_indexer_tensor from sglang.srt.layers.attention.dsv4.candidate_indexer import ( PrefillIndexerInputs, + make_candidate_indexer, select_candidate_blocks, ) from sglang.srt.layers.attention.dsv4.dense_prefill_indexer import ( DenseCandidateIndexer, ) +from sglang.srt.runtime_context import get_parallel from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.test_utils import CustomTestCase @@ -225,5 +228,209 @@ def test_prefill_tail_rebuilds_the_last_rows(self): self.assertGreaterEqual(len(a & b), MIN_TAIL_OVERLAP * len(a), r) +def _cp_indexer(cp_size=1): + if torch.cuda.get_device_capability()[0] < 10: + pytest.skip("DeepGEMM paged sparse MQA logits need SM100") + with get_parallel().override(attn_cp_size=cp_size): + return make_candidate_indexer(TOPK_BLOCKS, BLOCK) + + +def _take_prefill_rows(inputs, rows, counts): + return msgspec.structs.replace( + inputs, + q_fp4=inputs.q_fp4[rows].contiguous(), + q_sf=inputs.q_sf[rows].contiguous(), + weights=inputs.weights[rows].contiguous(), + compress_lens=inputs.compress_lens[rows].contiguous(), + request_starts=inputs.request_starts[rows].contiguous(), + kv_page_table=inputs.kv_page_table[rows].contiguous(), + rows_per_request=counts, + ) + + +def _cp_case(ratio): + counts, contexts = [1, 47, 48], [256, 20224, 17920] + cases = [ + make_case((n + 3) // 4 * 4, ctx, seed=ctx).inputs + for n, ctx in zip(counts, contexts) + ] + kv_page = 2 * PAGE + pages_per_request = [ctx * ratio // kv_page for ctx in contexts] + physical_pages = torch.randperm(sum(pages_per_request), device="cuda") + packed = torch.cat([c.k_cache for c in cases]) + cache = torch.empty_like(packed) + cache.view(-1, 2 // ratio, PAGE, 1, 68)[physical_pages] = packed.view( + -1, 2 // ratio, PAGE, 1, 68 + ) + req_to_token = torch.zeros( + 3, max(contexts) * ratio, dtype=torch.int64, device="cuda" + ) + page_table = torch.zeros( + 3, max(pages_per_request), dtype=torch.int32, device="cuda" + ) + for request, pages in enumerate(physical_pages.split(pages_per_request)): + page_table[request, : pages.numel()] = pages + slots = pages[:, None] * kv_page + torch.arange(kv_page, device="cuda") + req_to_token[request, : slots.numel()] = slots.flatten() + request_ids = torch.repeat_interleave( + torch.arange(3, device="cuda"), torch.tensor(counts, device="cuda") + ) + seq_lens = torch.cat( + [ + torch.arange(ctx * ratio - n + 1, ctx * ratio + 1, device="cuda") + for n, ctx in zip(counts, contexts) + ] + ).int() + inputs = PrefillIndexerInputs( + q_fp4=torch.cat([c.q_fp4[:n] for c, n in zip(cases, counts)]), + q_sf=torch.cat([c.q_sf[:n] for c, n in zip(cases, counts)]), + weights=torch.cat([c.weights[:n] for c, n in zip(cases, counts)]), + compress_lens=seq_lens // ratio, + request_starts=torch.tensor( + [0, contexts[0], sum(contexts[:2])], device="cuda", dtype=torch.int32 + )[request_ids], + lens_per_request=contexts, + rows_per_request=counts, + kv=tuple(torch.cat([c.kv[i] for c in cases]) for i in range(2)), + k_cache=cache, + page_size=PAGE, + kv_page_table=page_table[request_ids], + kv_page_size=kv_page, + compress_ratio=ratio, + ) + return inputs, req_to_token, request_ids + + +@pytest.mark.parametrize("ratio,cp_size", [(1, 2), (2, 4)]) +@torch.inference_mode() +def test_cp_prefill_matches_dense_with_shuffled_pages(ratio, cp_size): + """Sparse selection on CP-local rows must match dense scores on one GPU.""" + indexer, dense = _cp_indexer(cp_size), DenseCandidateIndexer(TOPK_BLOCKS, BLOCK) + inputs, req_to_token, request_ids = _cp_case(ratio) + for rank in range(cp_size): + rows = torch.arange(rank, inputs.num_rows, cp_size, device="cuda") + ids = request_ids[rows] + counts = torch.bincount(ids, minlength=3).tolist() + local = _take_prefill_rows(inputs, rows, counts) + table, own = publish(indexer, local) + reference, own_dense = publish(dense, local) + assert torch.equal(own.sort().values, own_dense.sort().values) + consumer = msgspec.structs.replace( + local, + q_fp4=local.q_fp4.roll(1, 0), + q_sf=local.q_sf.roll(1, 0), + weights=local.weights.roll(1, 0), + ) + scores = torch.empty(0, device="cuda") + tail_counts = [0, min(4, counts[1]), min(4, counts[2])] + ends = torch.tensor(counts, device="cuda").cumsum(0).tolist() + tail_rows = torch.cat( + [ + torch.arange(end - n, end, device="cuda") + for end, n in zip(ends, tail_counts) + if n + ] + ) + for tail, selected_rows, lengths in ( + (False, torch.arange(local.num_rows, device="cuda"), counts), + (True, tail_rows, tail_counts), + ): + sparse_table = indexer.prefill_tail(table, lengths) if tail else table + dense_table = dense.prefill_tail(reference, lengths) if tail else reference + selected = _take_prefill_rows(consumer, selected_rows, lengths) + positions = select(indexer, sparse_table, selected) + expected = select(dense, dense_table, selected) + for row, source_row in enumerate(selected_rows.tolist()): + start = selected.request_starts[row] + got = positions[row][positions[row] >= 0].long() - start + want = expected[row][expected[row] >= 0].long() - start + assert ( + got.numel() + == want.numel() + == min(TOPK, selected.compress_lens[row].item()) + ) + assert ((got >= 0) & (got < selected.compress_lens[row])).all() + assert got.unique().numel() == got.numel() + assert ( + len(set(got.tolist()) & set(want.tolist())) + >= 0.95 * want.numel() + ) + blocks = sparse_table.blocks[row] + columns = torch.searchsorted(blocks, got // BLOCK) + assert torch.equal(blocks[columns].long(), got // BLOCK) + slots = ( + sparse_table.phys_blocks[row, columns].long() * BLOCK + got % BLOCK + ) + assert torch.equal( + slots, req_to_token[ids[source_row], got * ratio] // ratio + ) + + +@pytest.mark.parametrize("rank", [0, 1]) +@torch.inference_mode() +def test_cp_signed_prefill_selects_topk_of_consumed_logits(rank, monkeypatch): + """Check signed BF16 selection and CP page mapping using consumed logits.""" + from sglang.srt.layers.attention.dsv4 import candidate_indexer_deep_gemm as sparse + + indexer = _cp_indexer(4) + inputs, req_to_token, request_ids = _cp_case(2) + rows = torch.arange(rank, inputs.num_rows, 4, device="cuda") + ids = request_ids[rows] + counts = torch.bincount(ids, minlength=3).tolist() + local = _take_prefill_rows(inputs, rows, counts) + local = msgspec.structs.replace( + local, weights=(2 * local.weights - 1).bfloat16().float() + ) + table, own = publish(indexer, local) + reference, own_dense = publish(DenseCandidateIndexer(TOPK_BLOCKS, BLOCK), local) + assert torch.equal(own.sort().values, own_dense.sort().values) + source_blocks = [row for blocks in reference.request_blocks for row in blocks] + consumer = msgspec.structs.replace( + local, + q_fp4=local.q_fp4.roll(1, 0), + q_sf=local.q_sf.roll(1, 0), + weights=local.weights.roll(1, 0), + ) + captured = [] + original = sparse.sparse_logits + + def capture_logits(*args, **kwargs): + logits = original(*args, **kwargs) + captured.append(logits) + return logits + + monkeypatch.setattr(sparse, "sparse_logits", capture_logits) + positions = select(indexer, table, consumer) + assert len(captured) == 1 and captured[0].dtype == torch.bfloat16 + for row, length in enumerate(local.compress_lens.tolist()): + blocks = table.blocks[row] + want_blocks = source_blocks[row] + assert torch.equal( + blocks[blocks < (length + BLOCK - 1) // BLOCK], + want_blocks[want_blocks >= 0].sort().values, + ) + valid = positions[row] >= 0 + assert (positions[row, ~valid] == -1).all() + got = positions[row, valid].long() - local.request_starts[row] + assert got.numel() == got.unique().numel() == min(TOPK, length) + assert ((got >= 0) & (got < length)).all() + columns = torch.searchsorted(blocks, got // BLOCK) + assert (columns < blocks.numel()).all() + assert torch.equal(blocks[columns].long(), got // BLOCK) + sparse_columns = columns * BLOCK + got % BLOCK + valid_length = int(table.valid_lens[row]) + assert (sparse_columns < valid_length).all() + scores = captured[0][row, :valid_length] + # Equal scores may choose different indices; compare their multisets. + torch.testing.assert_close( + scores[sparse_columns].sort().values, + torch.topk(scores, got.numel()).values.sort().values, + rtol=0, + atol=0, + ) + slots = table.phys_blocks[row, columns].long() * BLOCK + got % BLOCK + assert torch.equal(slots, req_to_token[ids[row], got * 2] // 2) + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py b/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py index ad06b05e8402..26cf78196fda 100644 --- a/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py +++ b/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py @@ -875,6 +875,192 @@ def run(rows_per_chunk): self.assertTrue(torch.equal(run(rows_per_chunk), expected)) +class TestCandidateIndexerCPPages(CustomTestCase): + @staticmethod + def _metadata(req_to_token, requests, lengths, num_tokens): + from sglang.kernels.ops.attention.dsv4_attn_metadata_kernels import ( + BuildPageTablePositions, + ) + from sglang.srt.layers.attention.deepseek_v4_backend import DSV4AttnMetadata + + prep = BuildPageTablePositions.execute( + req_to_token=req_to_token, + req_pool_indices_repeated=requests, + seq_lens_casual=lengths, + max_seq_len=int(lengths.max()), + page_size=256, + swa_window=128, + ) + return DSV4AttnMetadata( + page_size=256, + page_table=prep.page_table, + raw_out_loc=torch.zeros(num_tokens, dtype=torch.int32), + cuda_int32_kwargs={"dtype": torch.int32}, + seq_lens_casual=prep.seq_lens_casual, + positions_casual=prep.positions_casual, + swa_page_indices=torch.empty((lengths.numel(), 0), dtype=torch.int32), + swa_topk_lengths=prep.swa_topk_lengths, + index_topk=512, + present_ratios=(1, 2), + low_ratios=(1, 2), + ) + + def test_cp_pages_address_global_compressed_kv(self): + """Local queries, including compact tails, address the original request's KV.""" + from sglang.srt.layers.attention import deepseek_v4_backend as be + from sglang.srt.layers.attention.dsv4.candidate_indexer import ( + expand_index_page_table, + ) + from sglang.srt.layers.cp.interleave import interleave_rows_per_request + + extend_lens = [7, 1, 6] + requests = torch.tensor([3, 1, 4]) + positions = torch.cat( + [torch.arange(254, 261), torch.arange(1), torch.arange(509, 515)] + ) + global_requests = torch.repeat_interleave(requests, torch.tensor(extend_lens)) + pages = torch.tensor([[9, 2, 11], [4, 14, 1], [8, 5, 12]]) + req_to_token = torch.zeros((5, 768), dtype=torch.int64) + req_to_token[requests] = (pages[:, :, None] * 256 + torch.arange(256)).flatten( + 1 + ) + tail_indices = torch.tensor([4, 5, 6, 7, 11, 12, 13]) + forward_batch = SimpleNamespace(positions=positions, batch_size=3) + backend = be.DeepseekV4AttnBackend.__new__(be.DeepseekV4AttnBackend) + + for cp_size in (2, 4): + for rank, tail in product(range(cp_size), (False, True)): + parallel = SimpleNamespace(attn_cp_rank=rank, attn_cp_size=cp_size) + with ( + self.subTest(cp_size=cp_size, rank=rank, tail=tail), + patch.object(be, "get_parallel", return_value=parallel), + ): + metadata_rows = ( + tail_indices if tail else torch.arange(positions.numel()) + ) + local_rows = metadata_rows[metadata_rows % cp_size == rank] + if tail: + layout = backend._late_layer_tail_cp_layout( + forward_batch, tail_indices, torch.tensor([3, 1, 3]) + ) + cp_metadata = layout["cp_metadata"] + padded_rows = sum(cp_metadata.per_rank_actual_token) + local_index = cp_metadata.local_index + query_lens = layout["local_lens_cpu"] + else: + padded_rows = -(-positions.numel() // cp_size) * cp_size + local_index = None + query_lens = interleave_rows_per_request( + extend_lens, rank, cp_size + ) + padding = padded_rows - metadata_rows.numel() + metadata = self._metadata( + req_to_token, + torch.nn.functional.pad( + global_requests[metadata_rows], (0, padding) + ), + torch.nn.functional.pad( + positions[metadata_rows] + 1, (0, padding) + ), + metadata_rows.numel(), + ) + metadata.apply_cp_reindex( + num_tokens=metadata_rows.numel(), local_index=local_index + ) + local_requests = torch.repeat_interleave( + requests, torch.tensor(query_lens) + ) + torch.testing.assert_close( + local_requests, global_requests[local_rows] + ) + torch.testing.assert_close( + metadata.positions_casual[: local_rows.numel()].long(), + positions[local_rows], + ) + local_pages = metadata.page_table[: local_rows.numel()] + # DeepGEMM pairs adjacent rows of a request using one page table. + torch.testing.assert_close( + local_pages.long(), + torch.repeat_interleave(pages, torch.tensor(query_lens), dim=0), + ) + for ratio in (1, 2): + with self.subTest(ratio=ratio): + index_pages = expand_index_page_table( + local_pages, + full_page_size=256, + compress_ratio=ratio, + index_page_size=128, + ) + for row, global_row in enumerate(local_rows): + compressed = torch.arange( + (positions[global_row] + 1) // ratio + ) + slots = ( + index_pages[row, compressed // 128].long() * 128 + + compressed % 128 + ) + expected = ( + req_to_token[ + global_requests[global_row], compressed * ratio + ] + // ratio + ) + torch.testing.assert_close(slots, expected) + + def test_empty_cp_tail_clears_candidates_without_losing_full_publication(self): + """A rank with no tail queries must not build an empty sparse schedule.""" + from sglang.srt.layers.attention import deepseek_v4_backend as be + from sglang.srt.layers.attention.dsv4.candidate_indexer import ( + PrefillCandidateBlocks, + ) + from sglang.srt.layers.attention.dsv4.dense_prefill_indexer import ( + DenseCandidateIndexer, + ) + + positions, tail_indices = torch.arange(512), torch.arange(384, 512) + batch = SimpleNamespace(positions=positions, batch_size=1) + backend = be.DeepseekV4AttnBackend.__new__(be.DeepseekV4AttnBackend) + backend.candidate_indexer = DenseCandidateIndexer(2048, 8) + for rank in (0, 128): + with ( + self.subTest(rank=rank), + patch.object( + be, + "get_parallel", + return_value=SimpleNamespace(attn_cp_rank=rank, attn_cp_size=256), + ), + ): + layout = backend._late_layer_tail_cp_layout( + batch, tail_indices, torch.tensor([128]) + ) + tail = be.LateLayerTail( + token_indices=layout["local_token_indices"], + positions=layout["local_positions"], + extend_seq_lens=torch.tensor([128]), + extend_seq_lens_cpu=[128], + swa_out_cache_loc=tail_indices + 256, + pad_rows=layout["pad_rows"], + cp_metadata=layout["cp_metadata"], + local_lens_cpu=layout["local_lens_cpu"], + ) + published = PrefillCandidateBlocks( + request_blocks=[torch.tensor([[0], [1]], dtype=torch.int32)] + ) + backend.forward_metadata = be.DSV4Metadata(None, None) + backend.tail_forward_metadata = be.DSV4Metadata( + None, None, late_layer_tail=tail, candidate_metadata=published + ) + backend._publish_prefill(published) + self.assertIs(backend.forward_metadata.candidate_metadata, published) + result = backend.tail_forward_metadata.candidate_metadata + if rank == 0: + self.assertIsNone(result) + else: + torch.testing.assert_close( + result.request_blocks[0], published.request_blocks[0][-1:] + ) + + class TestCandidateIndexerGating(CustomTestCase): def test_candidate_indexer_gating(self): from sglang.srt.layers.attention.dsv4 import candidate_indexer