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:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user