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.
242 lines
9.5 KiB
Python
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()
|