Preserve CP ragged topk coordinates under batch planning
Batch-planning metadata can now be attached even for single-request CP prefill, which routed FP8 flashmla_sparse through the batch cp_index path. That path used compact MQA-buffer row bases for score lookup but did not override the final ragged topk coordinate base consumed by attention, so topk indices could point at the wrong KV rows and produce low accept length or meaningless output.\n\nThis keeps ordinary long page tails out of compute padding, reserves compute padding for truly tiny suffixes, and makes cp_index RAGGED topk emit request-ragged offsets while preserving the compact buffer descriptors used for score materialization. The debug ledger records the rejected intermediate diagnoses and the confirmed coordinate-space failure.\n\nConstraint: CP shared-KV cache residency is page-granular, but attention/index compute must not consume synthetic long-tail rows.\nConstraint: FP8 CP prefill uses flashmla_sparse/RAGGED, where fused topk output is consumed directly as attention page_table_1.\nRejected: Disable current reuse or batch planning | would hide the regression and lose the intended bs>1 fast path.\nRejected: Treat all page tails as compute padding | regresses bs=1 semantics and can corrupt query/topk row contracts.\nConfidence: medium\nScope-risk: moderate\nDirective: Do not change cp_index RAGGED topk offset handling without verifying score-buffer row bases and final attention KV coordinate space independently.\nTested: python -m py_compile on touched Python/test files; git diff --check; remote targeted ragged cp_index offset regression test; remote test_nsa_cp_utils.py; remote test_cp_shared_kv_runtime.py; user-reported ETE output recovered after restart.\nNot-tested: Agent-driven full ETE traffic run; broad multi-request bs>1 production soak.
This commit is contained in:
@@ -586,18 +586,17 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
plan.request_compute_seq_q_prev[0] + plan.request_compute_seq_q_next[0]
|
||||
)
|
||||
self.assertEqual(valid_local_rows, 4995)
|
||||
self.assertEqual(compute_local_rows, 5056)
|
||||
self.assertTrue(plan.compute_padding_enabled)
|
||||
self.assertEqual(compute_local_rows, valid_local_rows)
|
||||
self.assertFalse(plan.compute_padding_enabled)
|
||||
|
||||
local_compute = split_tensor_by_cp_batch_plan(
|
||||
local = split_tensor_by_cp_batch_plan(
|
||||
torch.arange(40387, dtype=torch.int64),
|
||||
plan,
|
||||
mode="1d",
|
||||
)
|
||||
self.assertEqual(local_compute.numel(), compute_local_rows)
|
||||
self.assertEqual(local_compute[:2560].tolist(), list(range(2560)))
|
||||
self.assertEqual(local_compute[2560:4995].tolist(), list(range(37952, 40387)))
|
||||
self.assertEqual(local_compute[4995:].tolist(), [0] * 61)
|
||||
self.assertEqual(local.numel(), valid_local_rows)
|
||||
self.assertEqual(local[:2560].tolist(), list(range(2560)))
|
||||
self.assertEqual(local[2560:4995].tolist(), list(range(37952, 40387)))
|
||||
|
||||
local_valid = select_cp_local_valid_rows_for_cache_write(
|
||||
SimpleNamespace(
|
||||
@@ -606,10 +605,10 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
batch_plan=plan,
|
||||
)
|
||||
),
|
||||
local_compute,
|
||||
local,
|
||||
)
|
||||
self.assertEqual(local_valid.numel(), valid_local_rows)
|
||||
self.assertEqual(local_valid.tolist(), local_compute[:valid_local_rows].tolist())
|
||||
self.assertEqual(local_valid.tolist(), local.tolist())
|
||||
|
||||
selected = _select_batch_topk_query_lengths(
|
||||
cp_metadata=NSAContextParallelMetadata(batch_size=1, batch_plan=plan),
|
||||
@@ -629,7 +628,7 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
selected.request_valid_seq_q_next, plan.request_valid_seq_q_next
|
||||
)
|
||||
|
||||
selected_compute = _select_batch_topk_query_lengths(
|
||||
selected_compute_alias = _select_batch_topk_query_lengths(
|
||||
cp_metadata=NSAContextParallelMetadata(batch_size=1, batch_plan=plan),
|
||||
batch_plan=plan,
|
||||
batch_size=1,
|
||||
@@ -637,20 +636,79 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
weights_tokens=compute_local_rows,
|
||||
)
|
||||
|
||||
self.assertTrue(selected_compute.uses_compute_query_rows)
|
||||
self.assertFalse(selected_compute_alias.uses_compute_query_rows)
|
||||
self.assertEqual(
|
||||
selected_compute.request_seq_q_prev, plan.request_compute_seq_q_prev
|
||||
selected_compute_alias.request_seq_q_prev, plan.request_compute_seq_q_prev
|
||||
)
|
||||
self.assertEqual(
|
||||
selected_compute.request_seq_q_next, plan.request_compute_seq_q_next
|
||||
selected_compute_alias.request_seq_q_next, plan.request_compute_seq_q_next
|
||||
)
|
||||
self.assertEqual(
|
||||
selected_compute.request_valid_seq_q_prev, plan.request_valid_seq_q_prev
|
||||
selected_compute_alias.request_valid_seq_q_prev, plan.request_valid_seq_q_prev
|
||||
)
|
||||
self.assertEqual(
|
||||
selected_compute.request_valid_seq_q_next, plan.request_valid_seq_q_next
|
||||
selected_compute_alias.request_valid_seq_q_next, plan.request_valid_seq_q_next
|
||||
)
|
||||
|
||||
def test_batch_plan_keeps_long_page_tail_out_of_compute_padding(self):
|
||||
plan = build_batch_page_aligned_in_seq_split_plan(
|
||||
extend_lens=[40387],
|
||||
prefix_lens=[0],
|
||||
page_size=64,
|
||||
cp_size=8,
|
||||
cp_rank=0,
|
||||
)
|
||||
|
||||
self.assertFalse(plan.compute_padding_enabled)
|
||||
self.assertEqual(plan.request_valid_padded_pages, [632])
|
||||
self.assertEqual(plan.request_valid_padded_tokens, [40448])
|
||||
self.assertEqual(plan.request_compute_padded_pages, [632])
|
||||
self.assertEqual(plan.request_compute_padded_tokens, [40387])
|
||||
self.assertEqual(plan.request_compute_padding_tokens, [0])
|
||||
self.assertEqual(
|
||||
plan.request_compute_split_lists,
|
||||
plan.request_valid_split_lists,
|
||||
)
|
||||
|
||||
def test_batch_plan_compute_padding_only_pads_tiny_request_in_mixed_batch(self):
|
||||
import torch
|
||||
|
||||
plan = build_batch_page_aligned_in_seq_split_plan(
|
||||
extend_lens=[65, 40387],
|
||||
prefix_lens=[40320, 0],
|
||||
page_size=64,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
)
|
||||
|
||||
self.assertTrue(plan.compute_padding_enabled)
|
||||
self.assertEqual(plan.request_valid_padded_pages, [2, 632])
|
||||
self.assertEqual(plan.request_compute_padded_pages, [8, 632])
|
||||
self.assertEqual(plan.request_compute_padded_tokens, [512, 40387])
|
||||
self.assertEqual(plan.request_compute_padding_tokens, [447, 0])
|
||||
self.assertEqual(
|
||||
plan.request_compute_split_lists[1],
|
||||
plan.request_valid_split_lists[1],
|
||||
)
|
||||
|
||||
local = split_tensor_by_cp_batch_plan(
|
||||
torch.arange(65 + 40387, dtype=torch.int64),
|
||||
plan,
|
||||
mode="1d",
|
||||
)
|
||||
self.assertEqual(local.numel(), sum(plan.request_compute_rank_local_tokens))
|
||||
|
||||
valid = select_cp_local_valid_rows_for_cache_write(
|
||||
SimpleNamespace(
|
||||
nsa_cp_metadata=NSAContextParallelMetadata(
|
||||
batch_size=2,
|
||||
batch_plan=plan,
|
||||
)
|
||||
),
|
||||
local,
|
||||
)
|
||||
self.assertEqual(valid.numel(), sum(plan.request_valid_rank_local_tokens))
|
||||
|
||||
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],
|
||||
@@ -687,14 +745,14 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
self.assertEqual(build_flat_page_owner_plan(plan), [0, 0, 1, 1])
|
||||
|
||||
local_1d = split_tensor_by_cp_batch_plan(torch.arange(13), plan, mode="1d")
|
||||
self.assertEqual(local_1d.tolist(), [0, 0, 0, 0, 8, 9, 10, 11, 12, 0, 0, 0])
|
||||
self.assertEqual(local_1d.tolist(), [0, 0, 0, 0, 8, 9, 10, 11, 12])
|
||||
|
||||
local_data = split_tensor_by_cp_batch_plan(
|
||||
torch.arange(13 * 2).view(13, 2), plan, mode="data"
|
||||
)
|
||||
self.assertEqual(
|
||||
local_data[:, 0].tolist(),
|
||||
[0, 0, 0, 0, 16, 18, 20, 22, 24, 0, 0, 0],
|
||||
[0, 0, 0, 0, 16, 18, 20, 22, 24],
|
||||
)
|
||||
|
||||
def test_collect_last_token_hidden_uses_batch_owner_metadata(self):
|
||||
@@ -3227,6 +3285,83 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
self.assertEqual(deep_gemm_calls[0]["ks"], [0, 0, 3, 6, 10, 10, 10])
|
||||
self.assertEqual(deep_gemm_calls[0]["ke"], [2, 3, 6, 10, 12, 13, 14])
|
||||
|
||||
def test_indexer_ragged_cp_index_batch_uses_request_ragged_offsets(self):
|
||||
import contextlib
|
||||
import torch
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from sglang.srt.layers.attention.nsa import nsa_indexer
|
||||
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
|
||||
|
||||
indexer = object.__new__(Indexer)
|
||||
indexer.index_topk = 2
|
||||
indexer._with_real_sm_count = lambda: contextlib.nullcontext()
|
||||
|
||||
topk_kwargs = []
|
||||
|
||||
def fake_prepare(**kwargs):
|
||||
total_kv_len = int(kwargs["total_kv_len"])
|
||||
return (
|
||||
torch.zeros((total_kv_len, 1), dtype=torch.uint8),
|
||||
torch.zeros((total_kv_len,), dtype=torch.float32),
|
||||
torch.tensor([0, 0, 3, 6, 10, 10, 10], dtype=torch.int32),
|
||||
torch.tensor([2, 3, 3, 4, 2, 3, 4], dtype=torch.int32),
|
||||
)
|
||||
|
||||
def fake_logits(q_fp8, kv_fp8, weights, ks, ke, clean_logits=False):
|
||||
return torch.zeros((int(q_fp8.shape[0]), 8), dtype=torch.float32)
|
||||
|
||||
class Metadata:
|
||||
def get_page_table_64(self):
|
||||
raise AssertionError("current cp_index path must not materialize index pages")
|
||||
|
||||
def topk_transform(self, logits, topk, **kwargs):
|
||||
topk_kwargs.append(kwargs)
|
||||
return torch.zeros((int(logits.shape[0]), topk), dtype=torch.int32)
|
||||
|
||||
forward_batch = SimpleNamespace(
|
||||
token_to_kv_pool=SimpleNamespace(page_size=64, index_head_dim=1),
|
||||
seq_lens_cpu=torch.tensor([3, 4], dtype=torch.int64),
|
||||
extend_seq_lens_cpu=[3, 4],
|
||||
)
|
||||
current_index_kv = (
|
||||
torch.arange(7, dtype=torch.uint8).view(7, 1),
|
||||
torch.arange(7, dtype=torch.float32).view(7, 1),
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
nsa_indexer,
|
||||
"try_tai_prepare_cp_mqa_current_index_batch",
|
||||
side_effect=fake_prepare,
|
||||
create=True,
|
||||
), patch.object(
|
||||
nsa_indexer,
|
||||
"deep_gemm",
|
||||
SimpleNamespace(fp8_mqa_logits=fake_logits),
|
||||
):
|
||||
Indexer._get_topk_ragged_with_cp(
|
||||
indexer,
|
||||
forward_batch,
|
||||
layer_id=7,
|
||||
q_fp8=torch.empty((7, 1), dtype=torch.float32),
|
||||
weights=torch.empty((7, 1, 1), dtype=torch.float32),
|
||||
metadata=Metadata(),
|
||||
kv_len=0,
|
||||
actual_seq_q=7,
|
||||
cp_index=[(0, 1, 3), (0, 2, 3), (1, 3, 4), (1, 1, 4)],
|
||||
current_index_kv=current_index_kv,
|
||||
)
|
||||
|
||||
self.assertEqual(len(topk_kwargs), 1)
|
||||
offset = topk_kwargs[0].get("topk_indices_offset_override")
|
||||
self.assertIsNotNone(offset)
|
||||
# Batch cp_index compacts each CP segment's K into a temporary buffer,
|
||||
# but flashmla_sparse consumes the normal ragged KV layout. The fused
|
||||
# ragged topk offset therefore has to stay in request/KV coordinates,
|
||||
# not compact segment coordinates or compact-q cu-seqlens.
|
||||
self.assertEqual(offset.tolist(), [0, 0, 0, 3, 3, 3, 3])
|
||||
|
||||
def test_indexer_ragged_cp_index_shared_batch_uses_tai_prepare_once(self):
|
||||
import contextlib
|
||||
import torch
|
||||
|
||||
@@ -1044,10 +1044,10 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
extend_seq_lens_cpu=[40387],
|
||||
seq_lens_cpu=torch.tensor([40387], dtype=torch.int32),
|
||||
out_cache_loc=torch.arange(40392, dtype=torch.int64),
|
||||
cp_local_out_cache_loc=torch.arange(5056, dtype=torch.int64),
|
||||
cp_local_out_cache_loc=torch.arange(4995, dtype=torch.int64),
|
||||
)
|
||||
local_k = torch.empty((5056, 1), dtype=torch.float32)
|
||||
local_rope = torch.empty((5056, 1), dtype=torch.float32)
|
||||
local_k = torch.empty((4995, 1), dtype=torch.float32)
|
||||
local_rope = torch.empty((4995, 1), dtype=torch.float32)
|
||||
|
||||
with envs.SGLANG_CP_SHARED_KV_CURRENT_REUSE.override(True):
|
||||
self.assertTrue(runtime.should_reuse_current_extend_kv(forward_batch))
|
||||
@@ -1057,12 +1057,12 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
local_k,
|
||||
local_rope,
|
||||
),
|
||||
5056,
|
||||
4995,
|
||||
)
|
||||
self.assertIsNone(
|
||||
runtime.current_extend_kv_rows_for_reuse(
|
||||
forward_batch,
|
||||
local_k[:5055],
|
||||
local_k[:4994],
|
||||
local_rope,
|
||||
)
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user