feat: add page-aligned index validation for CP HiCache direct transfers

Add validate_page_aligned_token_indices utility and apply it across
HiCache write/load paths, NSA indexer transfers, and CUDA direct copy
kernels to reject malformed (partial, misaligned, non-contiguous) page
groups before they reach native transfer code. Also validate supported
CP HiCache backend/layout combinations at server startup.
This commit is contained in:
2026-05-10 02:02:54 +08:00
parent f38937723e
commit 3a5928ef53
10 changed files with 472 additions and 18 deletions
@@ -371,6 +371,26 @@ class TestHiCacheArgs(CustomTestCase):
if expected_decode_backend is not None:
self.assertEqual(args.decode_attention_backend, expected_decode_backend)
def _make_cp_hicache_args(self, **overrides) -> ServerArgs:
args = self._make_args(
enable_hierarchical_cache=True,
enable_nsa_prefill_context_parallel=True,
enable_nsa_prefill_cp_shared_kv=True,
nsa_prefill_cp_mode="in-seq-split",
disaggregation_mode="prefill",
page_size=64,
enable_hisparse=False,
)
for key, value in overrides.items():
setattr(args, key, value)
return args
def _normalize_and_validate_cp_hicache_args(self, **overrides) -> ServerArgs:
args = self._make_cp_hicache_args(**overrides)
args._handle_hicache()
args._handle_cp_hicache_layout_validation()
return args
def test_cp_shared_kv_rejects_hicache_storage_backend(self):
with (
self.assertRaisesRegex(
@@ -468,6 +488,61 @@ class TestHiCacheArgs(CustomTestCase):
self.assertEqual(args.decode_attention_backend, "triton")
def test_cp_hicache_accepts_supported_backend_layout_pairs(self):
cases = [
("kernel", "layer_first"),
("kernel", "page_first"),
("direct", "layer_first"),
("direct", "page_first_direct"),
]
for io_backend, mem_layout in cases:
with self.subTest(io_backend=io_backend, mem_layout=mem_layout):
args = self._normalize_and_validate_cp_hicache_args(
hicache_io_backend=io_backend,
hicache_mem_layout=mem_layout,
)
self.assertEqual(args.hicache_io_backend, io_backend)
self.assertEqual(args.hicache_mem_layout, mem_layout)
def test_cp_hicache_normalizes_supported_alias_pairs(self):
args = self._normalize_and_validate_cp_hicache_args(
hicache_io_backend="kernel",
hicache_mem_layout="page_first_direct",
)
self.assertEqual(args.hicache_io_backend, "direct")
self.assertEqual(args.hicache_mem_layout, "page_first_direct")
args = self._normalize_and_validate_cp_hicache_args(
hicache_io_backend="direct",
hicache_mem_layout="page_first",
)
self.assertEqual(args.hicache_io_backend, "direct")
self.assertEqual(args.hicache_mem_layout, "page_first_direct")
def test_cp_hicache_rejects_page_head_layout(self):
with self.assertRaisesRegex(ValueError, "CP shared KV HiCache.*page_head"):
self._normalize_and_validate_cp_hicache_args(
hicache_io_backend="kernel",
hicache_mem_layout="page_head",
)
def test_cp_hicache_rejects_cuda_kv_split_layout(self):
with self.assertRaisesRegex(
ValueError, "CP shared KV HiCache.*page_first_kv_split"
):
self._normalize_and_validate_cp_hicache_args(
hicache_io_backend="kernel",
hicache_mem_layout="page_first_kv_split",
)
def test_cp_hicache_rejects_kernel_ascend_backend(self):
with self.assertRaisesRegex(ValueError, "CP shared KV HiCache.*kernel_ascend"):
self._normalize_and_validate_cp_hicache_args(
hicache_io_backend="kernel_ascend",
hicache_mem_layout="page_first_kv_split",
)
if __name__ == "__main__":
unittest.main()