Lock bs>1 CP compute-padding invariants

The GSM8K warm-cache investigation exposed several tiny-extend shapes where compute padding, valid-row selection, last-token collect, and all-gather rerange must stay request-slot aware. These tests pin those invariants without changing runtime behavior.

Constraint: bs>1 tiny extends use page-granular compute slots while only a subset of rows are semantically valid
Rejected: Rely on scalar non-padded token counts for these layouts | valid rows are not always a simple suffix mask
Confidence: medium
Scope-risk: narrow
Directive: Do not weaken these tests unless the replacement path proves equivalent request-slot and valid-row semantics
Tested: Remote pytest test_nsa_cp_utils.py as part of the 263-test CP regression suite
Not-tested: CUDA graph paths; prefill does not use cuda graph in the current CP setup
This commit is contained in:
laoyao0822
2026-06-08 20:56:48 +08:00
parent b58513cba4
commit 3698b0c22c
@@ -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,
*,