Fix RAGGED CP cache-hit current KV reuse

FP8 flashmla_sparse uses flattened RAGGED page tables that include both cached prefix and the just-computed current suffix. The old cache-hit path materialized the whole flattened range from persistent KV, which could read current rows through the wrong contract under CP shared-KV and compute padding.\n\nThis change makes the RAGGED path use the page-slot partial-current compose contract: prefix pages are materialized from cache slots while current rows are sourced from fresh k/k_rope and packed for FP8 when needed. A new helper accepts the actual current-row contracts seen by attention code: already-local valid rows, CP-local compute-padded rows, or unsplit global valid rows.\n\nConstraint: CP shared-KV stores and consumes cache at page granularity, while attention current rows may be valid-token tensors rather than cache-write local compute rows.\nRejected: Full materialize prefix+current from persistent KV | it can read current suffix through stale or unordered persistent cache state.\nRejected: Reuse select_cp_local_valid_rows_for_cache_write directly | it only accepts CP-local compute-padded rows and killed prefill when RAGGED supplied global valid rows.\nConfidence: medium\nScope-risk: moderate\nDirective: Do not route RAGGED FP8 cache-hit current suffix through full materialize without proving current write ordering and row-contract compatibility.\nTested: Local py_compile for touched runtime/test files.\nTested: Remote container pytest for RAGGED current compose and compute-padding global-current row selection.\nNot-tested: Full GSM8K warm-cache ETE after restart.
This commit is contained in:
laoyao0822
2026-06-08 03:39:34 +08:00
parent b324407def
commit 2aa0b7313e
4 changed files with 261 additions and 9 deletions
@@ -32,6 +32,7 @@ from sglang.srt.layers.attention.nsa.utils import (
nsa_use_prefill_cp,
pad_cp_local_input_ids_for_embedding,
prepare_input_dp_with_cp_dsa,
select_cp_current_valid_rows_for_reuse,
select_cp_local_valid_rows_for_cache_write,
split_tensor_by_cp_batch_plan,
split_in_seq_cp_local_pair,
@@ -1703,6 +1704,43 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
self.assertEqual(selected.tolist(), [[50.0, 51.0]])
def test_select_cp_current_valid_rows_accepts_global_rows_under_compute_padding(
self,
):
import torch
plan = build_batch_page_aligned_in_seq_split_plan(
extend_lens=[65],
prefix_lens=[0],
page_size=64,
cp_size=8,
cp_rank=1,
)
self.assertTrue(plan.compute_padding_enabled)
global_current = torch.arange(65 * 2, dtype=torch.float32).view(65, 2)
expected = split_tensor_by_cp_batch_plan(
global_current,
plan,
mode="data",
split_kind="valid",
)
forward_batch = SimpleNamespace(
extend_seq_lens_cpu=[65],
cp_local_out_cache_loc=torch.arange(expected.shape[0], dtype=torch.long),
nsa_cp_metadata=NSAContextParallelMetadata(
batch_size=1,
batch_plan=plan,
),
)
selected = select_cp_current_valid_rows_for_reuse(
forward_batch,
global_current,
)
self.assertIsNotNone(selected)
self.assertTrue(torch.equal(selected, expected))
def test_cp_split_and_rebuild_position_is_batch_aware_and_compute_padded(self):
import torch
@@ -1330,7 +1330,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
self.assertIsNone(result)
def test_fp8_ragged_mla_defers_cp_materialize_to_flattened_path(self):
def test_fp8_ragged_mla_uses_page_slot_current_compose_for_cache_hit(self):
from pathlib import Path
source = (
@@ -1353,7 +1353,13 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
ragged_end = source.index(" attn_output = self._forward_flashmla_sparse", ragged_start)
ragged_source = source[ragged_start:ragged_end]
self.assertIn("page_table_1_flattened", ragged_source)
self.assertIn("materialize_shared_token_kv_buffer", ragged_source)
self.assertIn(
"materialize_prefix_and_reuse_current_kv_page_slots",
ragged_source,
)
self.assertIn("select_cp_current_valid_rows_for_reuse", ragged_source)
self.assertIn("prefix_slot_spans=", ragged_source)
self.assertIn("current_slot_spans=", ragged_source)
@unittest.skipIf(not torch.cuda.is_available(), "CUDA is required")
def test_tai_current_slot_fill_sparse_page_self_test_passes_on_installed_kernel(