Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions python/sglang/srt/layers/attention/deepseek_v4_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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:
Expand Down
18 changes: 4 additions & 14 deletions python/sglang/srt/layers/attention/dsv4/candidate_indexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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 (
Expand All @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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()
Loading
Loading