diff --git a/test/registered/unit/layers/test_nsa_cp_utils.py b/test/registered/unit/layers/test_nsa_cp_utils.py index c1ff1d18f..cacc82485 100644 --- a/test/registered/unit/layers/test_nsa_cp_utils.py +++ b/test/registered/unit/layers/test_nsa_cp_utils.py @@ -743,6 +743,76 @@ class TestNSAInSeqCPUtils(unittest.TestCase): ) self.assertEqual(valid.numel(), sum(plan.request_valid_rank_local_tokens)) + def test_compute_padding_tiny_batch_valid_rows_are_not_suffix_maskable_for_moe( + self, + ): + # Regression invariant for the warm-cache GSM8K bs>1 failure: + # compute padding gives every tiny request a fixed local page slot, but + # the valid rows inside those slots are per-request prefixes, not one + # contiguous tensor prefix. DeepEP/MoE top-k's scalar + # num_token_non_padded can only mask a suffix, so it cannot represent + # this CP-local row layout. + import torch + + plan = build_batch_page_aligned_in_seq_split_plan( + extend_lens=[29, 9, 60, 9, 17], + prefix_lens=[704, 704, 640, 704, 704], + page_size=64, + cp_size=8, + cp_rank=0, + ) + + compute_rows = sum(plan.request_compute_rank_local_tokens) + valid_rows = sum(plan.request_valid_rank_local_tokens) + self.assertTrue(plan.compute_padding_enabled) + self.assertEqual(plan.request_compute_rank_local_tokens, [64] * 5) + self.assertEqual(plan.request_valid_rank_local_tokens, [29, 9, 60, 9, 17]) + self.assertEqual(compute_rows, 320) + self.assertEqual(valid_rows, 124) + + valid_mask = torch.zeros(compute_rows, dtype=torch.bool) + for req_offset, spans in zip( + plan.request_compute_rank_local_offsets, + plan.request_valid_query_row_spans, + ): + for start, end in spans: + if end > start: + valid_mask[req_offset + start : req_offset + end] = True + + # A suffix-padding scalar would keep rows [0, valid_rows) and mask the + # rest. The CP layout has dummy page-tail rows inside that prefix and + # later valid rows after it. + scalar_prefix_mask = torch.arange(compute_rows) < valid_rows + self.assertFalse(torch.equal(valid_mask, scalar_prefix_mask)) + self.assertGreater(int((scalar_prefix_mask & ~valid_mask).sum().item()), 0) + self.assertGreater(int((valid_mask & ~scalar_prefix_mask).sum().item()), 0) + + def test_compute_padding_non_owner_rank_scalar_non_padded_would_unmask_dummy_rows( + self, + ): + # Same regression shape on a non-owner CP rank: the local tensor has + # compute rows but no valid rows. Reusing the global input-token count + # as num_token_non_padded would route dummy rows through MoE. + extend_lens = [29, 9, 60, 9, 17] + plan = build_batch_page_aligned_in_seq_split_plan( + extend_lens=extend_lens, + prefix_lens=[704, 704, 640, 704, 704], + page_size=64, + cp_size=8, + cp_rank=1, + ) + + compute_rows = sum(plan.request_compute_rank_local_tokens) + valid_rows = sum(plan.request_valid_rank_local_tokens) + global_num_token_non_padded = sum(extend_lens) + + self.assertTrue(plan.compute_padding_enabled) + self.assertEqual(plan.request_compute_rank_local_tokens, [64] * 5) + self.assertEqual(plan.request_valid_rank_local_tokens, [0] * 5) + self.assertEqual(compute_rows, 320) + self.assertEqual(valid_rows, 0) + self.assertGreater(global_num_token_non_padded, valid_rows) + def test_batch_plan_compute_padding_is_per_request_not_batch_total(self): plan = build_batch_page_aligned_in_seq_split_plan( extend_lens=[65, 1024], @@ -912,6 +982,64 @@ class TestNSAInSeqCPUtils(unittest.TestCase): self.assertEqual(collected.tolist(), [[10.0], [20.0]]) + def test_collect_last_token_hidden_matches_parallel20_tiny_extend_layout(self): + import torch + + # Regression shape from the parallel=20 GSM8K warm-cache probe: + # tiny extend requests are compute-padded to one page per request. + # All real current tokens and all last tokens are owned by CP0, while + # the local hidden rows remain laid out as fixed 64-row request slots. + plan = build_batch_page_aligned_in_seq_split_plan( + extend_lens=[19, 21, 61, 33], + prefix_lens=[704, 704, 640, 704], + page_size=64, + cp_size=8, + cp_rank=0, + ) + self.assertEqual(plan.request_last_token_owner, [0, 0, 0, 0]) + self.assertEqual(plan.request_last_token_local_offset, [18, 20, 60, 32]) + self.assertEqual(plan.request_compute_rank_local_offsets, [0, 64, 128, 192]) + + hidden_states = torch.zeros((256, 1), dtype=torch.float32) + expected = [[100.0], [200.0], [300.0], [400.0]] + for value, rank_offset, local_offset in zip( + [100.0, 200.0, 300.0, 400.0], + plan.request_compute_rank_local_offsets, + plan.request_last_token_local_offset, + ): + hidden_states[rank_offset + local_offset] = value + + forward_batch = SimpleNamespace( + extend_seq_lens_cpu=[19, 21, 61, 33], + nsa_cp_metadata=NSAContextParallelMetadata( + batch_size=4, + batch_plan=plan, + ), + ) + + def fake_all_gather(output, local_last): + self.assertEqual(local_last.tolist(), expected) + output.zero_() + output[:4] = local_last + + 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.get_attention_cp_rank", + return_value=0, + ), + patch( + "sglang.srt.layers.attention.nsa.utils.attn_cp_all_gather_into_tensor", + side_effect=fake_all_gather, + ), + ): + collected = cp_collect_last_token_hidden(hidden_states, forward_batch, 8) + + self.assertEqual(collected.tolist(), expected) + def test_collect_last_token_hidden_fails_fast_without_batch_owner_metadata(self): import torch @@ -1373,6 +1501,79 @@ class TestNSAInSeqCPUtils(unittest.TestCase): self.assertTrue(torch.equal(actual, expected)) + def test_batch_in_seq_all_gather_rerange_matches_parallel20_tiny_extend_layout(self): + import torch + + cp_size = 8 + extend_lens = [19, 21, 61, 33] + prefix_lens = [704, 704, 640, 704] + plans = [ + build_batch_page_aligned_in_seq_split_plan( + extend_lens=extend_lens, + prefix_lens=prefix_lens, + page_size=64, + cp_size=cp_size, + cp_rank=rank, + ) + for rank in range(cp_size) + ] + valid_split_lists = plans[0].request_split_lists + compute_split_lists = plans[0].request_compute_split_lists + + self.assertEqual(plans[0].request_valid_rank_local_tokens, extend_lens) + for rank in range(1, cp_size): + self.assertEqual(plans[rank].request_valid_rank_local_tokens, [0, 0, 0, 0]) + + max_rank_token = max( + sum( + split[rank] + split[cp_size * 2 - rank - 1] + for split in compute_split_lists + ) + for rank in range(cp_size) + ) + self.assertEqual(max_rank_token, 64 * len(extend_lens)) + + input_tensor_all = torch.zeros((max_rank_token * cp_size, 1), dtype=torch.float32) + expected_rows = [] + rank0_cursor = 0 + for req_id, extend_len in enumerate(extend_lens): + valid_rows = torch.arange( + req_id * 1000, + req_id * 1000 + extend_len, + dtype=torch.float32, + ).view(-1, 1) + input_tensor_all[rank0_cursor : rank0_cursor + extend_len] = valid_rows + # Fill the rest of the 64-row compute slot with poison values that + # must not appear after valid-output rerange. + pad_len = 64 - extend_len + if pad_len: + input_tensor_all[ + rank0_cursor + extend_len : rank0_cursor + 64 + ] = 900000.0 + req_id + expected_rows.append(valid_rows) + rank0_cursor += 64 + # Non-owner ranks only have compute padding in this tiny-extend case. + input_tensor_all[max_rank_token:] = -777.0 + + forward_batch = SimpleNamespace( + nsa_cp_metadata=NSAContextParallelMetadata( + batch_size=len(extend_lens), + request_split_lists=valid_split_lists, + request_compute_split_lists=compute_split_lists, + compute_padding_enabled=True, + max_rank_len=[max_rank_token] * cp_size, + ) + ) + + actual = _torch_batch_in_seq_all_gather_rerange( + input_tensor_all, + forward_batch, + cp_size=cp_size, + ) + + expected = torch.cat(expected_rows, dim=0) + self.assertTrue(torch.equal(actual, expected)) + def _build_batch_rerange_case( self, *,