Files
sglang/test/registered/unit/layers/test_nsa_topk_transform.py
laoyao0822 f75ffff8d9 Protect CP shared-KV cache-hit correctness under batched FP8 reuse
Cache-hit GSM8K regressions only appeared after the second pass reused request-specific suffix pages, so this change adds fail-fast transfer validation, masks stale rectangular page-table tails, and extends CUDA/unit coverage across FP8 CP shared-KV write, load, top-k, and materialization paths. The temporary ledger records eliminated hypotheses to prevent re-debugging the same L2 and persistent-cache paths.\n\nConstraint: CP shared KV stores physical pages but scheduler-visible semantics must remain valid-token/page-bounded.\nConstraint: bs>1 FP8 prefill must preserve existing CP shared-KV fast paths without silent fallback.\nRejected: Blame raw HiCache L2 load without tests | L2 KV and index backup/load/materialize roundtrips pass on remote CUDA.\nRejected: Disable current/partial reuse broadly | hides the cache-hit contract regression and costs performance.\nConfidence: medium\nScope-risk: moderate\nDirective: Do not weaken CP shared-KV fail-fast or rectangular-tail masking without rerunning second-pass cache-hit accuracy tests.\nTested: remote CUDA pytest for fused FP8 MLA store, fused persistent index store, L2-loaded FP8 KV materialize, L2-loaded index materialize, ragged top-k offset, TAI batched index MQA prepare.\nTested: local py_compile for touched test files and git diff --check.\nNot-tested: full second-pass GSM8K accuracy after these diagnostic tests; root cause remains under investigation.
2026-06-07 13:26:49 +08:00

242 lines
9.5 KiB
Python

