From 1f074f434efc79cdf76795aca564fbb3d3409622 Mon Sep 17 00:00:00 2001 From: ThomasX Date: Fri, 8 May 2026 00:12:39 +0800 Subject: [PATCH] docs: add CP HiCache host plan --- .gitignore | 3 + ...26-05-07-cp-hicache-host-implementation.md | 1902 +++++++++++++++++ .../2026-05-07-cp-hicache-host-design.md | 429 ++++ 3 files changed, 2334 insertions(+) create mode 100644 docs/superpowers/plans/2026-05-07-cp-hicache-host-implementation.md create mode 100644 docs/superpowers/specs/2026-05-07-cp-hicache-host-design.md diff --git a/.gitignore b/.gitignore index 9f11f1a36..778cc7b43 100644 --- a/.gitignore +++ b/.gitignore @@ -272,3 +272,6 @@ sgl-kernel/csrc/**/*_musa/ *.glb *.ply *.npz + +# Git worktrees +.worktrees/ diff --git a/docs/superpowers/plans/2026-05-07-cp-hicache-host-implementation.md b/docs/superpowers/plans/2026-05-07-cp-hicache-host-implementation.md new file mode 100644 index 000000000..e09f6150e --- /dev/null +++ b/docs/superpowers/plans/2026-05-07-cp-hicache-host-implementation.md @@ -0,0 +1,1902 @@ +# CP HiCache Host Integration Implementation Plan + +> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. + +**Goal:** Implement host-only HiCache support for NSA CP shared KV while keeping radix/scheduler state in logical KV space and host/device transfers in per-rank physical KV space. + +**Architecture:** Add explicit CP host metadata for `HiRadixCache` nodes and a CP-aware transfer adapter in `HiCacheController`. `TreeNode.value`, scheduler admission, and req-to-token remain full logical locs; `HostKVCache` and `NSATokenToKVPool` remain physical pools. Storage backends stay unsupported with CP shared KV in this stage and must fail fast. + +**Tech Stack:** Python, unittest, SGLang `HiRadixCache`, `HiCacheController`, `CpSharedKVLayout`, `CPSharedPagedTokenToKVPoolAllocator`, `NSATokenToKVPool`. + +--- + +## Reference Documents + +- Design spec: `docs/superpowers/specs/2026-05-07-cp-hicache-host-design.md` +- Existing CP shared KV layout: `python/sglang/srt/mem_cache/cp_shared_kv_layout.py` +- Existing HiCache radix path: `python/sglang/srt/mem_cache/hiradix_cache.py` +- Existing HiCache transfer controller: `python/sglang/srt/managers/cache_controller.py` +- Existing CP allocator: `python/sglang/srt/mem_cache/allocator.py::CPSharedPagedTokenToKVPoolAllocator` + +## File Map + +- Create `test/registered/unit/mem_cache/test_cp_hicache_metadata.py`: metadata, `HiRadixCache` helper, split, eviction, and load-back unit coverage. +- Create `test/registered/unit/managers/test_hicache_controller_cp.py`: CP controller write/load mapping unit coverage with fake pools. +- Modify `python/sglang/srt/mem_cache/hiradix_cache.py`: add `CpHiCacheNodeMetadata`, CP metadata helpers, CP-aware write/load/split/evict behavior. +- Modify `python/sglang/srt/managers/cache_controller.py`: add CP adapter state, write/load mapping, allocation-failure result, and zero-owned no-op ack behavior. +- Modify `python/sglang/srt/managers/schedule_policy.py`: port HiCache load-back `mem_quota` budgeting. +- Modify `python/sglang/srt/mem_cache/allocator.py`: port allocator OOM accounting helpers and pre-allocation page checks. +- Modify `python/sglang/srt/mem_cache/common.py`: port OOM accounting and preserve existing CP compute-owner lane eviction. +- Modify `python/sglang/srt/mem_cache/base_prefix_cache.py`: include allocator state details in OOM strings. +- Modify `python/sglang/srt/mem_cache/radix_cache.py`: replace `defaultdict(TreeNode)` with `{}` and add CP node fields if fields are stored on `TreeNode`. +- Modify `python/sglang/srt/mem_cache/swa_radix_cache.py`: replace `defaultdict(TreeNode)` with `{}`. +- Modify `python/sglang/srt/mem_cache/mamba_radix_cache.py`: replace `defaultdict(TreeNode)` with `{}`. +- Modify `python/sglang/srt/server_args.py`: add CP shared KV + HiCache storage fail-fast. +- Modify `test/registered/unit/server_args/test_server_args.py`: add fail-fast coverage and convert `TestHiCacheArgs` to `CustomTestCase`. +- Modify `test/registered/unit/managers/test_prefill_adder.py`: add `mem_quota` scheduling coverage. + +--- + +### Task 1: Port HiCache Correctness Fixes + +**Files:** +- Modify: `python/sglang/srt/managers/schedule_policy.py` +- Modify: `python/sglang/srt/mem_cache/allocator.py` +- Modify: `python/sglang/srt/mem_cache/common.py` +- Modify: `python/sglang/srt/mem_cache/base_prefix_cache.py` +- Modify: `python/sglang/srt/mem_cache/radix_cache.py` +- Modify: `python/sglang/srt/mem_cache/swa_radix_cache.py` +- Modify: `python/sglang/srt/mem_cache/mamba_radix_cache.py` +- Test: `test/registered/unit/managers/test_prefill_adder.py` +- Test: `test/registered/unit/mem_cache/test_cp_shared_kv_layout.py` + +- [ ] **Step 1: Add failing `mem_quota` test** + +Add this test method to `test/registered/unit/managers/test_prefill_adder.py::TestPrefillAdder`: + +```python +def test_host_load_back_passes_mem_quota(self): + running_batch = self.create_running_batch() + self.mock_token_allocator.available_size.return_value = 512 + self.mock_tree_cache.init_load_back.return_value = ( + __import__("torch").tensor([1, 2, 3, 4], dtype=__import__("torch").int64), + "loaded_node", + ) + adder = self.create_adder( + running_batch, + page_size=64, + rem_input_tokens=4096, + rem_total_tokens=4096, + ) + req = self.create_mock_req("req", priority=0, max_new_tokens=16) + req.extend_input_len = 256 + req.host_hit_length = 128 + req.prefix_indices = __import__("torch").empty((0,), dtype=__import__("torch").int64) + req.last_node = object() + req.last_host_node = object() + req.fill_ids = list(range(256)) + req.cache_protected_len = 0 + req.set_extend_input_len = lambda value: setattr(req, "extend_input_len", value) + req.sampling_params.ignore_eos = False + + result = adder.add_one_req(req, has_chunked_req=False, truncation_align_size=None) + + self.assertNotEqual(result, AddReqResult.NO_TOKEN) + params = self.mock_tree_cache.init_load_back.call_args.args[0] + self.assertEqual(params.mem_quota, 256) +``` + +- [ ] **Step 2: Run the failing test** + +Run: `python3 test/registered/unit/managers/test_prefill_adder.py TestPrefillAdder.test_host_load_back_passes_mem_quota` + +Expected: FAIL because `params.mem_quota` is `None`. + +- [ ] **Step 3: Implement load-back quota helpers** + +Add these methods to `PrefillAdder` in `python/sglang/srt/managers/schedule_policy.py` near `ceil_paged_tokens`: + +```python +def _get_available_device_tokens_for_load_back(self) -> int: + if self.is_hybrid_swa: + return self.token_to_kv_pool_allocator.full_available_size() + return self.token_to_kv_pool_allocator.available_size() + +def _get_load_back_mem_quota(self, real_input_tokens: int) -> int: + reserve_tokens = real_input_tokens + self.page_size + return max(self._get_available_device_tokens_for_load_back() - reserve_tokens, 0) +``` + +Then update the `InitLoadBackParams` call in `add_one_req` to pass: + +```python +mem_quota=self._get_load_back_mem_quota(real_input_tokens), +``` + +- [ ] **Step 4: Verify quota test passes** + +Run: `python3 test/registered/unit/managers/test_prefill_adder.py TestPrefillAdder.test_host_load_back_passes_mem_quota` + +Expected: PASS. + +- [ ] **Step 5: Add failing allocator accounting tests** + +In `test/registered/unit/mem_cache/test_cp_shared_kv_layout.py`, add this import near the existing CI/test imports: + +```python +from sglang.test.test_utils import CustomTestCase +``` + +Change the touched allocator test class to use the project test base: + +```python +class TestCPSharedPagedAllocator(CustomTestCase): +``` + +Add this test method to `TestCPSharedPagedAllocator`: + +```python +def test_allocator_state_str_reports_free_and_release_pages(self): + from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator + + allocator = CPSharedPagedTokenToKVPoolAllocator( + logical_size=64 * 8, + physical_size=64 * 2, + page_size=64, + dtype=torch.bfloat16, + device="cpu", + kvcache=None, + need_sort=True, + cp_size=4, + cp_rank=0, + ) + allocator.free_pages = torch.tensor([1, 5], dtype=torch.int64) + allocator.release_pages = torch.tensor([9], dtype=torch.int64) + + state = allocator.allocator_state_str() + + self.assertIn("allocator_available_size=192", state) + self.assertIn("allocator_free_size=128", state) + self.assertIn("free_pages=2", state) + self.assertIn("allocator_release_size=64", state) + self.assertIn("release_pages=1", state) +``` + +- [ ] **Step 6: Run allocator accounting test** + +Run: `python3 test/registered/unit/mem_cache/test_cp_shared_kv_layout.py TestCPSharedPagedAllocator.test_allocator_state_str_reports_free_and_release_pages` + +Expected: FAIL because `allocator_state_str` does not exist. + +- [ ] **Step 7: Implement allocator accounting helpers** + +Add to `BaseTokenToKVPoolAllocator` in `python/sglang/srt/mem_cache/allocator.py`: + +```python +def immediate_available_pages(self) -> int: + if self.free_pages is None: + return 0 + return len(self.free_pages) + +def deferred_available_pages(self) -> int: + if self.release_pages is None: + return 0 + return len(self.release_pages) + +def immediate_available_size(self) -> int: + if self.free_pages is None: + return self.available_size() + return self.immediate_available_pages() * self.page_size + +def deferred_available_size(self) -> int: + if self.release_pages is None: + return 0 + return self.deferred_available_pages() * self.page_size + +def allocator_state_str(self) -> str: + if self.free_pages is None: + return f"allocator_available_size={self.available_size()}" + return ( + "allocator_available_size=" + f"{self.available_size()} " + f"(allocator_free_size={self.immediate_available_size()} " + f"[free_pages={self.immediate_available_pages()}] + " + f"allocator_release_size={self.deferred_available_size()} " + f"[release_pages={self.deferred_available_pages()}])" + ) +``` + +- [ ] **Step 8: Update paged allocation pre-checks** + +In `python/sglang/srt/mem_cache/allocator.py::PagedTokenToKVPoolAllocator.alloc_extend`, move the new-page count before `out_indices` allocation and replace the existing rough pre-check with this exact block after the debug assertion and `bs = len(prefix_lens)`: + +```python +num_new_pages = get_num_new_pages( + seq_lens=seq_lens_cpu, + page_size=self.page_size, + prefix_lens=prefix_lens_cpu, +) +if self.need_sort and num_new_pages > len(self.free_pages): + self.merge_and_sort_free() +if num_new_pages > len(self.free_pages): + return None +``` + +Then remove this duplicate block that currently appears after `alloc_extend_kernel`: + +```python +num_new_pages = get_num_new_pages( + seq_lens=seq_lens_cpu, + page_size=self.page_size, + prefix_lens=prefix_lens_cpu, +) +if num_new_pages > len(self.free_pages): + return None +``` + +Keep this tail: + +```python +self.free_pages = self.free_pages[num_new_pages:] +return out_indices +``` + +In `python/sglang/srt/mem_cache/allocator.py::PagedTokenToKVPoolAllocator.alloc_decode`, replace the `if self.need_sort and bs > len(self.free_pages):` pre-check with this exact block after `bs = len(seq_lens)`: + +```python +num_new_pages = get_num_new_pages( + seq_lens=seq_lens_cpu, + page_size=self.page_size, + decode=True, +) +if self.need_sort and num_new_pages > len(self.free_pages): + self.merge_and_sort_free() +if num_new_pages > len(self.free_pages): + return None +``` + +Then remove this duplicate block that currently appears after `alloc_decode_kernel`: + +```python +num_new_pages = get_num_new_pages( + seq_lens=seq_lens_cpu, + page_size=self.page_size, + decode=True, +) +if num_new_pages > len(self.free_pages): + return None +``` + +Keep this tail: + +```python +self.free_pages = self.free_pages[num_new_pages:] +return out_indices +``` + +- [ ] **Step 9: Preserve CP compute-owner allocation semantics** + +Do not change `CPSharedPagedTokenToKVPoolAllocator.alloc_extend_compute_owner` lane selection except where it benefits from inherited accounting helpers. Keep `_evict_for_compute_owner_lanes()` behavior in `common.py`. + +- [ ] **Step 10: Update OOM message helpers** + +In `python/sglang/srt/mem_cache/base_prefix_cache.py::BasePrefixCache.available_and_evictable_str`, replace the current method body with: + +```python +allocator = self.token_to_kv_pool_allocator +allocator_state_str = getattr(allocator, "allocator_state_str", None) +if allocator_state_str is None: + allocator_state = f"allocator_available_size={allocator.available_size()}" +else: + allocator_state = allocator_state_str() +evictable_size = self.evictable_size() +return ( + f"Available tokens: {allocator.available_size() + evictable_size} " + f"({allocator_state} + {evictable_size=})\n" +) +``` + +In `python/sglang/srt/mem_cache/common.py`, change the import to include `EvictResult`: + +```python +from sglang.srt.mem_cache.base_prefix_cache import ( + BasePrefixCache, + EvictParams, + EvictResult, +) +``` + +Add this helper near `evict_from_tree_cache`: + +```python +def _evict_result_str(evict_result: EvictResult | None) -> str: + if evict_result is None: + return "evict_result=None" + return ( + "evict_result=(" + f"num_tokens_evicted={evict_result.num_tokens_evicted}, " + f"swa_num_tokens_evicted={evict_result.swa_num_tokens_evicted}, " + f"mamba_num_evicted={evict_result.mamba_num_evicted})" + ) +``` + +Replace `evict_from_tree_cache` with this return-value-preserving version: + +```python +def evict_from_tree_cache( + tree_cache: BasePrefixCache | None, num_tokens: int +) -> EvictResult: + if tree_cache is None: + return EvictResult() + + if tree_cache.is_chunk_cache(): + return EvictResult() + + allocator = tree_cache.token_to_kv_pool_allocator + + if isinstance(allocator, SWATokenToKVPoolAllocator): + full_available_size = allocator.full_available_size() + swa_available_size = allocator.swa_available_size() + + if full_available_size < num_tokens or swa_available_size < num_tokens: + full_num_tokens = max(0, num_tokens - full_available_size) + swa_num_tokens = max(0, num_tokens - swa_available_size) + return tree_cache.evict( + EvictParams(num_tokens=full_num_tokens, swa_num_tokens=swa_num_tokens) + ) + return EvictResult() + + if allocator.available_size() < num_tokens: + return tree_cache.evict(EvictParams(num_tokens=num_tokens)) + return EvictResult() +``` + +In `alloc_token_slots` and `alloc_paged_token_slots_extend`, store the return value: + +```python +evict_result = evict_from_tree_cache(tree_cache, num_tokens) +``` + +Append this line to each OOM message after `available_and_evictable_str(tree_cache)`: + +```python +f"{_evict_result_str(evict_result)}\n" +``` + +Do not remove `_evict_for_compute_owner_lanes()` or any existing CP shared KV fallback logging. + +- [ ] **Step 11: Replace ghost-node defaultdicts** + +In `test/registered/unit/mem_cache/test_radix_cache_unit.py`, add this import near the other test imports: + +```python +from sglang.test.test_utils import CustomTestCase +``` + +Add this new test class near the existing `TestTreeNode` class before changing production code: + +```python +class TestTreeNodeChildren(CustomTestCase): + def test_missing_child_access_does_not_create_orphan_node(self): + node = TreeNode() + + with self.assertRaises(KeyError): + _ = node.children["missing"] + + self.assertEqual(node.children, {}) +``` + +Run: `python3 test/registered/unit/mem_cache/test_radix_cache_unit.py TestTreeNodeChildren.test_missing_child_access_does_not_create_orphan_node` + +Expected: FAIL because `defaultdict(TreeNode)` creates an orphan node instead of raising `KeyError`. + +Replace `self.children = defaultdict(TreeNode)` with `self.children = {}` in: + +```text +python/sglang/srt/mem_cache/radix_cache.py +python/sglang/srt/mem_cache/swa_radix_cache.py +python/sglang/srt/mem_cache/mamba_radix_cache.py +``` + +Remove `from collections import defaultdict` from `python/sglang/srt/mem_cache/radix_cache.py` if no other code in that file uses `defaultdict`. Repeat the same removal in `swa_radix_cache.py` and `mamba_radix_cache.py` only after replacing their `defaultdict(TreeNode)` children initialization. + +- [ ] **Step 12: Run task regression tests** + +Run: `python3 test/registered/unit/managers/test_prefill_adder.py` + +Expected: PASS. + +Run: `python3 test/registered/unit/mem_cache/test_cp_shared_kv_layout.py` + +Expected: PASS. + +- [ ] **Step 13: Commit Task 1** + +```bash +git add python/sglang/srt/managers/schedule_policy.py python/sglang/srt/mem_cache/allocator.py python/sglang/srt/mem_cache/common.py python/sglang/srt/mem_cache/base_prefix_cache.py python/sglang/srt/mem_cache/radix_cache.py python/sglang/srt/mem_cache/swa_radix_cache.py python/sglang/srt/mem_cache/mamba_radix_cache.py test/registered/unit/managers/test_prefill_adder.py test/registered/unit/mem_cache/test_cp_shared_kv_layout.py test/registered/unit/mem_cache/test_radix_cache_unit.py +git commit -m "fix: preserve HiCache load-back allocator headroom" +``` + +--- + +### Task 2: Add CP Shared KV Storage Fail-Fast + +**Files:** +- Modify: `python/sglang/srt/server_args.py` +- Modify: `test/registered/unit/server_args/test_server_args.py` + +- [ ] **Step 1: Update test base class** + +In `test/registered/unit/server_args/test_server_args.py`, change: + +```python +class TestHiCacheArgs(unittest.TestCase): +``` + +to: + +```python +class TestHiCacheArgs(CustomTestCase): +``` + +The file already imports `CustomTestCase` near the top. + +- [ ] **Step 2: Add failing fail-fast test** + +Add this method to `TestHiCacheArgs`: + +```python +def test_cp_shared_kv_rejects_hicache_storage_backend(self): + 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=None, + page_size=64, + enable_hisparse=False, + hicache_storage_backend="mooncake", + ) + + with self.assertRaisesRegex( + AssertionError, + "enable_nsa_prefill_cp_shared_kv.*hicache_storage_backend", + ): + args._handle_model_and_attention() +``` + +- [ ] **Step 3: Run the failing test** + +Run: `python3 test/registered/unit/server_args/test_server_args.py TestHiCacheArgs.test_cp_shared_kv_rejects_hicache_storage_backend` + +Expected: FAIL because the assertion is missing. + +- [ ] **Step 4: Implement fail-fast validation** + +In `python/sglang/srt/server_args.py`, inside the existing `if self.enable_nsa_prefill_cp_shared_kv:` block and after the `enable_hisparse` assertion, add: + +```python +assert self.hicache_storage_backend is None, ( + "--enable-nsa-prefill-cp-shared-kv does not support " + "hicache_storage_backend in the host-only CP HiCache stage. " + "Disable hicache_storage_backend or disable CP shared KV." +) +``` + +- [ ] **Step 5: Verify fail-fast test passes** + +Run: `python3 test/registered/unit/server_args/test_server_args.py TestHiCacheArgs.test_cp_shared_kv_rejects_hicache_storage_backend` + +Expected: PASS. + +- [ ] **Step 6: Commit Task 2** + +```bash +git add python/sglang/srt/server_args.py test/registered/unit/server_args/test_server_args.py +git commit -m "fix: reject CP shared KV with HiCache storage" +``` + +--- + +### Task 3: Add CP HiCache Metadata + +**Files:** +- Create: `test/registered/unit/mem_cache/test_cp_hicache_metadata.py` +- Modify: `python/sglang/srt/mem_cache/hiradix_cache.py` +- Modify: `python/sglang/srt/mem_cache/radix_cache.py` + +- [ ] **Step 1: Create failing metadata tests** + +Create `test/registered/unit/mem_cache/test_cp_hicache_metadata.py` with this structure: + +```python +import unittest +from unittest.mock import patch + +import torch + +from sglang.srt.mem_cache.hiradix_cache import CpHiCacheNodeMetadata, HiRadixCache +from sglang.srt.mem_cache.radix_cache import RadixKey, TreeNode +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=2, suite="stage-a-test-cpu") + + +class TestCpHiCacheNodeMetadata(CustomTestCase): + def test_split_zero_len_moves_all_positions_to_child(self): + metadata = CpHiCacheNodeMetadata( + logical_len=8, + owned_positions=torch.tensor([1, 3, 7], dtype=torch.int64), + host_indices=torch.tensor([10, 11, 12], dtype=torch.int64), + ) + + parent, child = metadata.split(0) + + self.assertEqual(parent.logical_len, 0) + self.assertEqual(parent.owned_positions.tolist(), []) + self.assertEqual(parent.host_indices.tolist(), []) + self.assertEqual(child.logical_len, 8) + self.assertEqual(child.owned_positions.tolist(), [1, 3, 7]) + self.assertEqual(child.host_indices.tolist(), [10, 11, 12]) + + def test_split_rebases_child_positions(self): + metadata = CpHiCacheNodeMetadata( + logical_len=10, + owned_positions=torch.tensor([0, 2, 5, 9], dtype=torch.int64), + host_indices=torch.tensor([20, 21, 22, 23], dtype=torch.int64), + ) + + parent, child = metadata.split(5) + + self.assertEqual(parent.logical_len, 5) + self.assertEqual(parent.owned_positions.tolist(), [0, 2]) + self.assertEqual(parent.host_indices.tolist(), [20, 21]) + self.assertEqual(child.logical_len, 5) + self.assertEqual(child.owned_positions.tolist(), [0, 4]) + self.assertEqual(child.host_indices.tolist(), [22, 23]) + + def test_zero_owned_metadata_is_valid(self): + metadata = CpHiCacheNodeMetadata( + logical_len=64, + owned_positions=torch.empty((0,), dtype=torch.int64), + host_indices=torch.empty((0,), dtype=torch.int64), + ) + + self.assertEqual(metadata.logical_len, 64) + self.assertEqual(metadata.owned_positions.device.type, "cpu") + self.assertEqual(metadata.host_indices.device.type, "cpu") + + def test_invalid_split_raises(self): + metadata = CpHiCacheNodeMetadata( + logical_len=4, + owned_positions=torch.tensor([1], dtype=torch.int64), + host_indices=torch.tensor([9], dtype=torch.int64), + ) + + with self.assertRaisesRegex(ValueError, "split_len"): + metadata.split(5) + + def test_unsorted_positions_raise(self): + with self.assertRaisesRegex(ValueError, "sorted"): + CpHiCacheNodeMetadata( + logical_len=4, + owned_positions=torch.tensor([2, 1], dtype=torch.int64), + host_indices=torch.tensor([9, 10], dtype=torch.int64), + ) + + def test_length_mismatch_raises(self): + with self.assertRaisesRegex(ValueError, "same length"): + CpHiCacheNodeMetadata( + logical_len=4, + owned_positions=torch.tensor([1, 2], dtype=torch.int64), + host_indices=torch.tensor([9], dtype=torch.int64), + ) + + def test_out_of_range_positions_raise(self): + with self.assertRaisesRegex(ValueError, r"\[0, logical_len\)"): + CpHiCacheNodeMetadata( + logical_len=4, + owned_positions=torch.tensor([4], dtype=torch.int64), + host_indices=torch.tensor([9], dtype=torch.int64), + ) + + +if __name__ == "__main__": + unittest.main() +``` + +- [ ] **Step 2: Run failing metadata tests** + +Run: `python3 test/registered/unit/mem_cache/test_cp_hicache_metadata.py TestCpHiCacheNodeMetadata` + +Expected: FAIL because `CpHiCacheNodeMetadata` does not exist. + +- [ ] **Step 3: Add metadata dataclass** + +In `python/sglang/srt/mem_cache/hiradix_cache.py`, add this import near the existing imports: + +```python +from dataclasses import dataclass +``` + +Add this dataclass immediately before `class HiRadixCache` in `python/sglang/srt/mem_cache/hiradix_cache.py`: + +```python +@dataclass +class CpHiCacheNodeMetadata: + logical_len: int + owned_positions: torch.Tensor + host_indices: torch.Tensor + + def __post_init__(self): + if self.logical_len < 0: + raise ValueError(f"logical_len must be non-negative, got {self.logical_len}") + self.owned_positions = self.owned_positions.to(device="cpu", dtype=torch.int64) + self.host_indices = self.host_indices.to(device="cpu", dtype=torch.int64) + if self.owned_positions.numel() != self.host_indices.numel(): + raise ValueError( + "owned_positions and host_indices must have same length, got " + f"{self.owned_positions.numel()} and {self.host_indices.numel()}" + ) + if self.owned_positions.numel() > 0: + if torch.any(self.owned_positions < 0) or torch.any( + self.owned_positions >= self.logical_len + ): + raise ValueError("owned_positions must be in [0, logical_len)") + if torch.any(self.owned_positions[1:] < self.owned_positions[:-1]): + raise ValueError("owned_positions must be sorted") + + def split( + self, split_len: int + ) -> tuple["CpHiCacheNodeMetadata", "CpHiCacheNodeMetadata"]: + if split_len < 0 or split_len > self.logical_len: + raise ValueError( + f"split_len must be in [0, {self.logical_len}], got {split_len}" + ) + parent_mask = self.owned_positions < split_len + child_mask = ~parent_mask + return ( + CpHiCacheNodeMetadata( + logical_len=split_len, + owned_positions=self.owned_positions[parent_mask], + host_indices=self.host_indices[parent_mask], + ), + CpHiCacheNodeMetadata( + logical_len=self.logical_len - split_len, + owned_positions=self.owned_positions[child_mask] - split_len, + host_indices=self.host_indices[child_mask], + ), + ) +``` + +- [ ] **Step 4: Add TreeNode fields** + +In `python/sglang/srt/mem_cache/radix_cache.py::TreeNode.__init__`, add: + +```python +self.host_len: int = 0 +self.cp_hicache = None +``` + +- [ ] **Step 5: Verify metadata tests pass** + +Run: `python3 test/registered/unit/mem_cache/test_cp_hicache_metadata.py TestCpHiCacheNodeMetadata` + +Expected: PASS. + +- [ ] **Step 6: Commit Task 3** + +```bash +git add python/sglang/srt/mem_cache/hiradix_cache.py python/sglang/srt/mem_cache/radix_cache.py test/registered/unit/mem_cache/test_cp_hicache_metadata.py +git commit -m "feat: add CP HiCache node metadata" +``` + +--- + +### Task 4: Add Controller CP Write Mapping + +**Files:** +- Create: `test/registered/unit/managers/test_hicache_controller_cp.py` +- Modify: `python/sglang/srt/managers/cache_controller.py` +- Modify: `python/sglang/srt/mem_cache/hiradix_cache.py` + +- [ ] **Step 1: Create fake-pool tests for CP write** + +Create `test/registered/unit/managers/test_hicache_controller_cp.py` with fakes that do not launch a server. Include `CustomTestCase` and `register_cpu_ci`. + +Create the file with this complete content: + +```python +import unittest +from unittest.mock import patch + +import torch + +from sglang.srt.managers.cache_controller import HiCacheController +from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=2, suite="stage-a-test-cpu") + + +class FakeHostPool: + def __init__(self, alloc_result): + self.alloc_result = alloc_result + self.alloc_calls = [] + self.backups = [] + self.loads = [] + self.frees = [] + self.page_size = 4 + self.layout = "page_first_direct" + + def alloc(self, need_size): + self.alloc_calls.append(need_size) + if self.alloc_result is None: + return None + return self.alloc_result[:need_size].clone() + + def backup_from_device_all_layer(self, device_pool, host_indices, device_indices, io_backend): + self.backups.append((host_indices.clone(), device_indices.clone())) + + def load_to_device_per_layer(self, device_pool, host_indices, device_indices, layer_id, io_backend): + self.loads.append((host_indices.clone(), device_indices.clone(), layer_id)) + + def free(self, indices): + self.frees.append(indices.clone()) + return len(indices) + + +class FakeDevicePool: + device = "cpu" + layer_num = 1 + + def register_layer_transfer_counter(self, counter): + self.counter = counter + + +class FakeAllocator: + def __init__(self, alloc_result=None): + self.alloc_result = alloc_result + self.alloc_calls = [] + self.cp_size = 4 + self.cp_rank = 1 + + def get_kvcache(self): + return FakeDevicePool() + + def alloc(self, need_size): + self.alloc_calls.append(need_size) + if self.alloc_result is None: + return None + return self.alloc_result[:need_size].clone() + + +class DummyEvent: + def record(self): + pass + + def wait(self, stream): + pass + + def query(self): + return True + + def synchronize(self): + pass + + +class DummyStream: + def __enter__(self): + return self + + def __exit__(self, exc_type, exc, tb): + return False + + +class DummyDeviceModule: + Event = DummyEvent + Stream = DummyStream + + @staticmethod + def stream(stream): + return stream + + +class DummyLayerDoneCounter: + def __init__(self): + self.events = [ + type( + "ProducerEvent", + (), + { + "start_event": DummyEvent(), + "finish_event": DummyEvent(), + "complete": lambda self, layer_id: None, + }, + )() + ] + + def update_producer(self): + return 0 + + +class TestHiCacheControllerCPWrite(CustomTestCase): + def setUp(self): + self.device_module_patcher = patch( + "sglang.srt.managers.cache_controller.device_module", + DummyDeviceModule, + ) + self.nsa_pool_patcher = patch( + "sglang.srt.managers.cache_controller.NSATokenToKVPool", + FakeDevicePool, + ) + self.device_module_patcher.start() + self.nsa_pool_patcher.start() + self.addCleanup(self.device_module_patcher.stop) + self.addCleanup(self.nsa_pool_patcher.stop) + + def make_controller(self, host_pool, allocator=None, cp_rank=1): + allocator = allocator or FakeAllocator() + controller = HiCacheController( + token_to_kv_pool_allocator=allocator, + mem_pool_host=host_pool, + page_size=4, + tp_group=None, + load_cache_event=__import__("threading").Event(), + io_backend="direct", + cp_shared_kv_layout=CpSharedKVLayout( + page_size=4, cp_size=4, cp_rank=cp_rank + ), + ) + controller.layer_done_counter = DummyLayerDoneCounter() + return controller + + def test_cp_write_filters_to_owned_physical_locs(self): + host_pool = FakeHostPool(torch.tensor([100, 101, 102, 103], dtype=torch.int64)) + controller = self.make_controller(host_pool, cp_rank=1) + logical_locs = torch.arange(4, 20, dtype=torch.int64) + + result = controller.write(logical_locs, node_id=7) + + self.assertEqual(result.metadata.logical_len, 16) + self.assertEqual(result.metadata.owned_positions.tolist(), [4, 5, 6, 7]) + self.assertEqual(result.metadata.host_indices.tolist(), [100, 101, 102, 103]) + self.assertEqual(host_pool.alloc_calls, [4]) + self.assertEqual(host_pool.backups[0][1].tolist(), [4, 5, 6, 7]) + + def test_cp_write_zero_owned_returns_metadata_and_noop_ack(self): + host_pool = FakeHostPool(torch.tensor([], dtype=torch.int64)) + controller = self.make_controller(host_pool, cp_rank=3) + logical_locs = torch.arange(4, 8, dtype=torch.int64) + + result = controller.write(logical_locs, node_id=8) + + self.assertEqual(result.metadata.logical_len, 4) + self.assertEqual(result.metadata.host_indices.tolist(), []) + self.assertEqual(host_pool.alloc_calls, []) + self.assertEqual(len(controller.ack_write_queue), 1) + + def test_cp_write_allocation_failure_reports_required_host_slots(self): + host_pool = FakeHostPool(None) + controller = self.make_controller(host_pool, cp_rank=1) + logical_locs = torch.arange(4, 20, dtype=torch.int64) + + result = controller.write(logical_locs, node_id=9) + + self.assertEqual(result.required_host_slots, 4) + + +if __name__ == "__main__": + unittest.main() +``` + +Use a fake device pool with `device="cpu"`, `layer_num=1`, and `register_layer_transfer_counter` no-op. `make_controller` sets `controller.layer_done_counter = DummyLayerDoneCounter()` so zero-owned load tests can verify no-op load acks without real device events. + +- [ ] **Step 2: Run failing CP write tests** + +Run: `python3 test/registered/unit/managers/test_hicache_controller_cp.py TestHiCacheControllerCPWrite` + +Expected: FAIL because CP adapter methods/results do not exist. + +- [ ] **Step 3: Add CP result types** + +In `python/sglang/srt/managers/cache_controller.py`, extend the imports at the top: + +```python +from dataclasses import dataclass +``` + +Add this import beside the existing `MLATokenToKVPool` import at the top of `python/sglang/srt/managers/cache_controller.py`: + +```python +from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool, NSATokenToKVPool +``` + +Add this non-type-checking import near the existing `hicache_storage` import because `__init__` uses it at runtime: + +```python +from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout +``` + +Add these dataclasses before `class HiCacheController`: + +```python +@dataclass +class HiCacheWriteResult: + metadata: object + required_host_slots: int = 0 + + +@dataclass +class HiCacheWriteFailure: + required_host_slots: int + metadata: object = None +``` + +Use these names consistently in `HiRadixCache.write_backup()`: success results expose `metadata`, and allocation failures expose `required_host_slots`. + +- [ ] **Step 4: Add CP adapter state** + +Extend `HiCacheController.__init__` with this keyword parameter after `enable_storage_metrics`: + +```python +cp_shared_kv_layout: Optional[CpSharedKVLayout] = None, +``` + +Set these fields after `self.mem_pool_device = mem_pool_device`: + +```python +self.cp_shared_kv_layout = cp_shared_kv_layout +self.uses_cp_hicache = cp_shared_kv_layout is not None +``` + +Then add this validation: + +```python +if self.uses_cp_hicache and not isinstance(self.mem_pool_device, NSATokenToKVPool): + raise ValueError( + "CP shared KV HiCache host integration requires NSATokenToKVPool." + ) +``` + +The tests in Step 1 patch `sglang.srt.managers.cache_controller.NSATokenToKVPool` to `FakeDevicePool` so the production validation is exercised without constructing a real NSA pool. They also patch `device_module` so CPU CI does not create real CUDA streams or events. + +- [ ] **Step 5: Implement CP write branch** + +At the top of `write()`, if `self.uses_cp_hicache`, call a new private method: + +```python +return self._write_cp(device_indices=device_indices, priority=priority, node_id=node_id) +``` + +Implement `_append_completed_write_ack` and `_write_cp` in `HiCacheController`: + +```python +def _append_completed_write_ack(self, node_id: int) -> None: + event = device_module.Event() + event.record() + self.ack_write_queue.append(HiCacheAck(event, event, [node_id])) + +def _append_completed_load_ack(self, node_id: int) -> int: + producer_id = self.layer_done_counter.update_producer() + producer_event = self.layer_done_counter.events[producer_id] + producer_event.start_event.record() + for layer_id in range(self.layer_num): + producer_event.complete(layer_id) + self.ack_load_queue.append( + HiCacheAck( + start_event=producer_event.start_event, + finish_event=producer_event.finish_event, + node_ids=[node_id], + ) + ) + return producer_id + +def _write_cp( + self, + device_indices: torch.Tensor, + priority: Optional[int] = None, + node_id: int = -1, +) -> HiCacheWriteResult | HiCacheWriteFailure: + from sglang.srt.mem_cache.hiradix_cache import CpHiCacheNodeMetadata + + owned_mask = self.cp_shared_kv_layout.owned_by_this_rank(device_indices) + owned_positions = owned_mask.nonzero(as_tuple=True)[0].cpu() + owned_logical_indices = device_indices[owned_mask] + if owned_logical_indices.numel() == 0: + self._append_completed_write_ack(node_id) + return HiCacheWriteResult( + metadata=CpHiCacheNodeMetadata( + logical_len=len(device_indices), + owned_positions=owned_positions, + host_indices=torch.empty((0,), dtype=torch.int64), + ) + ) + + physical_device_indices = self.cp_shared_kv_layout.logical_locs_to_physical( + owned_logical_indices + ) + host_indices = self.mem_pool_host.alloc(len(physical_device_indices)) + if host_indices is None: + return HiCacheWriteFailure(required_host_slots=len(physical_device_indices)) + + self.write_queue.append( + CacheOperation(host_indices, physical_device_indices, node_id, priority) + ) + self.start_writing() + return HiCacheWriteResult( + metadata=CpHiCacheNodeMetadata( + logical_len=len(device_indices), + owned_positions=owned_positions, + host_indices=host_indices.cpu(), + ) + ) +``` + +Keep the existing non-CP `write()` behavior unchanged when `self.uses_cp_hicache` is false. + +- [ ] **Step 6: Verify CP write tests pass** + +Run: `python3 test/registered/unit/managers/test_hicache_controller_cp.py TestHiCacheControllerCPWrite` + +Expected: PASS. + +- [ ] **Step 7: Commit Task 4** + +```bash +git add python/sglang/srt/managers/cache_controller.py python/sglang/srt/mem_cache/hiradix_cache.py test/registered/unit/managers/test_hicache_controller_cp.py +git commit -m "feat: map CP HiCache writes to owned physical slots" +``` + +--- + +### Task 5: Wire HiRadixCache CP Write State + +**Files:** +- Modify: `python/sglang/srt/mem_cache/hiradix_cache.py` +- Modify: `test/registered/unit/mem_cache/test_cp_hicache_metadata.py` + +- [ ] **Step 1: Add failing backup-state tests** + +Add this `TestHiRadixCacheCPBackup` class to `test/registered/unit/mem_cache/test_cp_hicache_metadata.py` after `TestCpHiCacheNodeMetadata`: + +```python +class TestHiRadixCacheCPBackup(CustomTestCase): + def test_node_backuped_uses_cp_metadata(self): + cache = HiRadixCache.__new__(HiRadixCache) + cache._uses_cp_hicache = True + node = TreeNode() + node.host_len = 8 + node.cp_hicache = CpHiCacheNodeMetadata( + logical_len=8, + owned_positions=torch.tensor([1, 2], dtype=torch.int64), + host_indices=torch.tensor([10, 11], dtype=torch.int64), + ) + + self.assertTrue(cache._node_backuped(node)) + + def test_inc_hit_count_does_not_rewrite_cp_backed_node(self): + cache = HiRadixCache.__new__(HiRadixCache) + cache._uses_cp_hicache = True + cache.write_through_threshold = 1 + cache.cache_controller = type("Controller", (), {"write_policy": "write_through"})() + cache.write_backup = lambda node: (_ for _ in ()).throw(AssertionError("must not rewrite")) + node = TreeNode() + node.host_len = 4 + node.cp_hicache = CpHiCacheNodeMetadata( + logical_len=4, + owned_positions=torch.tensor([], dtype=torch.int64), + host_indices=torch.tensor([], dtype=torch.int64), + ) + + cache._inc_hit_count(node) + + self.assertEqual(node.hit_count, 1) + + def test_write_backup_retries_by_required_physical_slots(self): + cache = HiRadixCache.__new__(HiRadixCache) + cache._uses_cp_hicache = True + cache.cache_controller = FakeWriteController(required_host_slots=3) + cache._evict_host_for_physical_slots = lambda required: setattr(cache, "evicted_required", required) + cache.ongoing_write_through = {} + cache.inc_node_lock_ref = lambda node: None + node = TreeNode() + node.value = torch.arange(16, dtype=torch.int64) + + cache.write_backup(node) + + self.assertEqual(cache.evicted_required, 3) + self.assertEqual(node.host_len, 16) +``` + +Define `FakeWriteController` in the test file with a first `write()` call returning an allocation-failure object containing `required_host_slots=3`, and a second `write()` call returning metadata. + +Add these fake classes before `TestHiRadixCacheCPBackup` in `test/registered/unit/mem_cache/test_cp_hicache_metadata.py`: + +```python +class FakeWriteFailure: + metadata = None + + def __init__(self, required_host_slots): + self.required_host_slots = required_host_slots + + +class FakeWriteSuccess: + required_host_slots = 0 + + def __init__(self, metadata): + self.metadata = metadata + + +class FakeWriteController: + def __init__(self, required_host_slots): + self.required_host_slots = required_host_slots + self.calls = 0 + self.write_policy = "write_through" + + def write(self, device_indices, node_id=-1, priority=None): + self.calls += 1 + if self.calls == 1: + return FakeWriteFailure(self.required_host_slots) + return FakeWriteSuccess( + CpHiCacheNodeMetadata( + logical_len=len(device_indices), + owned_positions=torch.tensor([0], dtype=torch.int64), + host_indices=torch.tensor([99], dtype=torch.int64), + ) + ) +``` + +The backup tests use `HiRadixCache.__new__(HiRadixCache)` to avoid server initialization and set every attribute they read inside the test body. + +- [ ] **Step 2: Run failing backup-state tests** + +Run: `python3 test/registered/unit/mem_cache/test_cp_hicache_metadata.py TestHiRadixCacheCPBackup` + +Expected: FAIL because helpers are missing or old `node.backuped` path is used. + +- [ ] **Step 3: Add CP init wiring** + +In `python/sglang/srt/mem_cache/hiradix_cache.py`, add these imports near the existing mem-cache imports: + +```python +from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator +from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout +``` + +In `HiRadixCache.__init__`, after `self.page_size = params.page_size`, determine CP mode: + +```python +self._uses_cp_hicache = isinstance( + params.token_to_kv_pool_allocator, + CPSharedPagedTokenToKVPoolAllocator, +) +if self._uses_cp_hicache: + if server_args.hicache_storage_backend is not None: + raise ValueError( + "CP shared KV HiCache host integration does not support storage backends." + ) + if not isinstance(params.token_to_kv_pool_allocator.get_kvcache(), NSATokenToKVPool): + raise ValueError( + "CP shared KV HiCache host integration requires NSATokenToKVPool." + ) +``` + +Immediately before constructing `HiCacheController`, build: + +```python +cp_shared_kv_layout = None +if self._uses_cp_hicache: + cp_shared_kv_layout = CpSharedKVLayout( + page_size=self.page_size, + cp_size=params.token_to_kv_pool_allocator.cp_size, + cp_rank=params.token_to_kv_pool_allocator.cp_rank, + ) +``` + +Pass it into `HiCacheController`: + +```python +cp_shared_kv_layout=cp_shared_kv_layout, +``` + +- [ ] **Step 4: Add backup helpers** + +Add methods with these exact semantics: + +```python +def _node_backuped(self, node: TreeNode) -> bool: + if self._uses_cp_hicache: + return node.host_len > 0 and node.cp_hicache is not None + return node.host_value is not None + +def _node_host_len(self, node: TreeNode) -> int: + if self._uses_cp_hicache: + return node.host_len + return len(node.host_value) + +def _node_host_evict_indices(self, node: TreeNode) -> torch.Tensor: + if self._uses_cp_hicache: + return node.cp_hicache.host_indices + return node.host_value +``` + +- [ ] **Step 5: Update write-through hit path** + +In `_inc_hit_count`, replace: + +```python +if not node.backuped: +``` + +with: + +```python +if not self._node_backuped(node): +``` + +- [ ] **Step 6: Update `write_backup` CP branch** + +At the top of `HiRadixCache.write_backup`, add this CP branch before the current non-CP implementation: + +```python +if self._uses_cp_hicache: + result = self.cache_controller.write( + device_indices=node.value, + node_id=node.id, + ) + if getattr(result, "metadata", None) is None: + self._evict_host_for_physical_slots(result.required_host_slots) + result = self.cache_controller.write( + device_indices=node.value, + node_id=node.id, + ) + if getattr(result, "metadata", None) is None: + return 0 + + node.host_len = len(node.value) + node.cp_hicache = result.metadata + node.host_value = None + self.ongoing_write_through[node.id] = node + if not write_back: + self.inc_node_lock_ref(node) + return len(node.cp_hicache.host_indices) +``` + +On success this branch sets: + +```python +node.host_len = len(node.value) +node.cp_hicache = result.metadata +node.host_value = None +``` + +On allocation failure, it calls a CP host eviction helper that evicts until physical freed slots reach `result.required_host_slots`. + +- [ ] **Step 7: Verify backup-state tests pass** + +Run: `python3 test/registered/unit/mem_cache/test_cp_hicache_metadata.py TestHiRadixCacheCPBackup` + +Expected: PASS. + +- [ ] **Step 8: Commit Task 5** + +```bash +git add python/sglang/srt/mem_cache/hiradix_cache.py test/registered/unit/mem_cache/test_cp_hicache_metadata.py +git commit -m "feat: track CP HiCache write-back state" +``` + +--- + +### Task 6: Wire CP Split and Host Eviction + +**Files:** +- Modify: `python/sglang/srt/mem_cache/hiradix_cache.py` +- Modify: `test/registered/unit/mem_cache/test_cp_hicache_metadata.py` + +- [ ] **Step 1: Add failing split and eviction tests** + +Add this class to `test/registered/unit/mem_cache/test_cp_hicache_metadata.py`: + +```python +class TestHiRadixCacheCPSplitEvict(CustomTestCase): + def test_split_node_splits_cp_metadata_by_owned_positions(self): + cache = HiRadixCache.__new__(HiRadixCache) + cache._uses_cp_hicache = True + cache.get_child_key_fn = lambda key: key.token_ids[0] + root = TreeNode() + root.key = RadixKey([]) + child = TreeNode() + child.parent = root + child.key = RadixKey(list(range(10))) + child.value = None + child.host_len = 10 + child.cp_hicache = CpHiCacheNodeMetadata( + logical_len=10, + owned_positions=torch.tensor([0, 2, 5, 9], dtype=torch.int64), + host_indices=torch.tensor([20, 21, 22, 23], dtype=torch.int64), + ) + root.children[0] = child + + new_node = cache._split_node(child.key, child, 5) + + self.assertEqual(new_node.host_len, 5) + self.assertEqual(child.host_len, 5) + self.assertEqual(new_node.cp_hicache.owned_positions.tolist(), [0, 2]) + self.assertEqual(new_node.cp_hicache.host_indices.tolist(), [20, 21]) + self.assertEqual(child.cp_hicache.owned_positions.tolist(), [0, 4]) + self.assertEqual(child.cp_hicache.host_indices.tolist(), [22, 23]) + + def test_cp_host_eviction_uses_physical_freed_slots_for_progress(self): + cache = HiRadixCache.__new__(HiRadixCache) + cache._uses_cp_hicache = True + cache.root_node = TreeNode() + cache.root_node.key = RadixKey([]) + cache.evictable_host_leaves = set() + cache.get_child_key_fn = lambda key: key.token_ids[0] + cache.eviction_strategy = type( + "Strategy", (), {"get_priority": lambda self, node: 0} + )() + cache._clear_pin = lambda node: None + cache._record_remove_event = lambda node: None + cache._update_host_leaf_status = lambda node: None + freed = [] + cache.cache_controller = type( + "Controller", + (), + {"evict_host": lambda self, indices: freed.append(indices.clone()) or len(indices)}, + )() + node = TreeNode() + node.parent = cache.root_node + node.key = RadixKey([1, 2, 3, 4]) + node.value = None + node.host_len = 4 + node.cp_hicache = CpHiCacheNodeMetadata( + logical_len=4, + owned_positions=torch.tensor([1], dtype=torch.int64), + host_indices=torch.tensor([70], dtype=torch.int64), + ) + cache.root_node.children[1] = node + cache.evictable_host_leaves.add(node) + + physical_freed = cache._evict_host_for_physical_slots(1) + + self.assertEqual(physical_freed, 1) + self.assertEqual(freed[0].tolist(), [70]) + self.assertEqual(node.host_len, 0) + self.assertIsNone(node.cp_hicache) +``` + +- [ ] **Step 2: Run failing tests** + +Run: `python3 test/registered/unit/mem_cache/test_cp_hicache_metadata.py TestHiRadixCacheCPSplitEvict` + +Expected: FAIL. + +- [ ] **Step 3: Update `_split_node`** + +In CP mode, when `self._node_backuped(child)` is true, call: + +```python +new_node.cp_hicache, child.cp_hicache = child.cp_hicache.split(split_len) +new_node.host_len = split_len +child.host_len = child.host_len - split_len +``` + +Keep non-CP `host_value` slicing unchanged. + +- [ ] **Step 4: Update host eviction** + +Add this helper to `HiRadixCache` near `evict_host`: + +```python +def _clear_cp_host_state(self, node: TreeNode) -> tuple[int, int]: + physical_indices = node.cp_hicache.host_indices + physical_count = len(physical_indices) + logical_count = node.host_len + if physical_count > 0: + self.cache_controller.evict_host(physical_indices) + node.host_len = 0 + node.cp_hicache = None + return physical_count, logical_count +``` + +Add this retry helper to `HiRadixCache` near `evict_host`: + +```python +def _evict_host_for_physical_slots(self, required_host_slots: int) -> int: + if required_host_slots <= 0: + return 0 + + leaves = list(self.evictable_host_leaves) + eviction_heap = [ + (self.eviction_strategy.get_priority(node), node) for node in leaves + ] + heapq.heapify(eviction_heap) + + physical_freed = 0 + while physical_freed < required_host_slots and eviction_heap: + _priority, node = heapq.heappop(eviction_heap) + if node == self.root_node: + break + if not node.evicted or not self._node_backuped(node): + continue + if node.pin_expiry > 0 and time.monotonic() > node.pin_expiry: + self._clear_pin(node) + if node.host_ref_counter > 0: + continue + + self._record_remove_event(node) + freed, _logical_count = self._clear_cp_host_state(node) + physical_freed += freed + + key = self.get_child_key_fn(node.key) + removed = node.parent.children.pop(key, None) + assert removed == node, f"parent does not have child key, {key}" + self.evictable_host_leaves.discard(node) + self._update_host_leaf_status(node.parent) + + if len(node.parent.children) == 0 and node.parent.evicted: + new_priority = self.eviction_strategy.get_priority(node.parent) + heapq.heappush(eviction_heap, (new_priority, node.parent)) + + return physical_freed +``` + +At the top of `evict_host`, add: + +```python +if self._uses_cp_hicache: + return self._evict_host_for_physical_slots(num_tokens) +``` + +The CP branch frees `node.cp_hicache.host_indices`, not `node.host_value`, and tracks physical freed slots for write retry progress. The legacy non-CP body remains unchanged. + +- [ ] **Step 5: Verify split and eviction tests pass** + +Run: `python3 test/registered/unit/mem_cache/test_cp_hicache_metadata.py TestHiRadixCacheCPSplitEvict` + +Expected: PASS. + +- [ ] **Step 6: Commit Task 6** + +```bash +git add python/sglang/srt/mem_cache/hiradix_cache.py test/registered/unit/mem_cache/test_cp_hicache_metadata.py +git commit -m "feat: split and evict CP HiCache metadata" +``` + +--- + +### Task 7: Add Controller CP Load Mapping + +**Files:** +- Modify: `python/sglang/srt/managers/cache_controller.py` +- Modify: `test/registered/unit/managers/test_hicache_controller_cp.py` + +- [ ] **Step 1: Add failing CP load tests** + +In `test/registered/unit/managers/test_hicache_controller_cp.py`, update the imports from `hiradix_cache` and `radix_cache` if they are not already present: + +```python +from sglang.srt.mem_cache.hiradix_cache import CpHiCacheNodeMetadata +from sglang.srt.mem_cache.radix_cache import TreeNode +``` + +Add this class after `TestHiCacheControllerCPWrite` and before the `if __name__ == "__main__":` block: + +```python +class TestHiCacheControllerCPLoad(TestHiCacheControllerCPWrite): + def test_cp_load_allocates_full_logical_locs_and_transfers_owned_physical_locs(self): + host_pool = FakeHostPool(torch.tensor([100, 101, 102, 103], dtype=torch.int64)) + allocator = FakeAllocator(alloc_result=torch.arange(64, 80, dtype=torch.int64)) + controller = self.make_controller(host_pool, allocator=allocator, cp_rank=1) + node = TreeNode() + node.host_len = 16 + node.cp_hicache = CpHiCacheNodeMetadata( + logical_len=16, + owned_positions=torch.tensor([4, 5, 6, 7], dtype=torch.int64), + host_indices=torch.tensor([100, 101, 102, 103], dtype=torch.int64), + ) + + device_indices = controller.load_cp([node], node_id=11) + controller.start_loading() + + self.assertEqual(device_indices.tolist(), list(range(64, 80))) + self.assertEqual(allocator.alloc_calls, [16]) + self.assertEqual(host_pool.loads[0][1].tolist(), [16, 17, 18, 19]) + + def test_cp_load_zero_owned_returns_full_logical_locs_and_noop_ack(self): + host_pool = FakeHostPool(torch.tensor([], dtype=torch.int64)) + allocator = FakeAllocator(alloc_result=torch.arange(64, 68, dtype=torch.int64)) + controller = self.make_controller(host_pool, allocator=allocator, cp_rank=3) + node = TreeNode() + node.host_len = 4 + node.cp_hicache = CpHiCacheNodeMetadata( + logical_len=4, + owned_positions=torch.empty((0,), dtype=torch.int64), + host_indices=torch.empty((0,), dtype=torch.int64), + ) + + device_indices = controller.load_cp([node], node_id=12) + + self.assertEqual(device_indices.tolist(), [64, 65, 66, 67]) + self.assertEqual(host_pool.loads, []) + self.assertEqual(len(controller.ack_load_queue), 1) +``` + +- [ ] **Step 2: Run failing CP load tests** + +Run: `python3 test/registered/unit/managers/test_hicache_controller_cp.py TestHiCacheControllerCPLoad` + +Expected: FAIL because `load_cp` is missing. + +- [ ] **Step 3: Implement `load_cp(nodes_to_load, node_id=-1)`** + +In `python/sglang/srt/managers/cache_controller.py`, add: + +```python +def load_cp(self, nodes_to_load, node_id: int = -1) -> Optional[torch.Tensor]: + logical_len = sum(node.host_len for node in nodes_to_load) + device_indices = self.mem_pool_device_allocator.alloc(logical_len) + if device_indices is None: + return None + + host_chunks = [] + physical_chunks = [] + offset = 0 + for node in nodes_to_load: + node_device_indices = device_indices[offset : offset + node.host_len] + offset += node.host_len + owned_positions = node.cp_hicache.owned_positions.to(device_indices.device) + if owned_positions.numel() == 0: + continue + selected_logical_locs = node_device_indices[owned_positions] + physical_chunks.append( + self.cp_shared_kv_layout.logical_locs_to_physical(selected_logical_locs) + ) + host_chunks.append(node.cp_hicache.host_indices.to(device_indices.device)) + + if not host_chunks: + self._append_completed_load_ack(node_id) + return device_indices + + self.load_queue.append( + CacheOperation( + torch.cat(host_chunks), + torch.cat(physical_chunks), + node_id, + ) + ) + + return device_indices +``` + +Implementation requirements: + +- `logical_len = sum(node.host_len for node in nodes_to_load)`. +- `device_indices = self.mem_pool_device_allocator.alloc(logical_len)`. +- Return `None` on allocation failure. +- For each node, slice `device_indices` by `node.host_len`. +- Select local logical locs by `node.cp_hicache.owned_positions`. +- Convert selected locs to physical locs. +- Concatenate host indices and physical device locs. +- Queue real transfer when owned count is nonzero. Do not call `start_loading()` inside `load_cp`; `ready_to_load_host_cache()` calls `start_loading()` after the scheduler prepares the batch. +- Produce a completed no-op load ack or producer event when owned count is zero by calling `_append_completed_load_ack(node_id)`. +- Return full logical `device_indices`. + +- [ ] **Step 4: Verify CP load tests pass** + +Run: `python3 test/registered/unit/managers/test_hicache_controller_cp.py TestHiCacheControllerCPLoad` + +Expected: PASS. + +- [ ] **Step 5: Commit Task 7** + +```bash +git add python/sglang/srt/managers/cache_controller.py test/registered/unit/managers/test_hicache_controller_cp.py +git commit -m "feat: map CP HiCache loads to owned physical slots" +``` + +--- + +### Task 8: Wire HiRadixCache CP Load Back + +**Files:** +- Modify: `python/sglang/srt/mem_cache/hiradix_cache.py` +- Modify: `test/registered/unit/mem_cache/test_cp_hicache_metadata.py` + +- [ ] **Step 1: Add failing CP load-back tests** + +Add `TestHiRadixCacheCPLoadBack` with tests using `HiRadixCache.__new__(HiRadixCache)` and fake controllers returning deterministic full logical `device_indices`: + +```python +class TestHiRadixCacheCPLoadBack(CustomTestCase): + def test_cp_load_back_uses_host_len_not_host_value(self): + cache = HiRadixCache.__new__(HiRadixCache) + cache._uses_cp_hicache = True + cache.root_node = TreeNode() + cache.device = "cpu" + cache.load_back_threshold = 1 + cache.evictable_size_ = 0 + cache.metrics_collector = None + cache.ongoing_load_back = {} + cache.cache_controller = type( + "Controller", + (), + {"load_cp": lambda self, nodes, node_id=-1: torch.arange(32, 40, dtype=torch.int64)}, + )() + cache.inc_lock_ref = lambda node: type("Result", (), {"delta": 0})() + cache.dec_lock_ref = lambda node: None + cache.evict = lambda params: None + parent = cache.root_node + parent.key = RadixKey([]) + parent.value = torch.empty((0,), dtype=torch.int64) + node = TreeNode() + node.parent = parent + node.key = RadixKey(list(range(8))) + node.value = None + node.host_value = None + node.host_len = 8 + node.cp_hicache = CpHiCacheNodeMetadata( + logical_len=8, + owned_positions=torch.tensor([0, 1], dtype=torch.int64), + host_indices=torch.tensor([50, 51], dtype=torch.int64), + ) + + loaded = cache.load_back(node) + + self.assertEqual(loaded.tolist(), list(range(32, 40))) + self.assertEqual(node.value.tolist(), list(range(32, 40))) + + def test_cp_load_back_threshold_uses_logical_length(self): + cache = HiRadixCache.__new__(HiRadixCache) + cache._uses_cp_hicache = True + cache.root_node = TreeNode() + cache.root_node.value = torch.empty((0,), dtype=torch.int64) + cache.load_back_threshold = 5 + cache.evictable_size_ = 0 + cache.metrics_collector = None + cache.ongoing_load_back = {} + cache.inc_lock_ref = lambda node: type("Result", (), {"delta": 0})() + cache.dec_lock_ref = lambda node: None + cache.cache_controller = type( + "Controller", + (), + {"load_cp": lambda self, nodes, node_id=-1: torch.arange(10, 16, dtype=torch.int64)}, + )() + node = TreeNode() + node.parent = cache.root_node + node.value = None + node.host_len = 6 + node.cp_hicache = CpHiCacheNodeMetadata( + logical_len=6, + owned_positions=torch.empty((0,), dtype=torch.int64), + host_indices=torch.empty((0,), dtype=torch.int64), + ) + + loaded = cache.load_back(node) + + self.assertEqual(loaded.tolist(), [10, 11, 12, 13, 14, 15]) +``` + +- [ ] **Step 2: Run failing load-back tests** + +Run: `python3 test/registered/unit/mem_cache/test_cp_hicache_metadata.py TestHiRadixCacheCPLoadBack` + +Expected: FAIL because current load-back reads `node.host_value`. + +- [ ] **Step 3: Implement CP load-back branch** + +At the top of `HiRadixCache.load_back`, keep the current non-CP body under `if not self._uses_cp_hicache:` and add this CP branch: + +```python +if self._uses_cp_hicache: + start_time = time.perf_counter() + last_hit_node = node + nodes_to_load = [] + while node.evicted: + assert self._node_backuped( + node + ), "No backup available on evicted nodes, should not happen" + nodes_to_load.insert(0, node) + node = node.parent + else: + ancester_node = node + + result = self.inc_lock_ref(ancester_node) + delta = result.delta + host_hit_len = sum(self._node_host_len(n) for n in nodes_to_load) + if host_hit_len < self.load_back_threshold or ( + host_hit_len > mem_quota + delta if mem_quota is not None else False + ): + self.dec_lock_ref(ancester_node) + return None + + device_indices = self.cache_controller.load_cp( + nodes_to_load, node_id=last_hit_node.id + ) + if device_indices is None: + self.evict(EvictParams(num_tokens=host_hit_len)) + device_indices = self.cache_controller.load_cp( + nodes_to_load, node_id=last_hit_node.id + ) + self.dec_lock_ref(ancester_node) + if device_indices is None: + logger.warning( + "load_back: FAILED to load %d CP logical tokens for node %d " + "even after eviction (evictable_size=%d)", + host_hit_len, + last_hit_node.id, + self.evictable_size_, + ) + return None + + self.ongoing_load_back[last_hit_node.id] = last_hit_node + offset = 0 + for loaded_node in nodes_to_load: + host_len = self._node_host_len(loaded_node) + loaded_node.value = device_indices[offset : offset + host_len].clone() + offset += host_len + self.evictable_size_ += len(device_indices) + self.inc_lock_ref(last_hit_node) + + if self.metrics_collector is not None: + self.metrics_collector.observe_load_back_duration( + time.perf_counter() - start_time + ) + self.metrics_collector.increment_load_back_num_tokens(len(device_indices)) + + return device_indices +``` + +This branch uses logical `host_len` for thresholding, memory quota, eviction retry size, metrics, and restored `TreeNode.value` slices. The existing non-CP branch continues to use `node.host_value`. + +- [ ] **Step 4: Verify CP load-back tests pass** + +Run: `python3 test/registered/unit/mem_cache/test_cp_hicache_metadata.py TestHiRadixCacheCPLoadBack` + +Expected: PASS. + +- [ ] **Step 5: Commit Task 8** + +```bash +git add python/sglang/srt/mem_cache/hiradix_cache.py test/registered/unit/mem_cache/test_cp_hicache_metadata.py +git commit -m "feat: restore CP HiCache host hits as logical locs" +``` + +--- + +### Task 9: Integration Regression + +**Files:** +- All modified files. + +- [ ] **Step 1: Run new CP HiCache metadata tests** + +Run: `python3 test/registered/unit/mem_cache/test_cp_hicache_metadata.py` + +Expected: PASS. + +- [ ] **Step 2: Run new CP controller tests** + +Run: `python3 test/registered/unit/managers/test_hicache_controller_cp.py` + +Expected: PASS. + +- [ ] **Step 3: Run server args tests** + +Run: `python3 test/registered/unit/server_args/test_server_args.py` + +Expected: PASS. + +- [ ] **Step 4: Run prefill adder tests** + +Run: `python3 test/registered/unit/managers/test_prefill_adder.py` + +Expected: PASS. + +- [ ] **Step 5: Run CP shared KV layout tests** + +Run: `python3 test/registered/unit/mem_cache/test_cp_shared_kv_layout.py` + +Expected: PASS. + +- [ ] **Step 6: Run radix cache tests** + +Run: `python3 test/registered/unit/mem_cache/test_radix_cache_unit.py` + +Expected: PASS. + +- [ ] **Step 7: Run import smoke** + +Run: `python3 -c "from sglang.srt.mem_cache.hiradix_cache import HiRadixCache, CpHiCacheNodeMetadata; from sglang.srt.managers.cache_controller import HiCacheController; print('OK')"` + +Expected: prints `OK`. + +- [ ] **Step 8: Commit integration verification fixes** + +Check the diff from integration verification: + +Run: `git diff -- python/sglang/srt/mem_cache/hiradix_cache.py python/sglang/srt/managers/cache_controller.py test/registered/unit/mem_cache/test_cp_hicache_metadata.py test/registered/unit/managers/test_hicache_controller_cp.py` + +Expected: either no output, or only fixes made while getting Steps 1-7 green. + +When the diff is non-empty, commit those fixes with: + +```bash +git add python/sglang/srt/mem_cache/hiradix_cache.py python/sglang/srt/managers/cache_controller.py test/registered/unit/mem_cache/test_cp_hicache_metadata.py test/registered/unit/managers/test_hicache_controller_cp.py +git commit -m "test: cover CP HiCache host integration" +``` + +When the diff is empty, do not create an empty commit. + +--- + +### Task 10: Record Deferred Performance Work + +**Files:** +- Modify: `docs/superpowers/specs/2026-05-07-cp-hicache-host-design.md` + +- [ ] **Step 1: Confirm Task 9 is green** + +Run the full Task 9 command list. Expected: every command reports PASS or has a recorded unrelated failure with file and test name. + +- [ ] **Step 2: Add explicit follow-up note** + +Search for the section first: + +Run: `git grep -n "Deferred Performance Follow-Up" -- docs/superpowers/specs/2026-05-07-cp-hicache-host-design.md` + +Expected before this step on the current spec: no output. + +Append this section to `docs/superpowers/specs/2026-05-07-cp-hicache-host-design.md`: + +```markdown +## Deferred Performance Follow-Up + +After host-only CP HiCache correctness lands, port HiCache performance work from +`25b6a957c` as separate changes with their own tests. Split cached timestamp use, +incremental eviction heap maintenance, persistent all-reduce tensors, and async +write acknowledgement into independent patches. Each patch must preserve CP +zero-owned write/load no-op acknowledgement cardinality and must rerun the CP +HiCache unit regression commands from the implementation plan. +``` + +- [ ] **Step 3: Commit deferral note if changed** + +Run this status check: + +Run: `git diff -- docs/superpowers/specs/2026-05-07-cp-hicache-host-design.md` + +Expected: diff contains the `Deferred Performance Follow-Up` section from Step 2. + +When the diff is non-empty, commit the note: + +```bash +git add docs/superpowers/specs/2026-05-07-cp-hicache-host-design.md +git commit -m "docs: defer CP HiCache performance follow-up" +``` + +When the diff is empty, do not create an empty commit. + +--- + +## Final Verification Checklist + +- [ ] Every new production method has a test that failed before implementation. +- [ ] `HostKVCache` remains unaware of `CpSharedKVLayout`. +- [ ] `TreeNode.value` remains full logical locs in CP mode. +- [ ] `TreeNode.host_value` remains legacy non-CP host indices and is not used for CP-owned subsets. +- [ ] CP host write retry progress uses physical freed host slots. +- [ ] CP write/load zero-owned ranks produce completed no-op logical ack/event behavior. +- [ ] CP load-back threshold uses logical length by design. +- [ ] CP shared KV + HiCache storage backend fails fast. +- [ ] New tests inherit from `CustomTestCase`. +- [ ] Required regression commands in Task 9 pass or have documented unrelated failures. + +## Execution Options + +1. **Subagent-Driven, recommended:** dispatch a fresh subagent per task and review after each task. +2. **Inline Execution:** execute tasks in this session using `executing-plans`, with checkpoints after each task. diff --git a/docs/superpowers/specs/2026-05-07-cp-hicache-host-design.md b/docs/superpowers/specs/2026-05-07-cp-hicache-host-design.md new file mode 100644 index 000000000..4df93ed95 --- /dev/null +++ b/docs/superpowers/specs/2026-05-07-cp-hicache-host-design.md @@ -0,0 +1,429 @@ +# CP Shared KV + HiCache Host Integration Design + +**Goal:** Deliver a first reliable integration between NSA CP shared KV and HiCache host caching for GLM 5 inference. + +**Architecture:** Keep scheduler, radix cache, and req-to-token in the CP-group logical address space. Keep device and host KV pools physical and per-rank sharded. Add a CP-aware HiCache controller adapter plus explicit per-node host metadata so only owned physical shards are transferred while radix state remains logically complete. + +**Tech Stack:** Python, SGLang server runtime, NSA `NSATokenToKVPool`, `CPSharedPagedTokenToKVPoolAllocator`, `CpSharedKVLayout`, `HiRadixCache`, `HiCacheController`. + +--- + +## Scope + +This design covers the first implementation stage only: + +- Support `--enable-nsa-prefill-cp-shared-kv` with HiCache host write/load/evict. +- Keep HiCache storage backends unsupported in this mode. +- Keep radix tree, scheduler admission, and req-to-token state in full logical KV loc space. +- Keep each CP rank's device KV pool and host KV pool sharded by physical owner pages. + +This design does not cover: + +- HiCache storage backend backup or prefetch. +- Storage key changes for File, NIXL, Mooncake, EIC, HF3FS, or other backends. +- Decode-side CP/shared KV. +- HiSparse NSA KV pools. +- A full replacement of HiRadixCache with upstream UnifiedRadixTree or HybridCacheController. + +If CP shared KV and `hicache_storage_backend` are enabled together, startup must fail fast with a clear error. Silent fallback is not acceptable because the existing storage path assumes page-hash sequences and contiguous host indices that do not describe CP-owned sparse shards. + +--- + +## Existing Code Facts + +The current CP shared KV branch already separates logical and physical KV locs: + +- `python/sglang/srt/mem_cache/cp_shared_kv_layout.py` provides owner masks and logical-to-physical conversions for torch and NumPy paths. +- `python/sglang/srt/mem_cache/allocator.py::CPSharedPagedTokenToKVPoolAllocator` tracks logical allocator size while storing physical pool size separately. +- `python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py` profiles physical token capacity and expands logical capacity by CP size. +- Runtime persistent writes already use `CpSharedKVLayout` to filter owned logical locs and convert to physical locs before touching local KV buffers. + +The current HiCache path is not CP-aware: + +- `HiRadixCache.write_backup()` passes `node.value` directly to `HiCacheController.write()`. +- `HiCacheController.write()` allocates host indices for the full tensor length and transfers those indices as if they were physical device indices. +- `HiCacheController.load()` allocates device indices for the full host tensor length and transfers those indices as if they were physical device indices. +- `HiRadixCache.load_back()` uses `len(node.host_value)` to compute host hit length and to slice restored device indices. +- `HiRadixCache._split_node()` slices `node.host_value` when radix nodes split. + +Under CP shared KV, `node.value` is a full logical loc tensor. It cannot be passed directly into physical host/device transfer functions. + +--- + +## Core Invariants + +The implementation must preserve these invariants: + +1. `TreeNode.value` remains a full logical loc tensor. +2. `req_to_token` remains full logical loc state. +3. Scheduler admission and radix matching remain based on full logical token lengths. +4. `HostKVCache` remains physical-only and does not import or depend on `CpSharedKVLayout`. +5. CP filtering and logical-to-physical conversion happen in `HiCacheController` or a small adapter owned by it. +6. `TreeNode.host_value` keeps the legacy non-CP meaning: physical host indices tensor. +7. CP mode does not store CP-owned host subsets in `TreeNode.host_value`. +8. CP mode uses explicit metadata to represent host backup state. +9. Host eviction and load-back accounting report logical token counts, even though local physical transfers are sparse. + +--- + +## Metadata Model + +Add explicit CP host metadata for HiRadixCache nodes. + +Recommended structure: + +```python +@dataclass +class CpHiCacheNodeMetadata: + logical_len: int + owned_positions: torch.Tensor + host_indices: torch.Tensor + + def split(self, split_len: int) -> tuple["CpHiCacheNodeMetadata", "CpHiCacheNodeMetadata"]: + if split_len < 0 or split_len > self.logical_len: + raise ValueError( + f"split_len must be in [0, {self.logical_len}], got {split_len}" + ) + parent_mask = self.owned_positions < split_len + child_mask = ~parent_mask + parent = CpHiCacheNodeMetadata( + logical_len=split_len, + owned_positions=self.owned_positions[parent_mask], + host_indices=self.host_indices[parent_mask], + ) + child = CpHiCacheNodeMetadata( + logical_len=self.logical_len - split_len, + owned_positions=self.owned_positions[child_mask] - split_len, + host_indices=self.host_indices[child_mask], + ) + return parent, child +``` + +Field semantics: + +- `logical_len` is the full logical token count represented by this node. +- `owned_positions` are node-local logical token positions owned by this CP rank. +- `host_indices` are local physical host slots storing the owned token KV data. +- `len(owned_positions) == len(host_indices)`. +- `owned_positions` can be empty on a rank that owns no page in this node. +- `owned_positions` are node-local positions, not old logical loc values. This avoids depending on evicted logical locs remaining reusable after load-back. +- `owned_positions` and `host_indices` are stored as CPU `torch.int64` tensors. Radix tree operations are CPU-side bookkeeping, and keeping metadata on CPU avoids retaining GPU tensors in long-lived tree nodes. + +Add fields to `TreeNode` or attach equivalent metadata from HiRadixCache: + +- `host_len: int = 0` +- `cp_hicache: Optional[CpHiCacheNodeMetadata] = None` + +For non-CP HiCache, existing `host_value` behavior remains unchanged. + +For CP HiCache, backup state is represented by `host_len > 0` and `cp_hicache is not None`. + +Radix backup checks must become cache-type aware. A helper is preferable to changing the global `TreeNode.backuped` property for all cache classes: + +```python +def _node_backuped(self, node: TreeNode) -> bool: + if self._uses_cp_hicache: + return node.host_len > 0 and node.cp_hicache is not None + return node.host_value is not None +``` + +Use this helper inside `HiRadixCache` paths that currently read `node.backuped`. Do not change SWA, Mamba, or plain Radix semantics as part of this integration unless a touched path requires it. + +--- + +## Controller Adapter + +`HiCacheController` is the transfer chokepoint and should own CP-aware translation. + +CP HiCache mode must be detected from the allocator or explicit layout, not from the KV pool alone. Use `isinstance(token_to_kv_pool_allocator, CPSharedPagedTokenToKVPoolAllocator)` or an explicitly passed `CpSharedKVLayout` to enable the CP adapter. After CP mode is detected, validate that the underlying KV pool is `NSATokenToKVPool` for this stage. + +Add initialization state: + +- `uses_cp_hicache: bool` +- `cp_shared_kv_layout: Optional[CpSharedKVLayout]` +- `page_size: int` + +The controller can derive CP mode from the allocator by checking for CP-specific fields such as `cp_size` and `cp_rank`, or it can receive an explicit layout from `HiRadixCache`. Prefer explicit construction in `HiRadixCache` and passing the layout into `HiCacheController` so the adapter has one source of truth. + +### CP Write + +External call shape stays close to the current API: + +```python +metadata = cache_controller.write(device_indices=node.value, node_id=node.id) +``` + +CP behavior: + +1. Treat `device_indices` as full logical locs. +2. Build `owned_mask = layout.owned_by_this_rank(device_indices)`. +3. Build `owned_positions = owned_mask.nonzero(as_tuple=True)[0]`. +4. Build `physical_device_indices = layout.logical_locs_to_physical(device_indices[owned_mask])`. +5. Allocate `host_indices = mem_pool_host.alloc(len(physical_device_indices))`. +6. If allocation fails, return a failure that exposes `required_host_slots=len(physical_device_indices)` so `HiRadixCache` can evict physical host slots and retry. +7. If the owned count is zero, skip host allocation and skip physical transfer, but still return metadata with `logical_len=len(device_indices)`. +8. Enqueue or run the existing host transfer using `host_indices` and `physical_device_indices`. +9. Store returned metadata tensors on CPU: `owned_positions.cpu()` and `host_indices.cpu()`. +10. Return `CpHiCacheNodeMetadata(logical_len, owned_positions, host_indices)`. + +Every logical write-back must produce the same acknowledgement cardinality on every TP rank. If a rank owns zero physical slots for the node, the controller must still create a completed no-op write acknowledgement when `HiRadixCache` adds the node to `ongoing_write_through`. This preserves the existing `writing_check()` invariant that all ranks pop the same logical ack ids after `all_reduce(MIN)`. + +The non-CP path continues returning the existing host indices tensor. + +### CP Load + +The load path must allocate full logical device locs but transfer only owned physical shards. + +Recommended API: + +```python +device_indices = cache_controller.load_cp(nodes_to_load) +``` + +CP behavior: + +1. Compute `logical_len = sum(node.host_len for node in nodes_to_load)`. +2. Allocate `device_indices = mem_pool_device_allocator.alloc(logical_len)`. +3. For each node, take its slice of `device_indices`. +4. Use `node.cp_hicache.owned_positions` to select the newly allocated logical locs owned by this rank. +5. Convert selected logical locs to physical locs with `layout.logical_locs_to_physical(selected_logical_locs)`. +6. Concatenate all owned host indices and owned physical device indices. +7. If the owned count is zero, skip physical transfer and still return the full logical `device_indices`. +8. Enqueue or run the existing host-to-device transfer for the owned physical pairs. +9. Return the full logical `device_indices` to HiRadixCache. + +CP load-back has the same acknowledgement cardinality requirement as CP write-back. A zero-owned rank must still produce a completed no-op load acknowledgement or equivalent completed producer event so lock release and producer ids remain consistent across ranks. + +The non-CP path continues using the existing `load(host_indices=host_indices)` API shape. + +--- + +## HiRadixCache Changes + +### Initialization + +`HiRadixCache.__init__()` must detect CP shared KV mode and construct the controller adapter. + +It should also enforce: + +- CP shared KV + HiCache storage backend is unsupported and must fail during server argument validation or cache initialization. +- CP shared KV + HiCache is valid only for `NSATokenToKVPool` in this stage. +- CP HiCache mode is enabled from `CPSharedPagedTokenToKVPoolAllocator` or an explicit `CpSharedKVLayout`, not merely because the device pool is `NSATokenToKVPool`. + +### Backup State + +Introduce helpers inside `HiRadixCache`: + +```python +def _node_backuped(self, node: TreeNode) -> bool: + if self._uses_cp_hicache: + return node.host_len > 0 and node.cp_hicache is not None + return node.host_value is not None + + +def _node_host_len(self, node: TreeNode) -> int: + if self._uses_cp_hicache: + return node.host_len + return len(node.host_value) + + +def _node_host_evict_indices(self, node: TreeNode) -> torch.Tensor: + if self._uses_cp_hicache: + return node.cp_hicache.host_indices + return node.host_value +``` + +Use these helpers in paths currently reading `node.backuped`, `len(node.host_value)`, or `node.host_value` for host eviction/load-back. + +### Write Backup + +Non-CP path remains unchanged. + +CP path: + +1. Call controller CP write with full logical `node.value`. +2. If write reports insufficient host slots, evict host entries until physical freed slots are at least the requested `required_host_slots`, then retry. +3. On success, set `node.host_len = len(node.value)` and `node.cp_hicache = metadata`. +4. Keep `node.host_value is None` in CP mode. +5. Add the node to `ongoing_write_through` if write-through acknowledgement is required. +6. Locking and metrics count logical tokens, not owned physical tokens. + +`_inc_hit_count()` must use `_node_backuped(node)` instead of `node.backuped`, otherwise CP-backed nodes would appear unbacked because `node.host_value` remains `None` in CP mode and the write-through path would repeatedly call `write_backup()`. + +### Split Node + +`_split_node()` must split CP metadata when a backed-up node is split: + +- Parent metadata keeps positions `< split_len`. +- Child metadata keeps positions `>= split_len`, rebased by subtracting `split_len`. +- Parent `host_len = split_len`. +- Child `host_len = old_host_len - split_len`. +- Host indices are split according to the owned positions mask, not by slicing first `split_len` entries. + +Non-CP `host_value` slicing remains unchanged. + +### Load Back + +Non-CP path remains unchanged. + +CP path: + +1. Collect `nodes_to_load` while walking evicted ancestors. +2. Compute `host_hit_len = sum(node.host_len for node in nodes_to_load)`. +3. Apply `load_back_threshold` and `mem_quota` checks to `host_hit_len`. +4. Call controller CP load with `nodes_to_load`. +5. If allocation fails, evict device logical tokens and retry. +6. Assign full logical slices from returned `device_indices` to each `node.value` using `node.host_len`. +7. Update `evictable_size_`, locks, and metrics by full logical lengths. + +The `load_back_threshold` check intentionally uses logical length in this first stage. Using local physical owned-token count would let different CP ranks make different load-back decisions. A physical-transfer threshold would require a CP-group reduction and is deferred. + +### Host Eviction + +Non-CP path remains unchanged. + +CP path: + +1. Host eviction candidate selection still uses radix logical nodes. +2. Free only `node.cp_hicache.host_indices` from the local host pool. +3. Track two counts separately: physical freed slots for allocator retry progress, and logical evicted tokens for metrics/debug output. +4. Clear `node.host_len` and `node.cp_hicache`. +5. Remove the node from the radix tree as the existing host eviction path does. + +Host eviction for CP write retry must progress by physical freed slots, not by logical `host_len`. A node can have a large logical `host_len` but zero local `host_indices` on a rank. Counting that logical length as eviction progress could stop the eviction loop without freeing the physical host slots needed by the retry. + +--- + +## Capacity Policy + +Host memory remains per-rank physical capacity based. + +The existing `HostKVCache` default sizing uses `device_pool.size * hicache_ratio` when `hicache_size <= 0`. Under CP shared KV, `device_pool.size` is already physical per-rank size, so this is the desired behavior. + +Do not change host pool sizing to logical capacity. That would multiply host memory by CP size per rank and defeat the CP memory-saving goal. + +Do not add a new CP-specific host ratio parameter in the first stage. Existing `hicache_ratio` and `hicache_size` remain the user-facing controls. + +--- + +## dev-hicache Fix Integration + +Do not blindly cherry-pick the `dev-hicache` commits onto `dev-cp-refactor`. The CP branch already modified allocator/common/model runner code, and direct cherry-picks risk overwriting CP shared KV logic. + +Manually port fixes in this order: + +1. Correctness fixes before CP HiCache integration: + - Load-back `mem_quota` scheduling budget from `a666ab552`. + - Paged allocator extend OOM accounting from `d6ca2fd0c`, merged carefully with compute-owner lane eviction in `common.py`. + - `defaultdict(TreeNode)` ghost-node fix from `e008d2cf3` for radix-like nodes. +2. CP host HiCache integration from this design. +3. HiCache performance fixes from `25b6a957c`: + - Incremental eviction heap. + - Cached timestamp in hot loops. + - Persistent all-reduce tensors. + - Async write acknowledgement thread. + +The short bigram page lookup fix from upstream commit `7be96c8ce` is a separate small correctness task and should be manually ported with a unit test. + +The upstream SWA HiCache storage type registration is out of scope for GLM 5 NSA host-only integration. + +--- + +## Error Handling + +Startup or initialization must fail with clear messages for unsupported combinations: + +- CP shared KV + HiCache storage backend. +- CP shared KV + HiSparse NSA pool. +- CP shared KV + non-NSA KV pool in HiCache host integration. + +Runtime allocation failure handling follows existing HiCache behavior: + +- Host write allocation failure returns `None`; HiRadixCache evicts host entries and retries. +- Device load allocation failure returns `None`; HiRadixCache evicts device entries and retries. +- If retry still fails, `load_back()` returns `None` and the request recomputes the prefix as today. + +For CP metadata consistency failures, fail loudly with assertions in unit-testable helper methods. Examples: + +- `metadata.logical_len` must match the node logical length. +- `len(metadata.owned_positions) == len(metadata.host_indices)`. +- split lengths must be within `[0, logical_len]`. +- `owned_positions` must be sorted and in range. +- `owned_positions` and `host_indices` must be CPU tensors before being stored on `TreeNode`. + +--- + +## Testing Strategy + +Use SGLang unit-test conventions for new tests: + +- Place tests under `test/registered/unit/`. +- Inherit from `CustomTestCase`. +- Register CPU tests with `register_cpu_ci` where possible. +- Avoid launching a server for metadata and controller mapping tests. + +Required unit coverage: + +1. `CpHiCacheNodeMetadata.split()`: + - page-aligned split. + - `split_len == 0`. + - split with no owned positions on this rank. + - split where owned positions are on both sides. + - invalid split values raise errors. +2. Controller CP write mapping: + - full logical locs produce only owned host allocation. + - physical device indices passed to host transfer match `CpSharedKVLayout`. + - zero-owned node returns metadata without allocating host slots. + - zero-owned node still produces one completed logical write acknowledgement if the node enters `ongoing_write_through`. +3. Controller CP load mapping: + - full logical device locs are allocated. + - owned positions map to newly allocated logical locs, then to physical locs. + - zero-owned metadata skips transfer but returns full logical locs. + - zero-owned metadata still produces a completed logical load acknowledgement or producer event. +4. HiRadixCache metadata behavior: + - `_node_backuped()` returns true for CP nodes with `host_len` and metadata. + - `_inc_hit_count()` uses `_node_backuped()` and does not repeatedly write back already CP-backed nodes. + - `_split_node()` splits `host_len` and CP metadata correctly. + - `load_back()` uses `host_len` rather than `len(host_value)` in CP mode. + - host eviction frees owned host indices and uses physical freed slots as retry progress. +5. Server args/config validation: + - CP shared KV + `hicache_storage_backend` fails fast. + +Required regression checks: + +- CP shared KV layout tests. +- Radix cache unit tests. +- Scheduler prefill-adder tests covering `mem_quota`. +- New CP HiCache metadata/controller tests. + +Optional manual GPU smoke: + +- Launch GLM5/NSA CP shared KV with HiCache host-only. +- Use repeated-prefix requests to trigger backup, device eviction, and host load-back. +- Confirm no logical/physical index assertion failures and no host/storage unsupported path is entered. + +--- + +## Implementation Order + +1. Add tests for existing correctness fixes that are currently missing from the CP branch. +2. Manually port load-back `mem_quota`, allocator OOM accounting, and `defaultdict(TreeNode)` fixes. +3. Add CP+storage fail-fast validation. +4. Add `CpHiCacheNodeMetadata` and unit tests. +5. Add CP adapter methods in `HiCacheController` with mock unit tests. +6. Wire `HiRadixCache` backup, split, load-back, and host eviction to CP metadata helpers. +7. Run CPU/unit regression tests. +8. Manually smoke test GPU CP+HiCache host-only if hardware/model access is available. +9. Port performance fixes from `25b6a957c` as separate patches with regression tests. + +--- + +## Open Follow-Up Work + +These items are deliberately deferred: + +- CP-aware HiCache storage key and prefetch design. +- Support for File/NIXL/Mooncake/EIC/HF3FS storage backends under CP HiCache. +- Independent CP-specific host ratio configuration. +- Decode-side CP/shared KV HiCache support. +- Upstream UnifiedRadixTree or HybridCacheController migration.