Keep current reuse independent of async CP prefetch
Async MLA/index prefetch is a scheduling optimization, not the correctness contract for target current reuse. Tiny cache-hit suffixes can skip async prefetcher creation while target partial-current reuse still composes page-slot prefix materialization with current KV rows synchronously. CP HiCache radix/device accounting now treats retained valid-tail pages as physical page spans so allocator state stays consistent when logical cache keys are shorter than the retained page. Constraint: CP shared KV ownership and HiCache residency are page-granular while request-visible cache lengths remain valid-token lengths. Constraint: Async prefetch can hang or regress on large-prefix tiny-extend traffic and must not be required for current reuse. Rejected: Treat missing prefetcher as fail-fast for target partial-current reuse | disabled useful current reuse and broke tiny-prefix/tiny-suffix traffic. Rejected: Keep async prefetcher object with synchronous consume mode | conflates prefetch object existence with current-layer correctness and hides fallback semantics. Confidence: medium Scope-risk: moderate Directive: Do not make current-only or target partial-current reuse depend on MLA/index prefetcher creation; prefetcher objects mean async next-layer work exists. Tested: Remote g0034 container py_compile for touched modules. Tested: Remote g0034 PYTHONPATH=python python -m pytest -q test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py -> 73 passed. Tested: Remote g0034 PYTHONPATH=python python -m pytest -q test/registered/unit/mem_cache/test_cp_hicache_metadata.py -> 90 passed. Not-tested: Latest full ETE traffic run with GLM-5.1 CP HiCache after this commit. Not-tested: CUDA kernel-level performance impact of synchronous no-prefetch partial-current compose.
This commit is contained in:
@@ -214,6 +214,7 @@ class Envs:
|
||||
SGLANG_CP_SHARED_KV_MATERIALIZE_NVTX = EnvBool(False)
|
||||
SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_PREFIX_TOKENS = EnvInt(512)
|
||||
SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_PREFIX_PAGES = EnvInt(-1)
|
||||
SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_EXTEND_TOKENS = EnvInt(-1)
|
||||
SGLANG_CP_DRAFT_SHARED_KV = EnvBool(False)
|
||||
SGLANG_CP_DRAFT_SHARED_KV_DEBUG = EnvBool(False)
|
||||
SGLANG_DISABLE_TAI_BIGRAM = EnvBool(False)
|
||||
|
||||
@@ -14,6 +14,7 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
|
||||
cp_shared_kv_mla_prefetch_enabled,
|
||||
cp_shared_kv_mla_prefetch_log,
|
||||
cp_shared_kv_mla_prefetch_log_enabled,
|
||||
cp_shared_kv_mla_prefetch_min_async_extend_tokens,
|
||||
cp_shared_kv_mla_prefetch_min_prefix_pages,
|
||||
cp_shared_kv_mla_prefetch_should_log_layer,
|
||||
filter_locs_mappable_to_physical_pool,
|
||||
@@ -430,6 +431,27 @@ class CpSharedKVMlaPrefetcher:
|
||||
)
|
||||
return None
|
||||
|
||||
extend_seq_lens_cpu = getattr(forward_batch, "extend_seq_lens_cpu", None)
|
||||
if extend_seq_lens_cpu is None or len(extend_seq_lens_cpu) != 1:
|
||||
_prefetch_log("create_skip reason=bad_extend_lens_metadata")
|
||||
return None
|
||||
extend_len = int(extend_seq_lens_cpu[0])
|
||||
min_async_extend_tokens = cp_shared_kv_mla_prefetch_min_async_extend_tokens(
|
||||
page_size=page_size
|
||||
)
|
||||
if extend_len < min_async_extend_tokens:
|
||||
_prefetch_log(
|
||||
"create_skip reason=extend_below_min cp_rank=%s cp_size=%s "
|
||||
"extend_len=%s min_async_extend_tokens=%s prefix_pages=%s page_size=%s",
|
||||
layout.cp_rank,
|
||||
layout.cp_size,
|
||||
extend_len,
|
||||
min_async_extend_tokens,
|
||||
prefix_pages,
|
||||
page_size,
|
||||
)
|
||||
return None
|
||||
|
||||
cp_group = get_attention_cp_group()
|
||||
if getattr(cp_group, "pynccl_comm", None) is None and layout.cp_size > 1:
|
||||
_prefetch_log(
|
||||
@@ -471,7 +493,8 @@ class CpSharedKVMlaPrefetcher:
|
||||
|
||||
_prefetch_log(
|
||||
"create cp_rank=%s cp_size=%s prefix_pages=%s total_slots=%s "
|
||||
"owned_prefix_pages=%s owned_total_pages=%s dense_pages=%s page_size=%s",
|
||||
"owned_prefix_pages=%s owned_total_pages=%s dense_pages=%s page_size=%s "
|
||||
"min_async_extend_tokens=%s extend_len=%s",
|
||||
layout.cp_rank,
|
||||
layout.cp_size,
|
||||
prefix_pages,
|
||||
@@ -480,6 +503,8 @@ class CpSharedKVMlaPrefetcher:
|
||||
owned_total_pages,
|
||||
remap.dense_num_pages,
|
||||
page_size,
|
||||
min_async_extend_tokens,
|
||||
extend_len,
|
||||
)
|
||||
create_total_ms = _cpu_timing_ms(create_cpu)
|
||||
_prefetch_log(
|
||||
@@ -1195,6 +1220,27 @@ class CpSharedKVIndexPrefetcher:
|
||||
)
|
||||
return None
|
||||
|
||||
extend_seq_lens_cpu = getattr(forward_batch, "extend_seq_lens_cpu", None)
|
||||
if extend_seq_lens_cpu is None or len(extend_seq_lens_cpu) != 1:
|
||||
_prefetch_log("index_create_skip reason=bad_extend_lens_metadata")
|
||||
return None
|
||||
extend_len = int(extend_seq_lens_cpu[0])
|
||||
min_extend_tokens = cp_shared_kv_mla_prefetch_min_async_extend_tokens(
|
||||
page_size=page_size
|
||||
)
|
||||
if extend_len < min_extend_tokens:
|
||||
_prefetch_log(
|
||||
"index_create_skip reason=extend_below_min cp_rank=%s cp_size=%s "
|
||||
"extend_len=%s min_extend_tokens=%s prefix_pages=%s page_size=%s",
|
||||
layout.cp_rank,
|
||||
layout.cp_size,
|
||||
extend_len,
|
||||
min_extend_tokens,
|
||||
prefix_pages,
|
||||
page_size,
|
||||
)
|
||||
return None
|
||||
|
||||
cp_group = get_attention_cp_group()
|
||||
if getattr(cp_group, "pynccl_comm", None) is None and layout.cp_size > 1:
|
||||
_index_prefetch_fallback_log(
|
||||
|
||||
@@ -96,6 +96,28 @@ def cp_shared_kv_mla_prefetch_min_prefix_pages(
|
||||
return max(int(configured), 0)
|
||||
|
||||
|
||||
def cp_shared_kv_mla_prefetch_min_async_extend_tokens(
|
||||
*, page_size: int | None = None
|
||||
) -> int:
|
||||
"""Minimum current-extend tokens required for async next-layer prefetch.
|
||||
|
||||
This threshold gates creation of the async next-layer MLA/index prefetcher
|
||||
only. It must not gate current-layer reuse: when no prefetcher exists, the
|
||||
backend can still synchronously materialize prefix pages and splice current
|
||||
KV rows for target partial-current reuse.
|
||||
|
||||
A negative env value uses the page size as the dynamic default. Set the env
|
||||
to 0 to allow async prefetch even for sub-page extends during experiments.
|
||||
"""
|
||||
|
||||
configured = envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_EXTEND_TOKENS.get()
|
||||
if configured < 0:
|
||||
if page_size is not None and int(page_size) > 0:
|
||||
return int(page_size)
|
||||
return 0
|
||||
return max(int(configured), 0)
|
||||
|
||||
|
||||
def cp_shared_kv_mla_prefetch_log(message: str, *args) -> None:
|
||||
if cp_shared_kv_mla_prefetch_log_enabled():
|
||||
logger.info("[CP_SHARED_KV_MLA_PREFETCH] " + message, *args)
|
||||
@@ -1945,6 +1967,82 @@ def materialize_local_token_kv_page_slots_into(
|
||||
dense_range.copy_(torch.where(owned_view, gathered, zero))
|
||||
|
||||
|
||||
def materialize_prefix_and_reuse_current_kv_page_slots(
|
||||
*,
|
||||
kv_cache: torch.Tensor,
|
||||
logical_locs: torch.Tensor,
|
||||
current_kv_cache: torch.Tensor,
|
||||
current_locs: torch.Tensor,
|
||||
slot_remap: SharedTokenKVSlotRemap,
|
||||
layout: CpSharedKVLayout,
|
||||
page_size: int,
|
||||
prefix_pages: int,
|
||||
layer_id: int | None = None,
|
||||
nvtx_source: str = "mla.partial_current_sync",
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Synchronously compose prefix materialization with current KV rows.
|
||||
|
||||
This is the non-prefetch partial-current path. It preserves the same
|
||||
page-slot layout used by async prefetch compose, but does not require a
|
||||
prefetcher object or an extra CUDA stream. Prefix pages are materialized and
|
||||
reduced immediately; current rows are then inserted into their padded suffix
|
||||
page slots and non-current tail slack is masked from the returned locs.
|
||||
"""
|
||||
|
||||
total_slots = int(slot_remap.slot_logical_pages.numel())
|
||||
if prefix_pages < 0 or prefix_pages > total_slots:
|
||||
raise ValueError(
|
||||
"Invalid CP shared KV partial-current prefix range: "
|
||||
f"prefix_pages={prefix_pages} total_slots={total_slots}"
|
||||
)
|
||||
|
||||
dense_kv_cache = kv_cache.new_zeros(
|
||||
(slot_remap.dense_num_pages * page_size, *kv_cache.shape[1:])
|
||||
)
|
||||
materialize_local_token_kv_page_slots_into(
|
||||
kv_cache=kv_cache,
|
||||
dense_kv_cache=dense_kv_cache,
|
||||
slot_logical_pages=slot_remap.slot_logical_pages,
|
||||
layout=layout,
|
||||
page_size=page_size,
|
||||
start_slot=0,
|
||||
end_slot=prefix_pages,
|
||||
)
|
||||
|
||||
prefix_rows = slot_range_to_token_slice(page_size, 0, prefix_pages)
|
||||
_all_reduce_materialized_buffer_range(
|
||||
dense_kv_cache,
|
||||
layout.cp_size,
|
||||
prefix_rows.start,
|
||||
prefix_rows.stop,
|
||||
nvtx_source=nvtx_source,
|
||||
nvtx_layer_id=layer_id,
|
||||
nvtx_cp_rank=layout.cp_rank,
|
||||
)
|
||||
|
||||
logical_locs = filter_locs_mappable_to_physical_pool(
|
||||
logical_locs=logical_locs,
|
||||
layout=layout,
|
||||
physical_token_capacity=kv_cache.shape[0],
|
||||
)
|
||||
dense_locs = remap_logical_locs_to_slot_dense_locs_optimized(
|
||||
logical_locs,
|
||||
page_inverse=slot_remap.page_inverse,
|
||||
page_size=page_size,
|
||||
)
|
||||
mixed_kv_cache, mixed_locs, _ = fill_current_kv_page_slots_and_remap_locs(
|
||||
dense_kv_cache=dense_kv_cache,
|
||||
materialized_dense_locs=dense_locs,
|
||||
current_kv_cache=current_kv_cache,
|
||||
logical_locs=logical_locs,
|
||||
current_locs=current_locs,
|
||||
page_inverse=slot_remap.page_inverse,
|
||||
page_size=page_size,
|
||||
mask_non_current_in_current_pages=True,
|
||||
)
|
||||
return mixed_kv_cache, mixed_locs
|
||||
|
||||
|
||||
def slot_range_to_token_slice(
|
||||
page_size: int,
|
||||
start_slot: int,
|
||||
|
||||
@@ -26,6 +26,7 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
|
||||
filter_owned_logical_locs,
|
||||
get_or_build_shared_token_kv_slot_remap,
|
||||
is_current_only_extend_batch,
|
||||
materialize_prefix_and_reuse_current_kv_page_slots,
|
||||
materialize_shared_token_kv_buffer,
|
||||
should_reuse_current_extend_kv,
|
||||
tensor_debug_checksum,
|
||||
@@ -1850,23 +1851,69 @@ class NativeSparseAttnBackend(
|
||||
if extend_lens_cpu is not None
|
||||
else None
|
||||
)
|
||||
raise RuntimeError(
|
||||
"[CP_SHARED_KV_FAIL_FAST][mla_partial_current_prefetch] "
|
||||
"CP shared KV MLA partial-current reuse requires "
|
||||
"page-slot prefetch compose. Compact "
|
||||
"materialize/current merge fallback is disabled "
|
||||
"because it can expose padded tail slack. "
|
||||
f"reason={reason} "
|
||||
f"cp_rank={forward_batch.cp_shared_kv_layout.cp_rank} "
|
||||
f"layer_id={layer.layer_id} "
|
||||
f"prefix_lens={prefix_lens} "
|
||||
f"extend_lens={extend_lens} "
|
||||
f"current_rows={int(current_kv_cache.shape[0])} "
|
||||
f"logical_page_table_shape={tuple(logical_page_table_1.shape)} "
|
||||
f"current_locs_shape={tuple(forward_batch.out_cache_loc.shape)} "
|
||||
f"page_size={current_remap_page_size} "
|
||||
f"logical_page_capacity={current_remap_logical_page_capacity}"
|
||||
page_size = int(forward_batch.token_to_kv_pool.page_size)
|
||||
if (
|
||||
prefix_lens_cpu is None
|
||||
or len(prefix_lens_cpu) != 1
|
||||
or int(prefix_lens_cpu[0]) <= 0
|
||||
or int(prefix_lens_cpu[0]) % page_size != 0
|
||||
):
|
||||
raise RuntimeError(
|
||||
"[CP_SHARED_KV_FAIL_FAST][mla_partial_current_sync] "
|
||||
"CP shared KV MLA partial-current sync compose "
|
||||
"requires one positive page-aligned prefix. "
|
||||
f"reason={reason} "
|
||||
f"cp_rank={forward_batch.cp_shared_kv_layout.cp_rank} "
|
||||
f"layer_id={layer.layer_id} "
|
||||
f"prefix_lens={prefix_lens} "
|
||||
f"extend_lens={extend_lens} "
|
||||
f"current_rows={int(current_kv_cache.shape[0])} "
|
||||
f"logical_page_table_shape={tuple(logical_page_table_1.shape)} "
|
||||
f"current_locs_shape={tuple(forward_batch.out_cache_loc.shape)} "
|
||||
f"page_size={page_size}"
|
||||
)
|
||||
prefix_pages = int(prefix_lens_cpu[0]) // page_size
|
||||
slot_remap = get_or_build_shared_token_kv_slot_remap(
|
||||
forward_batch,
|
||||
kv_cache=kv_cache,
|
||||
remap_logical_pages=metadata.real_page_table,
|
||||
layout=forward_batch.cp_shared_kv_layout,
|
||||
page_size=page_size,
|
||||
)
|
||||
kv_cache, page_table_1 = (
|
||||
materialize_prefix_and_reuse_current_kv_page_slots(
|
||||
kv_cache=kv_cache,
|
||||
logical_locs=logical_page_table_1,
|
||||
current_kv_cache=current_kv_cache,
|
||||
current_locs=forward_batch.out_cache_loc,
|
||||
slot_remap=slot_remap,
|
||||
layout=forward_batch.cp_shared_kv_layout,
|
||||
page_size=page_size,
|
||||
prefix_pages=prefix_pages,
|
||||
layer_id=layer.layer_id,
|
||||
)
|
||||
)
|
||||
if (
|
||||
cp_shared_kv_mla_prefetch_log_enabled()
|
||||
and cp_shared_kv_mla_prefetch_should_log_layer(
|
||||
layer.layer_id
|
||||
)
|
||||
):
|
||||
cp_shared_kv_mla_prefetch_log(
|
||||
"forward_partial_current_sync_compose cp_rank=%s "
|
||||
"layer=%s reason=%s prefix_lens=%s extend_lens=%s "
|
||||
"prefix_pages=%s current_rows=%s kv_rows=%s "
|
||||
"page_table_shape=%s",
|
||||
forward_batch.cp_shared_kv_layout.cp_rank,
|
||||
layer.layer_id,
|
||||
reason,
|
||||
prefix_lens,
|
||||
extend_lens,
|
||||
prefix_pages,
|
||||
int(current_kv_cache.shape[0]),
|
||||
int(kv_cache.shape[0]),
|
||||
tuple(page_table_1.shape),
|
||||
)
|
||||
if (
|
||||
cp_shared_kv_mla_prefetch_log_enabled()
|
||||
and cp_shared_kv_mla_prefetch_should_log_layer(layer.layer_id)
|
||||
|
||||
@@ -46,6 +46,7 @@ from sglang.srt.mem_cache.radix_cache import (
|
||||
RadixCache,
|
||||
RadixKey,
|
||||
TreeNode,
|
||||
ceil_to_page_len,
|
||||
compute_node_hash_values,
|
||||
split_node_hash_value,
|
||||
)
|
||||
@@ -2035,7 +2036,6 @@ class HiRadixCache(RadixCache):
|
||||
or self.page_size <= 1
|
||||
or prefix_len <= 0
|
||||
or prefix_len % self.page_size == 0
|
||||
or not self._node_backuped(child)
|
||||
):
|
||||
return prefix_len
|
||||
return prefix_len // self.page_size * self.page_size
|
||||
@@ -2380,6 +2380,27 @@ class HiRadixCache(RadixCache):
|
||||
node.pin_expiry = 0.0
|
||||
node.pin_ttl = 0
|
||||
|
||||
def _node_device_resident_len(self, node: TreeNode) -> int:
|
||||
"""Allocator-visible device residency for a radix node.
|
||||
|
||||
CP HiCache radix keys are valid-token lengths, but CP device KV
|
||||
ownership is page-granular. Residency accounting therefore uses the
|
||||
physical padded page span whenever CP HiCache is active, while normal
|
||||
HiCache keeps historical token-count accounting.
|
||||
"""
|
||||
|
||||
value = getattr(node, "value", None)
|
||||
if value is None:
|
||||
return 0
|
||||
if not getattr(self, "_uses_cp_hicache", False):
|
||||
return len(value)
|
||||
|
||||
metadata = getattr(node, "cp_hicache", None)
|
||||
padded_len = getattr(metadata, "padded_len", None)
|
||||
if padded_len is not None:
|
||||
return int(padded_len)
|
||||
return ceil_to_page_len(len(value), getattr(self, "page_size", 1))
|
||||
|
||||
def pin_prefix(
|
||||
self, token_ids: List[int], ttl_seconds: int = 300
|
||||
) -> Tuple[int, Optional[str]]:
|
||||
@@ -2455,10 +2476,11 @@ class HiRadixCache(RadixCache):
|
||||
|
||||
delta = 0
|
||||
while node != self.root_node:
|
||||
resident_len = self._node_device_resident_len(node)
|
||||
if node.lock_ref == 0:
|
||||
self.evictable_size_ -= len(node.key)
|
||||
self.protected_size_ += len(node.key)
|
||||
delta -= len(node.key)
|
||||
self.evictable_size_ -= resident_len
|
||||
self.protected_size_ += resident_len
|
||||
delta -= resident_len
|
||||
node.lock_ref += 1
|
||||
self._update_leaf_status(node)
|
||||
self._update_host_leaf_status(node)
|
||||
@@ -2473,10 +2495,11 @@ class HiRadixCache(RadixCache):
|
||||
|
||||
delta = 0
|
||||
while node != self.root_node:
|
||||
resident_len = self._node_device_resident_len(node)
|
||||
if node.lock_ref == 1:
|
||||
self.evictable_size_ += len(node.key)
|
||||
self.protected_size_ -= len(node.key)
|
||||
delta += len(node.key)
|
||||
self.evictable_size_ += resident_len
|
||||
self.protected_size_ -= resident_len
|
||||
delta += resident_len
|
||||
node.lock_ref -= 1
|
||||
self._update_leaf_status(node)
|
||||
self._update_host_leaf_status(node)
|
||||
@@ -2487,6 +2510,34 @@ class HiRadixCache(RadixCache):
|
||||
node = node.parent
|
||||
return DecLockRefResult(delta=delta)
|
||||
|
||||
def inc_node_lock_ref(self, node: TreeNode):
|
||||
if self.disable:
|
||||
return
|
||||
if node == self.root_node:
|
||||
return
|
||||
resident_len = self._node_device_resident_len(node)
|
||||
if node.lock_ref == 0:
|
||||
self.evictable_size_ -= resident_len
|
||||
self.protected_size_ += resident_len
|
||||
node.lock_ref += 1
|
||||
self._update_leaf_status(node)
|
||||
if hasattr(self, "evictable_host_leaves"):
|
||||
self._update_host_leaf_status(node)
|
||||
|
||||
def dec_node_lock_ref(self, node: TreeNode):
|
||||
if self.disable:
|
||||
return
|
||||
if node == self.root_node:
|
||||
return
|
||||
resident_len = self._node_device_resident_len(node)
|
||||
if node.lock_ref == 1:
|
||||
self.evictable_size_ += resident_len
|
||||
self.protected_size_ -= resident_len
|
||||
node.lock_ref -= 1
|
||||
self._update_leaf_status(node)
|
||||
if hasattr(self, "evictable_host_leaves"):
|
||||
self._update_host_leaf_status(node)
|
||||
|
||||
def _update_host_leaf_status(self, node: TreeNode):
|
||||
if not node.evicted or node.lock_ref > 0:
|
||||
if node in self.evictable_host_leaves:
|
||||
@@ -2587,13 +2638,15 @@ class HiRadixCache(RadixCache):
|
||||
|
||||
def _evict_backuped(self, node: TreeNode):
|
||||
# GPU -> CPU demotion: no BlockRemoved since block is still reachable via load_back
|
||||
num_evicted = self.cache_controller.evict_device(node.value)
|
||||
assert num_evicted > 0
|
||||
self.evictable_size_ -= num_evicted
|
||||
device_resident_len = self._node_device_resident_len(node)
|
||||
freed_len = self.cache_controller.evict_device(node.value)
|
||||
assert freed_len > 0
|
||||
self.evictable_size_ -= device_resident_len
|
||||
logger.info(
|
||||
"[HiCache-evict] _evict_backuped: node_id=%d num_evicted=%d lock_ref=%d backed=%s",
|
||||
"[HiCache-evict] _evict_backuped: node_id=%d num_evicted=%d physical_tokens=%d lock_ref=%d backed=%s",
|
||||
node.id,
|
||||
num_evicted,
|
||||
freed_len,
|
||||
device_resident_len,
|
||||
node.lock_ref,
|
||||
self._node_backuped(node),
|
||||
)
|
||||
@@ -2603,11 +2656,11 @@ class HiRadixCache(RadixCache):
|
||||
# update leaf status for the parent because the node is evicted
|
||||
self._update_leaf_status(node.parent)
|
||||
self._update_host_leaf_status(node.parent)
|
||||
return num_evicted
|
||||
return device_resident_len
|
||||
|
||||
def _evict_regular(self, node: TreeNode):
|
||||
# evict a node not initiated write to host -- emit BlockRemoved
|
||||
num_evicted = len(node.value)
|
||||
num_evicted = self._node_device_resident_len(node)
|
||||
logger.info(
|
||||
"[HiCache-evict] _evict_regular: node_id=%d num_evicted=%d",
|
||||
node.id,
|
||||
@@ -2618,6 +2671,18 @@ class HiRadixCache(RadixCache):
|
||||
self._delete_leaf(node)
|
||||
return num_evicted
|
||||
|
||||
def _delete_leaf(self, node: TreeNode):
|
||||
key = self.get_child_key_fn(node.key)
|
||||
v = node.parent.children.pop(key, None)
|
||||
assert v == node, f"parent does not have child key, {key}"
|
||||
|
||||
self.evictable_size_ -= self._node_device_resident_len(node)
|
||||
if node in self.evictable_leaves:
|
||||
self.evictable_leaves.remove(node)
|
||||
self._update_leaf_status(node.parent)
|
||||
if hasattr(self, "evictable_host_leaves"):
|
||||
self._update_host_leaf_status(node.parent)
|
||||
|
||||
def _remove_host_leaf(self, node: TreeNode) -> TreeNode:
|
||||
parent = node.parent
|
||||
key = self.get_child_key_fn(node.key)
|
||||
@@ -3372,7 +3437,6 @@ class HiRadixCache(RadixCache):
|
||||
if (
|
||||
self._uses_cp_hicache
|
||||
and self.page_size > 1
|
||||
and self._node_backuped(child)
|
||||
and prefix_len % self.page_size != 0
|
||||
):
|
||||
prefix_len = self._cp_floor_backed_partial_split_len(
|
||||
@@ -3474,7 +3538,7 @@ class HiRadixCache(RadixCache):
|
||||
# change the reference if the node is evicted
|
||||
# this often happens in the case of KV cache recomputation
|
||||
node.value = value[:prefix_len].clone()
|
||||
self.evictable_size_ += len(node.value)
|
||||
self.evictable_size_ += self._node_device_resident_len(node)
|
||||
self._update_leaf_status(node)
|
||||
self._update_host_leaf_status(node)
|
||||
# update parent status as a new leaf is added into device
|
||||
@@ -3495,7 +3559,7 @@ class HiRadixCache(RadixCache):
|
||||
new_node.priority = max(new_node.priority, priority)
|
||||
if new_node.evicted:
|
||||
new_node.value = value[:prefix_len].clone()
|
||||
self.evictable_size_ += len(new_node.value)
|
||||
self.evictable_size_ += self._node_device_resident_len(new_node)
|
||||
self._update_leaf_status(new_node)
|
||||
self._update_host_leaf_status(new_node)
|
||||
# update parent status as a new leaf is added into device
|
||||
@@ -3529,7 +3593,7 @@ class HiRadixCache(RadixCache):
|
||||
new_node.key = key
|
||||
new_node.value = value.clone()
|
||||
node.children[child_key] = new_node
|
||||
self.evictable_size_ += len(value)
|
||||
self.evictable_size_ += self._node_device_resident_len(new_node)
|
||||
self._update_leaf_status(node)
|
||||
self._update_leaf_status(new_node)
|
||||
|
||||
|
||||
@@ -117,6 +117,12 @@ def page_align_keys(key: list, page_size) -> list:
|
||||
return key[:page_aligned_len]
|
||||
|
||||
|
||||
def ceil_to_page_len(length: int, page_size: int) -> int:
|
||||
if page_size <= 1:
|
||||
return length
|
||||
return ((length + page_size - 1) // page_size) * page_size
|
||||
|
||||
|
||||
class TreeNode:
|
||||
|
||||
counter = 0
|
||||
@@ -526,8 +532,19 @@ class RadixCache(BasePrefixCache):
|
||||
if prepared_cp_backup is not None:
|
||||
req.cp_hicache_prepared_backup = None
|
||||
|
||||
# free the unaligned tail
|
||||
self.token_to_kv_pool_allocator.free(kv_indices[len(keys) :])
|
||||
# Free the unaligned tail.
|
||||
#
|
||||
# CP HiCache radix keys are scheduler-visible valid lengths, while the
|
||||
# backing allocator is page-granular. Freeing a loc inside the retained
|
||||
# tail page releases the whole page and leaves radix accounting pointing
|
||||
# at freed memory (observed as available_size=max with evictable_size>0
|
||||
# after EAGLE warmup). Therefore CP starts tail free at the next page
|
||||
# boundary; non-CP keeps the historical token boundary because keys were
|
||||
# already floored by page_align_keys above.
|
||||
tail_free_start = len(keys)
|
||||
if getattr(self, "_uses_cp_hicache", False):
|
||||
tail_free_start = ceil_to_page_len(tail_free_start, self.page_size)
|
||||
self.token_to_kv_pool_allocator.free(kv_indices[tail_free_start:])
|
||||
|
||||
# Remove req slot release the cache lock
|
||||
self.dec_lock_ref(req.last_node)
|
||||
|
||||
Reference in New Issue
Block a user