Files
sglang/docs/superpowers/plans/2026-05-07-cp-hicache-host-implementation.md

65 KiB

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:

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:

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:

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:

from sglang.test.test_utils import CustomTestCase

Change the touched allocator test class to use the project test base:

class TestCPSharedPagedAllocator(CustomTestCase):

Add this test method to TestCPSharedPagedAllocator:

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:

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):

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:

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:

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):

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:

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:

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:

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:

from sglang.srt.mem_cache.base_prefix_cache import (
    BasePrefixCache,
    EvictParams,
    EvictResult,
)

Add this helper near evict_from_tree_cache:

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:

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:

evict_result = evict_from_tree_cache(tree_cache, num_tokens)

Append this line to each OOM message after available_and_evictable_str(tree_cache):

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:

from sglang.test.test_utils import CustomTestCase

Add this new test class near the existing TestTreeNode class before changing production code:

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:

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
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:

class TestHiCacheArgs(unittest.TestCase):

to:

class TestHiCacheArgs(CustomTestCase):

The file already imports CustomTestCase near the top.

  • Step 2: Add failing fail-fast test

Add this method to TestHiCacheArgs:

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:

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
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:

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:

from dataclasses import dataclass

Add this dataclass immediately before class HiRadixCache in python/sglang/srt/mem_cache/hiradix_cache.py:

@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:

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
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:

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:

from dataclasses import dataclass

Add this import beside the existing MLATokenToKVPool import at the top of python/sglang/srt/managers/cache_controller.py:

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:

from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout

Add these dataclasses before class HiCacheController:

@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:

cp_shared_kv_layout: Optional[CpSharedKVLayout] = None,

Set these fields after self.mem_pool_device = mem_pool_device:

self.cp_shared_kv_layout = cp_shared_kv_layout
self.uses_cp_hicache = cp_shared_kv_layout is not None

Then add this validation:

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:

return self._write_cp(device_indices=device_indices, priority=priority, node_id=node_id)

Implement _append_completed_write_ack and _write_cp in HiCacheController:

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
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:

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:

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:

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:

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:

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:

cp_shared_kv_layout=cp_shared_kv_layout,
  • Step 4: Add backup helpers

Add methods with these exact semantics:

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:

if not node.backuped:

with:

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:

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:

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
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:

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:

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:

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:

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:

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
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:

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:

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:

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
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:

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:

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
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:

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:

## 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:

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.