Stabilize CP HiCache page-tail ownership under EAGLE reuse
CP shared KV and HiCache now keep page-aligned physical ownership while preserving valid-token radix semantics. Repeated tiny EAGLE exact hits free duplicate tail pages instead of leaking one allocator page, owner-lane load-back uses page-vector admission/eviction, and single-DP idle schedulers avoid entering an unnecessary MLP-sync collective. The commit also records the current page-aligned cache contract and adds gated decode-side EAGLE accept diagnostics so future accept-length collapses can be tied to draft KV/state transfer evidence instead of more prefill cache speculation. Constraint: CP HiCache allocator ownership is page-granular while radix matching remains valid-token based. Constraint: New diagnostics must be gated and must not alter normal EAGLE, transfer, or cache behavior. Rejected: Padding short requests to cp_size or 2*cp_size pages | wastes KV capacity and still hides valid-tail lifecycle bugs. Rejected: Adding more unconditional collectives to prove CP consistency | hot-path collectives previously caused severe performance risk. Confidence: medium Scope-risk: broad Directive: Do not reintroduce silent fallback for CP shared KV/HiCache paths; warning-level fallback or fail-fast is intentional. Tested: git diff --check Tested: local py_compile for all modified Python files Tested: remote g0034 container py_compile for modified Python/test files Tested: remote g0034 container PYTHONPATH=python python -m pytest -q test/registered/unit/layers/test_nsa_cp_utils.py test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py test/registered/unit/mem_cache/test_cp_hicache_load_back_owner_lanes.py test/registered/unit/managers/test_scheduler_dp_attn_mixin.py => 114 passed, 5 warnings, 2 subtests passed Not-tested: full ETE traffic rerun after this commit Not-tested: CUDA/TAI kernel benchmark coverage for all production shapes
This commit is contained in:
@@ -89,11 +89,17 @@ for _schema in (
|
||||
raise
|
||||
|
||||
from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
|
||||
import sglang.srt.mem_cache.common as mem_cache_common
|
||||
from sglang.srt.mem_cache.hiradix_cache import (
|
||||
CpHiCacheNodeMetadata,
|
||||
HiRadixCache,
|
||||
)
|
||||
from sglang.srt.mem_cache.radix_cache import RadixKey, TreeNode, get_child_key
|
||||
from sglang.srt.mem_cache.radix_cache import (
|
||||
RadixKey,
|
||||
TreeNode,
|
||||
_key_match_paged,
|
||||
get_child_key,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
@@ -216,6 +222,57 @@ def _make_cache(allocator):
|
||||
return cache
|
||||
|
||||
|
||||
class _TinyReqToTokenPool:
|
||||
def __init__(self, size=8, max_context_len=32):
|
||||
self.req_to_token = torch.zeros((size, max_context_len), dtype=torch.int64)
|
||||
|
||||
def write(self, indices, values):
|
||||
self.req_to_token[indices] = values
|
||||
|
||||
def free(self, req):
|
||||
req.req_pool_idx = None
|
||||
|
||||
|
||||
def _make_tiny_eagle_req(cache, allocator, *, seq_len, req_pool_idx):
|
||||
page_size = allocator.page_size
|
||||
alloc_len = ((seq_len + page_size - 1) // page_size) * page_size
|
||||
locs = allocator.alloc(alloc_len)
|
||||
cache.req_to_token_pool.write(
|
||||
(req_pool_idx, slice(0, seq_len)), locs[:seq_len]
|
||||
)
|
||||
|
||||
req = types.SimpleNamespace(
|
||||
req_pool_idx=req_pool_idx,
|
||||
mamba_pool_idx=None,
|
||||
fill_ids=list(range(seq_len)),
|
||||
origin_input_ids=list(range(seq_len)),
|
||||
output_ids=[],
|
||||
extra_key=None,
|
||||
cache_protected_len=0,
|
||||
last_node=cache.root_node,
|
||||
priority=0,
|
||||
kv_committed_len=seq_len,
|
||||
kv_allocated_len=seq_len,
|
||||
kv_committed_freed=False,
|
||||
kv_overallocated_freed=False,
|
||||
cp_hicache_prepared_backup=None,
|
||||
)
|
||||
|
||||
def pop_committed_kv_cache():
|
||||
assert not req.kv_committed_freed
|
||||
req.kv_committed_freed = True
|
||||
return req.kv_committed_len
|
||||
|
||||
def pop_overallocated_kv_cache():
|
||||
assert not req.kv_overallocated_freed
|
||||
req.kv_overallocated_freed = True
|
||||
return req.kv_committed_len, req.kv_allocated_len
|
||||
|
||||
req.pop_committed_kv_cache = pop_committed_kv_cache
|
||||
req.pop_overallocated_kv_cache = pop_overallocated_kv_cache
|
||||
return req
|
||||
|
||||
|
||||
class TestCpHiCacheLoadBackOwnerLanes(CustomTestCase):
|
||||
def test_load_back_plan_reports_owner_lane_vectors(self):
|
||||
allocator = _make_allocator()
|
||||
@@ -302,6 +359,46 @@ class TestCpHiCacheLoadBackOwnerLanes(CustomTestCase):
|
||||
self.assertEqual(target.value.tolist(), loaded.tolist())
|
||||
self.assertIn(target.id, cache.ongoing_load_back)
|
||||
|
||||
def test_owner_lane_evict_params_choose_deficit_contributing_victim(self):
|
||||
from sglang.srt.mem_cache.base_prefix_cache import EvictParams
|
||||
|
||||
allocator = _make_allocator()
|
||||
cache = _make_cache(allocator)
|
||||
|
||||
non_contributing = _make_node(
|
||||
50,
|
||||
900,
|
||||
[1],
|
||||
value=torch.tensor([8, 9, 10, 11], dtype=torch.int64),
|
||||
priority=0,
|
||||
)
|
||||
contributing = _make_node(
|
||||
51,
|
||||
1000,
|
||||
[0],
|
||||
value=torch.tensor([4, 5, 6, 7], dtype=torch.int64),
|
||||
priority=10,
|
||||
)
|
||||
_attach_child(cache, cache.root_node, non_contributing)
|
||||
_attach_child(cache, cache.root_node, contributing)
|
||||
cache.evictable_leaves.update({non_contributing, contributing})
|
||||
cache.evictable_size_ = len(non_contributing.value) + len(contributing.value)
|
||||
|
||||
result = cache.evict(
|
||||
EvictParams(
|
||||
num_tokens=allocator.page_size,
|
||||
owner_lane_deficits=[1, 0, 0, 0],
|
||||
)
|
||||
)
|
||||
|
||||
self.assertEqual(result.num_tokens_evicted, allocator.page_size)
|
||||
self.assertEqual(
|
||||
[indices.tolist() for indices in cache.cache_controller.evicted_device_indices],
|
||||
[[4, 5, 6, 7]],
|
||||
)
|
||||
self.assertIsNone(contributing.value)
|
||||
self.assertIsNotNone(non_contributing.value)
|
||||
|
||||
def test_load_back_evicts_blocking_leaf_to_unlock_parent_owner_lane(self):
|
||||
allocator = _make_allocator()
|
||||
# Only owner lane 1 is initially available. The target needs owner lane
|
||||
@@ -384,6 +481,41 @@ class TestCpHiCacheLoadBackOwnerLanes(CustomTestCase):
|
||||
self.assertNotIn(target.id, cache.ongoing_load_back)
|
||||
self.assertEqual(cache.cache_controller.ack_load_queue, [])
|
||||
|
||||
def test_repeated_tiny_eagle_exact_hit_frees_duplicate_tail_page(self):
|
||||
allocator = _make_allocator(page_size=4, cp_size=4)
|
||||
cache = _make_cache(allocator)
|
||||
cache.req_to_token_pool = _TinyReqToTokenPool()
|
||||
cache.key_match_fn = functools.partial(
|
||||
_key_match_paged, page_size=allocator.page_size
|
||||
)
|
||||
cache.device = "cpu"
|
||||
cache.is_eagle = True
|
||||
cache.disable_finished_insert = False
|
||||
cache.enable_storage = False
|
||||
cache.write_through_threshold = 10**9
|
||||
|
||||
old_get_global_server_args = mem_cache_common.get_global_server_args
|
||||
mem_cache_common.get_global_server_args = lambda: types.SimpleNamespace(
|
||||
page_size=allocator.page_size,
|
||||
speculative_algorithm="EAGLE",
|
||||
)
|
||||
try:
|
||||
for req_pool_idx in (1, 2, 3):
|
||||
req = _make_tiny_eagle_req(
|
||||
cache, allocator, seq_len=2, req_pool_idx=req_pool_idx
|
||||
)
|
||||
cache.cache_unfinished_req(req)
|
||||
mem_cache_common.release_kv_cache(req, cache)
|
||||
|
||||
accounted = (
|
||||
allocator.available_size()
|
||||
+ cache.evictable_size()
|
||||
+ cache.protected_size()
|
||||
)
|
||||
self.assertEqual(accounted, allocator.size)
|
||||
finally:
|
||||
mem_cache_common.get_global_server_args = old_get_global_server_args
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -561,6 +561,14 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
forward_batch.out_cache_loc = torch.arange(127, dtype=torch.int64)
|
||||
self.assertFalse(can_reuse_current_extend_kv(forward_batch))
|
||||
|
||||
forward_batch.extend_seq_lens_cpu = [65]
|
||||
forward_batch.seq_lens_cpu = torch.tensor([40320 + 65], dtype=torch.int32)
|
||||
forward_batch.out_cache_loc = torch.arange(128, dtype=torch.int64)
|
||||
self.assertTrue(can_reuse_current_extend_kv(forward_batch))
|
||||
|
||||
forward_batch.out_cache_loc = torch.arange(64, dtype=torch.int64)
|
||||
self.assertFalse(can_reuse_current_extend_kv(forward_batch))
|
||||
|
||||
def test_should_reuse_current_extend_kv_disables_draft_cache_hit_suffix(self):
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
|
||||
Reference in New Issue
Block a user