import unittest
from unittest.mock import patch
import torch
from sglang.srt.layers.attention.nsa_backend import (
NSAIndexerMetadata,
NSAMetadata,
TopkTransformMethod,
)
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=1, suite="stage-a-test-cpu")
class TestNSATopkTransform(unittest.TestCase):
@unittest.skipIf(not torch.cuda.is_available(), "CUDA is required")
def test_ragged_topk_transform_offsets_are_request_relative_with_row_starts(self):
from sgl_kernel import fast_topk_transform_ragged_fused
# CP bs>1 compacts each request segment into a temporary K buffer.
# The selected compact column must be translated back to the normal
# request-concatenated ragged layout as:
# request_base + (selected_compact_col - row_start)
# not request_base + selected_compact_col.
columns = 5000
logits = torch.zeros((2, columns), device="cuda", dtype=torch.float32)
logits[0, 2999] = 10.0
logits[1, 1000 + 2999] = 10.0
lengths = torch.tensor([3000, 3000], device="cuda", dtype=torch.int32)
row_starts = torch.tensor([0, 1000], device="cuda", dtype=torch.int32)
request_offsets = torch.tensor(
[0, 100000], device="cuda", dtype=torch.int32
)
out = fast_topk_transform_ragged_fused(
score=logits,
lengths=lengths,
topk_indices_offset=request_offsets,
topk=2048,
row_starts=row_starts,
)
torch.cuda.synchronize()
self.assertEqual(int(out[0, 0].item()), 2999)
self.assertEqual(int(out[1, 0].item()), 100000 + 2999)
def test_paged_topk_transform_raises_when_fused_output_is_not_from_page_table(self):
page_table = torch.tensor(
[
[10, 11, 12, 13, 14, 15, 16],
[20, 21, 22, 23, 24, 25, 26],
],
dtype=torch.int32,
)
lengths_seen = {}
def fake_fast_topk_transform_fused(**kwargs):
lengths_seen["value"] = kwargs["lengths"].clone()
return torch.tensor(
[
[10, 1_039_799_618, 11, 12],
[20, 21, 22, 23],
],
dtype=torch.int32,
)
metadata = NSAMetadata(
page_size=1,
cache_seqlens_int32=torch.tensor([7], dtype=torch.int32),
max_seq_len_q=1,
max_seq_len_k=7,
cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32),
cu_seqlens_k=torch.tensor([0, 7], dtype=torch.int32),
page_table_1=page_table,
real_page_table=page_table,
nsa_cache_seqlens_int32=torch.tensor([4, 7], dtype=torch.int32),
nsa_cu_seqlens_q=torch.tensor([0, 1, 2], dtype=torch.int32),
nsa_cu_seqlens_k=torch.tensor([0, 4, 11], dtype=torch.int32),
nsa_extend_seq_lens_list=[2],
nsa_seqlens_expanded=torch.tensor([4, 7], dtype=torch.int32),
)
indexer_metadata = NSAIndexerMetadata(
attn_metadata=metadata,
topk_transform_method=TopkTransformMethod.PAGED,
validate_paged_topk=True,
)
with patch(
"sgl_kernel.fast_topk_transform_fused",
side_effect=fake_fast_topk_transform_fused,
), patch(
"sglang.srt.environ.envs.SGLANG_NSA_FUSE_TOPK.get", return_value=True
), patch(
"sglang.srt.environ.envs.SGLANG_DEBUG_CP_SHARED_KV.get", return_value=True
):
with self.assertRaisesRegex(
RuntimeError,
"NSA PAGED fused topk_transform produced values outside page_table_1",
):
indexer_metadata.topk_transform(
logits=torch.zeros((2, 7), dtype=torch.float32),
topk=4,
cu_seqlens_q=torch.tensor([1, 1], dtype=torch.int32),
)
self.assertEqual(lengths_seen["value"].tolist(), [4, 7])
def test_paged_topk_transform_rejects_lengths_exceeding_page_table_width(self):
page_table = torch.tensor([[10, 11, 12]], dtype=torch.int32)
metadata = NSAMetadata(
page_size=1,
cache_seqlens_int32=torch.tensor([3], dtype=torch.int32),
max_seq_len_q=1,
max_seq_len_k=3,
cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32),
cu_seqlens_k=torch.tensor([0, 3], dtype=torch.int32),
page_table_1=page_table,
real_page_table=page_table,
nsa_cache_seqlens_int32=torch.tensor([4], dtype=torch.int32),
nsa_cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32),
nsa_cu_seqlens_k=torch.tensor([0, 4], dtype=torch.int32),
nsa_extend_seq_lens_list=[1],
nsa_seqlens_expanded=torch.tensor([4], dtype=torch.int32),
)
indexer_metadata = NSAIndexerMetadata(
attn_metadata=metadata,
topk_transform_method=TopkTransformMethod.PAGED,
validate_paged_topk=True,
)
with patch(
"sgl_kernel.fast_topk_transform_fused",
side_effect=AssertionError("fused kernel should not be called"),
), patch(
"sglang.srt.environ.envs.SGLANG_NSA_FUSE_TOPK.get", return_value=True
), patch(
"sglang.srt.environ.envs.SGLANG_DEBUG_CP_SHARED_KV.get", return_value=True
):
with self.assertRaisesRegex(
RuntimeError,
"NSA PAGED fused topk lengths exceed page_table width",
):
indexer_metadata.topk_transform(
logits=torch.zeros((1, 4), dtype=torch.float32),
topk=4,
)
def test_paged_topk_transform_skips_validation_during_cuda_graph_capture(self):
page_table = torch.tensor([[10, 11, 12]], dtype=torch.int32)
def fake_fast_topk_transform_fused(**kwargs):
return torch.tensor([[1_039_799_618]], dtype=torch.int32)
metadata = NSAMetadata(
page_size=1,
cache_seqlens_int32=torch.tensor([3], dtype=torch.int32),
max_seq_len_q=1,
max_seq_len_k=3,
cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32),
cu_seqlens_k=torch.tensor([0, 3], dtype=torch.int32),
page_table_1=page_table,
real_page_table=page_table,
nsa_cache_seqlens_int32=torch.tensor([3], dtype=torch.int32),
nsa_cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32),
nsa_cu_seqlens_k=torch.tensor([0, 3], dtype=torch.int32),
nsa_extend_seq_lens_list=[1],
nsa_seqlens_expanded=torch.tensor([3], dtype=torch.int32),
)
indexer_metadata = NSAIndexerMetadata(
attn_metadata=metadata,
topk_transform_method=TopkTransformMethod.PAGED,
validate_paged_topk=True,
)
with patch(
"sgl_kernel.fast_topk_transform_fused",
side_effect=fake_fast_topk_transform_fused,
), patch(
"sglang.srt.layers.attention.nsa_backend._is_cuda_stream_capturing",
return_value=True,
), patch(
"sglang.srt.environ.envs.SGLANG_NSA_FUSE_TOPK.get", return_value=True
), patch(
"sglang.srt.environ.envs.SGLANG_DEBUG_CP_SHARED_KV.get", return_value=True
):
out = indexer_metadata.topk_transform(
logits=torch.zeros((1, 3), dtype=torch.float32),
topk=1,
)
self.assertEqual(out.tolist(), [[1_039_799_618]])
def test_paged_topk_transform_skips_validation_when_cp_shared_debug_disabled(self):
page_table = torch.tensor([[10, 11, 12]], dtype=torch.int32)
def fake_fast_topk_transform_fused(**kwargs):
return torch.tensor([[1_039_799_618]], dtype=torch.int32)
metadata = NSAMetadata(
page_size=1,
cache_seqlens_int32=torch.tensor([3], dtype=torch.int32),
max_seq_len_q=1,
max_seq_len_k=3,
cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32),
cu_seqlens_k=torch.tensor([0, 3], dtype=torch.int32),
page_table_1=page_table,
real_page_table=page_table,
nsa_cache_seqlens_int32=torch.tensor([3], dtype=torch.int32),
nsa_cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32),
nsa_cu_seqlens_k=torch.tensor([0, 3], dtype=torch.int32),
nsa_extend_seq_lens_list=[1],
nsa_seqlens_expanded=torch.tensor([3], dtype=torch.int32),
)
indexer_metadata = NSAIndexerMetadata(
attn_metadata=metadata,
topk_transform_method=TopkTransformMethod.PAGED,
validate_paged_topk=True,
)
with patch(
"sgl_kernel.fast_topk_transform_fused",
side_effect=fake_fast_topk_transform_fused,
), patch(
"sglang.srt.layers.attention.nsa_backend._is_cuda_stream_capturing",
return_value=False,
), patch(
"sglang.srt.environ.envs.SGLANG_NSA_FUSE_TOPK.get", return_value=True
), patch(
"sglang.srt.environ.envs.SGLANG_DEBUG_CP_SHARED_KV.get", return_value=False
):
out = indexer_metadata.topk_transform(
logits=torch.zeros((1, 3), dtype=torch.float32),
topk=1,
)
self.assertEqual(out.tolist(), [[1_039_799_618]])
if __name__ == "__main__":
unittest.main()