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()