import unittest import sys from types import SimpleNamespace from unittest.mock import patch from sglang.srt.layers.attention.nsa.utils import ( NSAContextParallelMetadata, PageAlignedCacheExtent, build_page_aligned_cache_extent, _get_in_seq_last_token_owner_and_offset, build_page_aligned_in_seq_split_list, build_token_balanced_in_seq_split_list, can_cp_split, cp_split_and_rebuild_1d, get_cp_shared_kv_local_out_cache_loc, get_cp_shared_kv_local_physical_out_cache_loc, get_cp_local_embedding_padded_token_count, pad_cp_local_input_ids_for_embedding, split_in_seq_cp_local_pair, ) from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=1, suite="stage-a-test-cpu") class TestPageAlignedCacheExtent(unittest.TestCase): def test_extent_uses_page_boundary_not_cp_size(self): extent = build_page_aligned_cache_extent(valid_tokens=100, page_size=64) self.assertEqual(extent.valid_tokens, 100) self.assertEqual(extent.padded_pages, 2) self.assertEqual(extent.padded_tokens, 128) self.assertEqual(extent.padding_tokens, 28) def test_extent_handles_empty_and_aligned_lengths(self): self.assertEqual( build_page_aligned_cache_extent(valid_tokens=0, page_size=64), PageAlignedCacheExtent( valid_tokens=0, padded_pages=0, padded_tokens=0, padding_tokens=0, ), ) self.assertEqual( build_page_aligned_cache_extent(valid_tokens=128, page_size=64), PageAlignedCacheExtent( valid_tokens=128, padded_pages=2, padded_tokens=128, padding_tokens=0, ), ) class TestNSAInSeqCPUtils(unittest.TestCase): def test_contiguous_valid_cp_query_count(self): from sglang.srt.layers.attention.nsa.nsa_indexer import ( _compute_contiguous_valid_cp_query_count, ) self.assertEqual( _compute_contiguous_valid_cp_query_count( cp_kv_end=1024, actual_seq_q=128, logical_kv_limit=1024, ), 128, ) self.assertEqual( _compute_contiguous_valid_cp_query_count( cp_kv_end=1100, actual_seq_q=128, logical_kv_limit=1024, ), 52, ) self.assertEqual( _compute_contiguous_valid_cp_query_count( cp_kv_end=1100, actual_seq_q=64, logical_kv_limit=1000, ), 0, ) self.assertEqual( _compute_contiguous_valid_cp_query_count( cp_kv_end=100, actual_seq_q=0, logical_kv_limit=100, ), 0, ) def assert_page_aligned_boundaries( self, split_list, *, extend_prefix_len, extend_len, page_size ): cursor = 0 for segment_len in split_list[:-1]: cursor += segment_len if cursor < extend_len: self.assertEqual((extend_prefix_len + cursor) % page_size, 0) def test_page_aligned_split_keeps_boundaries_on_pages(self): split_list, info = build_page_aligned_in_seq_split_list( total_len=32768, extend_len=32768, extend_prefix_len=0, page_size=64, cp_size=8, ) self.assertTrue(info.page_aligned) self.assertEqual(sum(split_list), 32768) self.assertEqual(len(split_list), 16) self.assertTrue(all(segment_len > 0 for segment_len in split_list)) self.assert_page_aligned_boundaries( split_list, extend_prefix_len=0, extend_len=32768, page_size=64 ) def test_page_aligned_split_uses_prefix_for_boundary_alignment(self): split_list, info = build_page_aligned_in_seq_split_list( total_len=1024, extend_len=1024, extend_prefix_len=128, page_size=64, cp_size=8, ) self.assertTrue(info.page_aligned) self.assertEqual(sum(split_list), 1024) self.assert_page_aligned_boundaries( split_list, extend_prefix_len=128, extend_len=1024, page_size=64 ) def test_page_aligned_split_keeps_tail_partial_page_unsplit(self): split_list, info = build_page_aligned_in_seq_split_list( total_len=1100, extend_len=1100, extend_prefix_len=0, page_size=64, cp_size=8, ) self.assertTrue(info.page_aligned) self.assertEqual(sum(split_list), 1100) self.assertEqual(split_list[-1], 12) self.assert_page_aligned_boundaries( split_list, extend_prefix_len=0, extend_len=1100, page_size=64 ) def test_page_aligned_split_exposes_padded_extent_without_padding_split_list(self): split_list, info = build_page_aligned_in_seq_split_list( total_len=100, extend_len=100, extend_prefix_len=0, page_size=64, cp_size=8, ) self.assertTrue(info.page_aligned) self.assertEqual(sum(split_list), 100) self.assertEqual(split_list[:2], [64, 36]) self.assertEqual(split_list[2:], [0] * 14) self.assertEqual(info.extend_valid_tokens, 100) self.assertEqual(info.extend_padded_pages, 2) self.assertEqual(info.extend_padded_tokens, 128) self.assertEqual(info.extend_padding_tokens, 28) def test_page_aligned_split_falls_back_when_prefix_is_not_page_aligned(self): split_list, info = build_page_aligned_in_seq_split_list( total_len=1024, extend_len=1024, extend_prefix_len=1, page_size=64, cp_size=8, ) self.assertFalse(info.page_aligned) self.assertEqual(split_list, build_token_balanced_in_seq_split_list(1024, 8)) def test_page_aligned_split_pads_zero_segments_when_page_units_are_short(self): split_list, info = build_page_aligned_in_seq_split_list( total_len=512, extend_len=512, extend_prefix_len=0, page_size=64, cp_size=8, ) self.assertTrue(info.page_aligned) self.assertEqual(sum(split_list), 512) self.assertEqual(split_list[:8], [64] * 8) self.assertEqual(split_list[8:], [0] * 8) self.assert_page_aligned_boundaries( split_list, extend_prefix_len=0, extend_len=512, page_size=64 ) def test_page_aligned_split_allows_radix_hit_suffix_with_one_page_per_rank(self): split_list, info = build_page_aligned_in_seq_split_list( total_len=512, extend_len=512, extend_prefix_len=54464, page_size=64, cp_size=8, ) self.assertTrue(info.page_aligned) self.assertEqual(sum(split_list), 512) self.assertEqual(split_list[:8], [64] * 8) self.assertEqual(split_list[8:], [0] * 8) self.assert_page_aligned_boundaries( split_list, extend_prefix_len=54464, extend_len=512, page_size=64 ) def test_page_aligned_split_keeps_short_radix_hit_suffix_page_aligned(self): split_list, info = build_page_aligned_in_seq_split_list( total_len=256, extend_len=256, extend_prefix_len=54464, page_size=64, cp_size=8, ) self.assertTrue(info.page_aligned) self.assertEqual(sum(split_list), 256) self.assertEqual(split_list[:4], [64] * 4) self.assertEqual(split_list[4:], [0] * 12) self.assert_page_aligned_boundaries( split_list, extend_prefix_len=54464, extend_len=256, page_size=64 ) def test_can_cp_split_keeps_cp_for_short_radix_hit_suffix(self): class Mode: def is_context_parallel_extend(self): return True forward_batch = SimpleNamespace( uses_cp_shared_kv=True, extend_seq_lens_cpu=[256], extend_prefix_lens_cpu=[54464], token_to_kv_pool=SimpleNamespace(page_size=64), forward_mode=Mode(), ) with ( patch( "sglang.srt.layers.attention.nsa.utils.is_nsa_prefill_cp_round_robin_split", return_value=False, ), patch( "sglang.srt.layers.attention.nsa.utils.is_nsa_enable_prefill_cp", return_value=True, ), ): self.assertTrue(can_cp_split(256, 8, True, forward_batch)) def test_can_cp_split_keeps_cp_for_radix_hit_suffix_with_one_page_per_rank(self): class Mode: def is_context_parallel_extend(self): return True forward_batch = SimpleNamespace( uses_cp_shared_kv=True, extend_seq_lens_cpu=[512], extend_prefix_lens_cpu=[54464], token_to_kv_pool=SimpleNamespace(page_size=64), forward_mode=Mode(), ) with ( patch( "sglang.srt.layers.attention.nsa.utils.is_nsa_prefill_cp_round_robin_split", return_value=False, ), patch( "sglang.srt.layers.attention.nsa.utils.is_nsa_enable_prefill_cp", return_value=True, ), ): self.assertTrue(can_cp_split(512, 8, True, forward_batch)) def test_page_aligned_split_adds_padding_tokens_to_last_segment(self): split_list, info = build_page_aligned_in_seq_split_list( total_len=1040, extend_len=1024, extend_prefix_len=0, page_size=64, cp_size=8, ) self.assertTrue(info.page_aligned) self.assertEqual(sum(split_list), 1040) self.assertEqual(split_list[-1], 80) self.assert_page_aligned_boundaries( split_list, extend_prefix_len=0, extend_len=1024, page_size=64 ) def test_last_token_owner_uses_actual_token_count_when_batch_is_padded(self): # Padded prefill can have 64 model tokens while the real prompt has only # 11 tokens. In in-seq split with cp_size=8, the real last token is in # segment 2, not in rank 0's trailing padded segment. split_list = [4] * 16 owner, local_offset = _get_in_seq_last_token_owner_and_offset( split_list=split_list, cp_size=8, actual_token_count=11, ) self.assertEqual(owner, 2) self.assertEqual(local_offset, 2) def test_last_token_owner_keeps_existing_unpadded_fast_path_location(self): split_list = [4] * 16 owner, local_offset = _get_in_seq_last_token_owner_and_offset( split_list=split_list, cp_size=8, actual_token_count=64, ) self.assertEqual(owner, 0) self.assertEqual(local_offset, 7) def test_local_pair_split_uses_metadata_lengths_not_half_split(self): import torch tensor = torch.arange(9) prev, next_ = split_in_seq_cp_local_pair(tensor, 6, 3) self.assertEqual(prev.tolist(), [0, 1, 2, 3, 4, 5]) self.assertEqual(next_.tolist(), [6, 7, 8]) def test_local_pair_split_rejects_stale_metadata(self): import torch with self.assertRaisesRegex(RuntimeError, "local in-seq CP length mismatch"): split_in_seq_cp_local_pair(torch.arange(9), 5, 5, name="q_fp8") def test_cp_split_and_rebuild_1d_matches_in_seq_zigzag_order(self): import torch from types import SimpleNamespace forward_batch = SimpleNamespace( nsa_cp_metadata=NSAContextParallelMetadata( split_list=[2, 2, 2, 2, 2, 2, 2, 2], zigzag_index=[1, 6], ) ) local_locs = cp_split_and_rebuild_1d(forward_batch, torch.arange(16)) self.assertEqual(local_locs.tolist(), [2, 3, 12, 13]) def test_cp_local_embedding_pad_len_uses_metadata_max_rank_len(self): from types import SimpleNamespace import torch forward_batch = SimpleNamespace( nsa_cp_metadata=NSAContextParallelMetadata(max_rank_len=[4096] * 8) ) self.assertEqual( get_cp_local_embedding_padded_token_count(forward_batch, 4040), 4096 ) self.assertEqual( get_cp_local_embedding_padded_token_count(forward_batch, 4096), 4096 ) self.assertEqual( pad_cp_local_input_ids_for_embedding( SimpleNamespace( nsa_cp_metadata=NSAContextParallelMetadata(max_rank_len=[6] * 8) ), torch.tensor([11, 12, 13, 14]), ).tolist(), [11, 12, 13, 14, 0, 0], ) self.assertEqual( pad_cp_local_input_ids_for_embedding( SimpleNamespace( nsa_cp_metadata=NSAContextParallelMetadata(max_rank_len=[4] * 8) ), torch.tensor([11, 12, 13, 14]), ).tolist(), [11, 12, 13, 14], ) missing_metadata = SimpleNamespace(nsa_cp_metadata=None) self.assertIsNone( get_cp_local_embedding_padded_token_count(missing_metadata, 4040) ) self.assertIsNone( pad_cp_local_input_ids_for_embedding( missing_metadata, torch.tensor([11, 12, 13, 14]) ) ) stale_metadata = SimpleNamespace( nsa_cp_metadata=NSAContextParallelMetadata(max_rank_len=[4039] * 8) ) self.assertIsNone( get_cp_local_embedding_padded_token_count(stale_metadata, 4040) ) def test_local_out_cache_loc_requires_compute_owner_pages(self): import torch from types import SimpleNamespace page_size = 4 # Segment order for cp_size=4, cp_rank=1 is segment 1 then 6. # The logical page ids below deliberately encode the same owners through # (logical_page - 1) % cp_size: # segment 1 -> page 2 owner 1 # segment 6 -> page 6 owner 1 segment_pages = [1, 2, 3, 4, 8, 7, 6, 5] out_cache_loc = torch.cat( [ torch.arange(page * page_size, (page + 1) * page_size) for page in segment_pages ] ) forward_batch = SimpleNamespace( uses_cp_shared_kv=True, cp_shared_kv_layout=CpSharedKVLayout( page_size=page_size, cp_size=4, cp_rank=1, ), nsa_cp_metadata=NSAContextParallelMetadata( split_list=[page_size] * 8, zigzag_index=[1, 6], page_aligned=True, page_size=page_size, extend_prefix_len=0, ), out_cache_loc=out_cache_loc, ) local_locs = get_cp_shared_kv_local_out_cache_loc(forward_batch) self.assertIsNotNone(local_locs) self.assertEqual( local_locs.tolist(), list(range(2 * page_size, 3 * page_size)) + list(range(6 * page_size, 7 * page_size)), ) def test_local_physical_out_cache_loc_is_cached(self): import torch from types import SimpleNamespace page_size = 4 segment_pages = [1, 2, 3, 4, 8, 7, 6, 5] out_cache_loc = torch.cat( [ torch.arange(page * page_size, (page + 1) * page_size) for page in segment_pages ] ) forward_batch = SimpleNamespace( uses_cp_shared_kv=True, cp_shared_kv_layout=CpSharedKVLayout( page_size=page_size, cp_size=4, cp_rank=1, ), nsa_cp_metadata=NSAContextParallelMetadata( split_list=[page_size] * 8, zigzag_index=[1, 6], page_aligned=True, page_size=page_size, extend_prefix_len=0, ), out_cache_loc=out_cache_loc, ) physical_locs = get_cp_shared_kv_local_physical_out_cache_loc(forward_batch) second_read = get_cp_shared_kv_local_physical_out_cache_loc(forward_batch) self.assertIs(physical_locs, second_read) self.assertEqual( physical_locs.tolist(), list(range(1 * page_size, 2 * page_size)) + list(range(2 * page_size, 3 * page_size)), ) def test_local_out_cache_loc_falls_back_when_owner_mismatch(self): import torch from types import SimpleNamespace page_size = 4 out_cache_loc = torch.arange(page_size * 8, page_size * 16) forward_batch = SimpleNamespace( uses_cp_shared_kv=True, cp_shared_kv_layout=CpSharedKVLayout( page_size=page_size, cp_size=4, cp_rank=1, ), nsa_cp_metadata=NSAContextParallelMetadata( split_list=[page_size] * 8, zigzag_index=[1, 6], page_aligned=True, page_size=page_size, extend_prefix_len=0, ), out_cache_loc=out_cache_loc, ) self.assertIsNone(get_cp_shared_kv_local_out_cache_loc(forward_batch)) def test_local_out_cache_loc_logs_every_fallback_event(self): import torch from types import SimpleNamespace from sglang.srt.layers.attention.nsa import utils as nsa_utils page_size = 4 forward_batch = SimpleNamespace( uses_cp_shared_kv=True, cp_shared_kv_layout=CpSharedKVLayout( page_size=page_size, cp_size=4, cp_rank=1, ), nsa_cp_metadata=NSAContextParallelMetadata( split_list=[page_size] * 8, zigzag_index=[1, 6], page_aligned=False, page_size=page_size, extend_prefix_len=0, ), out_cache_loc=torch.arange(page_size * 8, page_size * 16), ) with self.assertLogs( "sglang.srt.layers.attention.nsa.utils", level="WARNING" ) as cm: self.assertIsNone(get_cp_shared_kv_local_out_cache_loc(forward_batch)) self.assertIsNone(get_cp_shared_kv_local_out_cache_loc(forward_batch)) self.assertEqual(len(cm.output), 2) self.assertIn("[CP_SHARED_KV_FALLBACK][direct_write]", cm.output[0]) self.assertIn("metadata is not page-aligned", cm.output[0]) self.assertIn("[CP_SHARED_KV_FALLBACK][direct_write]", cm.output[1]) self.assertIn("metadata is not page-aligned", cm.output[1]) def test_indexer_direct_write_does_not_log_missing_metadata_for_non_cp_batch(self): import torch from types import SimpleNamespace from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer indexer = object.__new__(Indexer) indexer.nsa_enable_prefill_cp = True forward_batch = SimpleNamespace( uses_cp_shared_kv=True, nsa_cp_metadata=None, ) with self.assertNoLogs( "sglang.srt.layers.attention.nsa.utils", level="INFO" ): stored = Indexer._store_cp_shared_local_index_k_cache( indexer, forward_batch, layer_id=0, local_key=torch.empty(0), act_quant=None, ) self.assertFalse(stored) def test_indexer_in_seq_cp_pair_materializes_index_once_for_prev_next(self): import torch from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer indexer = object.__new__(Indexer) logical_pages = torch.tensor([[1, 2, 3, 4]], dtype=torch.int32) materialized_index = torch.tensor([11], dtype=torch.int32) dense_pages = torch.tensor([[1, 2, 3, 4]], dtype=torch.int32) materialize_calls = [] topk_calls = [] class Metadata: def get_page_table_64(self): return logical_pages def fake_materialize(forward_batch, layer_id, logical_page_table): materialize_calls.append((layer_id, logical_page_table)) return materialized_index, dense_pages def fake_get_topk( forward_batch, layer_id, q_fp8, weights, metadata, kv_len, actual_seq_q, cp_index=None, current_index_kv=None, shared_index_buffer=None, shared_block_tables=None, actual_seq_q_tensor=None, actual_seq_q_cu_tensor=None, ): topk_calls.append( { "kv_len": kv_len, "actual_seq_q": actual_seq_q, "actual_seq_q_tensor": actual_seq_q_tensor, "actual_seq_q_cu_tensor": actual_seq_q_cu_tensor, "shared_index_buffer": shared_index_buffer, "shared_block_tables": shared_block_tables, "current_index_kv": current_index_kv, } ) return torch.full((actual_seq_q, 2), len(topk_calls), dtype=torch.int32) indexer._maybe_materialize_shared_index_buffer = fake_materialize indexer._get_topk_ragged_with_cp = fake_get_topk forward_batch = type( "ForwardBatchStub", (), { "nsa_cp_metadata": NSAContextParallelMetadata( kv_len_prev=5, kv_len_next=9, actual_seq_q_prev=3, actual_seq_q_next=2, actual_seq_q_prev_cu_tensor=torch.tensor([0, 3], dtype=torch.int32), actual_seq_q_next_cu_tensor=torch.tensor([0, 2], dtype=torch.int32), ) }, )() q_fp8 = torch.arange(5 * 4, dtype=torch.float32).view(5, 4) weights = torch.arange(5 * 2, dtype=torch.float32).view(5, 2) result = Indexer._get_topk_in_seq_cp_pair( indexer, forward_batch, layer_id=7, q_fp8=q_fp8, weights=weights, metadata=Metadata(), current_index_kv=None, ) self.assertEqual(len(materialize_calls), 1) self.assertIs(materialize_calls[0][1], logical_pages) self.assertEqual(len(topk_calls), 2) self.assertIs(topk_calls[0]["shared_index_buffer"], materialized_index) self.assertIs(topk_calls[1]["shared_index_buffer"], materialized_index) self.assertIs(topk_calls[0]["shared_block_tables"], dense_pages) self.assertIs(topk_calls[1]["shared_block_tables"], dense_pages) self.assertIsNone(topk_calls[0]["current_index_kv"]) self.assertEqual(topk_calls[0]["kv_len"], 5) self.assertEqual(topk_calls[1]["kv_len"], 9) self.assertEqual(topk_calls[0]["actual_seq_q_cu_tensor"].tolist(), [0, 3]) self.assertEqual(topk_calls[1]["actual_seq_q_cu_tensor"].tolist(), [0, 2]) self.assertEqual(result.tolist(), [[1, 1], [1, 1], [1, 1], [2, 2], [2, 2]]) def test_indexer_in_seq_cp_pair_skips_materialize_when_current_index_reused(self): import torch from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer indexer = object.__new__(Indexer) current_index_kv = (torch.tensor([1]), torch.tensor([2])) materialize_calls = [] topk_calls = [] class Metadata: def get_page_table_64(self): raise AssertionError("current index reuse should not read page table") def fake_materialize(forward_batch, layer_id, logical_page_table): materialize_calls.append((layer_id, logical_page_table)) raise AssertionError("current index reuse should not materialize") def fake_get_topk( forward_batch, layer_id, q_fp8, weights, metadata, kv_len, actual_seq_q, cp_index=None, current_index_kv=None, shared_index_buffer=None, shared_block_tables=None, actual_seq_q_tensor=None, actual_seq_q_cu_tensor=None, ): topk_calls.append( { "current_index_kv": current_index_kv, "actual_seq_q_tensor": actual_seq_q_tensor, "actual_seq_q_cu_tensor": actual_seq_q_cu_tensor, "shared_index_buffer": shared_index_buffer, "shared_block_tables": shared_block_tables, } ) return torch.full((actual_seq_q, 2), len(topk_calls), dtype=torch.int32) indexer._maybe_materialize_shared_index_buffer = fake_materialize indexer._get_topk_ragged_with_cp = fake_get_topk forward_batch = type( "ForwardBatchStub", (), { "nsa_cp_metadata": NSAContextParallelMetadata( kv_len_prev=5, kv_len_next=9, actual_seq_q_prev=3, actual_seq_q_next=2, actual_seq_q_prev_cu_tensor=torch.tensor([0, 3], dtype=torch.int32), actual_seq_q_next_cu_tensor=torch.tensor([0, 2], dtype=torch.int32), ) }, )() result = Indexer._get_topk_in_seq_cp_pair( indexer, forward_batch, layer_id=7, q_fp8=torch.empty(5, 4), weights=torch.empty(5, 2), metadata=Metadata(), current_index_kv=current_index_kv, ) self.assertEqual(materialize_calls, []) self.assertEqual(len(topk_calls), 2) self.assertIs(topk_calls[0]["current_index_kv"], current_index_kv) self.assertIs(topk_calls[1]["current_index_kv"], current_index_kv) self.assertIsNone(topk_calls[0]["shared_index_buffer"]) self.assertIsNone(topk_calls[1]["shared_block_tables"]) self.assertEqual(topk_calls[0]["actual_seq_q_cu_tensor"].tolist(), [0, 3]) self.assertEqual(topk_calls[1]["actual_seq_q_cu_tensor"].tolist(), [0, 2]) self.assertEqual(result.tolist(), [[1, 1], [1, 1], [1, 1], [2, 2], [2, 2]]) def test_paged_topk_transform_uses_cu_override_without_scan_metadata_ops(self): import torch from sglang.srt.layers.attention.nsa_backend import ( NSAMetadata, NSAIndexerMetadata, TopkTransformMethod, ) cu_override = torch.tensor([0, 4], dtype=torch.int32) attn_metadata = NSAMetadata( page_size=64, cache_seqlens_int32=torch.tensor([4], dtype=torch.int32), max_seq_len_q=4, max_seq_len_k=8, cu_seqlens_q=torch.tensor([0, 4], dtype=torch.int32), cu_seqlens_k=torch.tensor([0, 8], dtype=torch.int32), page_table_1=torch.arange(8, dtype=torch.int32).view(1, 8), real_page_table=torch.arange(8, dtype=torch.int32).view(1, 8), nsa_cache_seqlens_int32=torch.tensor([4], dtype=torch.int32), nsa_cu_seqlens_q=torch.arange(2, dtype=torch.int32), nsa_cu_seqlens_k=torch.tensor([0, 4], dtype=torch.int32), nsa_extend_seq_lens_list=[4], nsa_seqlens_expanded=torch.arange(1, 5, dtype=torch.int32), topk_indices_offset=torch.zeros(4, dtype=torch.int32), ) metadata = NSAIndexerMetadata( attn_metadata=attn_metadata, topk_transform_method=TopkTransformMethod.PAGED, ) logits = torch.zeros(4, 8) lengths = torch.arange(1, 5, dtype=torch.int32) expected = torch.full((4, 2), 7, dtype=torch.int32) def fake_fused(**kwargs): self.assertIs(kwargs["cu_seqlens_q"], cu_override) self.assertIs(kwargs["lengths"], lengths) return expected fake_sgl_kernel = SimpleNamespace( fast_topk_transform_fused=fake_fused, fast_topk_transform_ragged_fused=lambda **_: (_ for _ in ()).throw( AssertionError("ragged path should not run") ), fast_topk_v2=lambda *_, **__: (_ for _ in ()).throw( AssertionError("unfused path should not run") ), ) with ( patch.dict(sys.modules, {"sgl_kernel": fake_sgl_kernel}), patch( "sglang.srt.layers.attention.nsa_backend.envs.SGLANG_NSA_FUSE_TOPK.get", return_value=True, ), patch( "sglang.srt.layers.attention.nsa_backend.compute_cu_seqlens", side_effect=AssertionError("paged override should skip cumsum"), ), patch( "torch.repeat_interleave", side_effect=AssertionError("paged topk should not build ragged offsets"), ), ): actual = metadata.topk_transform( logits, topk=2, cu_seqlens_q=torch.tensor([4], dtype=torch.int32), ke_offset=lengths, cu_seqlens_q_topk_override=cu_override, ) self.assertIs(actual, expected) if __name__ == "__main__": unittest.main()