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