import unittest from sglang.srt.layers.attention.nsa.utils import ( _get_in_seq_last_token_owner_and_offset, build_page_aligned_in_seq_split_list, build_token_balanced_in_seq_split_list, split_in_seq_cp_local_pair, ) from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=1, suite="stage-a-test-cpu") class TestNSAInSeqCPUtils(unittest.TestCase): 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_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_falls_back_when_page_units_are_too_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.assertFalse(info.page_aligned) self.assertEqual(split_list, build_token_balanced_in_seq_split_list(512, 8)) 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") if __name__ == "__main__": unittest.main()