Make CP HiCache residency owner-lane deterministic

CP shared KV cannot treat capacity as a scalar token count: cache-hit load-back and fresh extend allocation both have to preserve the logical page owner pattern or later direct writes, HiCache reload, and prefix materialization can read the wrong lane. This change moves the critical paths to owner-lane plans, makes owner-lane exhaustion recoverable during prefill scheduling, and routes shared-KV prefix prefetch through prefetch-stream-safe KV getters so HiCache layer-load waits do not attach to the forward stream.

Constraint: CP shared KV correctness depends on page owner lane preservation across allocation, backup, load, eviction, and prefix materialization.

Constraint: Avoid adding CP/global collectives for capacity agreement; derive capacity from deterministic local owner-lane state.

Rejected: Keep SGLANG_DISABLE_TAI_OWNER_SELECT fallback | legacy allocation can silently break owner-lane invariants.

Rejected: Scalar total-token eviction for CP HiCache load-back | total capacity can be sufficient while the required owner lane is exhausted.

Confidence: medium

Scope-risk: broad

Directive: Do not reintroduce silent legacy fallback in owner-lane paths; unexpected owner-lane failure must be warning-level fail-closed or recoverable capacity wait.

Tested: Remote g0034 container PYTHONPATH=python python -m pytest test/registered/unit/mem_cache/test_alloc_pages_with_owners.py test/registered/unit/mem_cache/test_cp_shared_kv_layout.py test/registered/unit/mem_cache/test_cp_hicache_load_back_owner_lanes.py test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py -q -> 95 passed.

Tested: Local py_compile for modified runtime/cache/scheduler modules.

Not-tested: Full CUDA ETE performance trace for cache-hit overlap and MTP accept-rate impact.

Co-authored-by: OmX <omx@oh-my-codex.dev>
This commit is contained in:
laoyao0822
2026-05-28 05:54:23 +08:00
parent 2c94b8de23
commit ff33446787
12 changed files with 2913 additions and 136 deletions

View File

@@ -45,6 +45,50 @@ def _index_prefetch_fallback_log(reason: str, message: str, *args) -> None:
)
def _prefetch_pool_get_key_buffer(
*,
token_to_kv_pool: Any,
layer_id: int,
stream: torch.cuda.Stream,
path: str,
) -> torch.Tensor:
getter = getattr(token_to_kv_pool, "get_key_buffer_for_prefetch", None)
if getter is not None:
return getter(layer_id, stream)
if getattr(token_to_kv_pool, "layer_transfer_counter", None) is not None:
logger.warning(
"[CP_SHARED_KV_FALLBACK][%s_prefetch] "
"reason=prefetch_safe_getter_unavailable layer_id=%s pool=%s "
"has_layer_transfer_counter=True",
path,
layer_id,
type(token_to_kv_pool).__name__,
)
return token_to_kv_pool.get_key_buffer(layer_id)
def _prefetch_pool_get_index_buffer(
*,
token_to_kv_pool: Any,
layer_id: int,
stream: torch.cuda.Stream,
) -> torch.Tensor:
getter = getattr(
token_to_kv_pool, "get_index_k_with_scale_buffer_for_prefetch", None
)
if getter is not None:
return getter(layer_id, stream)
if getattr(token_to_kv_pool, "layer_transfer_counter", None) is not None:
_index_prefetch_fallback_log(
"prefetch_safe_getter_unavailable",
"pool has active layer transfer counter but no prefetch-safe index "
"getter. layer_id=%s pool=%s",
layer_id,
type(token_to_kv_pool).__name__,
)
return token_to_kv_pool.get_index_k_with_scale_buffer(layer_id=layer_id)
def _is_cuda_stream_capturing() -> bool:
if not torch.cuda.is_available():
return False
@@ -235,9 +279,15 @@ class CpSharedKVMlaPrefetcher:
)
return None
prefetch_stream = stream if stream is not None else torch.cuda.Stream()
try:
first_layer_id = int(getattr(token_to_kv_pool, "start_layer", 0))
kv_cache = token_to_kv_pool.get_key_buffer(first_layer_id)
kv_cache = _prefetch_pool_get_key_buffer(
token_to_kv_pool=token_to_kv_pool,
layer_id=first_layer_id,
stream=prefetch_stream,
path="mla",
)
remap = get_or_build_shared_token_kv_slot_remap(
forward_batch,
kv_cache=kv_cache,
@@ -276,7 +326,7 @@ class CpSharedKVMlaPrefetcher:
dense_num_pages=remap.dense_num_pages,
owned_prefix_pages=owned_prefix_pages,
owned_total_pages=owned_total_pages,
stream=stream,
stream=prefetch_stream,
)
def _layer_in_pool(self, token_to_kv_pool: Any, layer_id: int) -> bool:
@@ -447,7 +497,12 @@ class CpSharedKVMlaPrefetcher:
return
try:
kv_cache = token_to_kv_pool.get_key_buffer(next_layer_id)
kv_cache = _prefetch_pool_get_key_buffer(
token_to_kv_pool=token_to_kv_pool,
layer_id=next_layer_id,
stream=self.stream,
path="mla",
)
except Exception:
logger.exception(
"Failed to get next-layer KV cache for CP shared KV MLA prefetch."
@@ -461,31 +516,61 @@ class CpSharedKVMlaPrefetcher:
return
try:
dense_kv_cache = kv_cache.new_zeros(
(self.dense_num_pages * self.page_size, *kv_cache.shape[1:])
)
self._log_next_layer(
next_layer_id,
"start_prefix_begin next_layer=%s start_slot=0 end_slot=%s "
"dense_rows=%s",
next_layer_id,
self.prefix_pages,
int(dense_kv_cache.shape[0]),
)
materialize_local_token_kv_page_slots_into(
kv_cache=kv_cache,
dense_kv_cache=dense_kv_cache,
slot_logical_pages=self.slot_logical_pages,
layout=self.layout,
page_size=self.page_size,
start_slot=0,
end_slot=self.prefix_pages,
)
current_stream = torch.cuda.current_stream()
self.stream.wait_stream(current_stream)
prefix_rows = slot_range_to_token_slice(
self.page_size,
0,
self.prefix_pages,
)
with torch.cuda.stream(self.stream):
dense_kv_cache = kv_cache.new_zeros(
(self.dense_num_pages * self.page_size, *kv_cache.shape[1:])
)
self._log_next_layer(
next_layer_id,
"start_prefix_begin next_layer=%s start_slot=0 end_slot=%s "
"dense_rows=%s",
next_layer_id,
self.prefix_pages,
int(dense_kv_cache.shape[0]),
)
materialize_local_token_kv_page_slots_into(
kv_cache=kv_cache,
dense_kv_cache=dense_kv_cache,
slot_logical_pages=self.slot_logical_pages,
layout=self.layout,
page_size=self.page_size,
start_slot=0,
end_slot=self.prefix_pages,
)
event = _all_reduce_materialized_buffer_async(
dense_kv_cache[prefix_rows],
cp_size=self.layout.cp_size,
stream=self.stream,
nvtx_source="mla.prefetch_prefix",
nvtx_layer_id=next_layer_id,
nvtx_cp_rank=self.layout.cp_rank,
nvtx_rows=(prefix_rows.start, prefix_rows.stop),
)
if event is None:
self.disabled = True
logger.warning(
"[CP_SHARED_KV_FALLBACK][mla_prefetch] "
"reason=async_reduce_unavailable layer_id=%s cp_rank=%s "
"cp_size=%s prefix_pages=%s total_slots=%s",
next_layer_id,
self.layout.cp_rank,
self.layout.cp_size,
self.prefix_pages,
self.total_slots,
)
self._log_next_layer(
next_layer_id,
"start_disable reason=async_reduce_unavailable next_layer=%s",
next_layer_id,
)
return
except Exception:
logger.exception("Failed to start CP shared KV MLA prefix prefetch.")
self.disabled = True
@@ -500,12 +585,14 @@ class CpSharedKVMlaPrefetcher:
layer_id=next_layer_id,
dense_kv_cache=dense_kv_cache,
prefix_rows=prefix_rows,
event=event,
)
self.handles[next_layer_id] = handle
self.pending_attention_handle = handle
self._log_next_layer(
next_layer_id,
"start next_layer=%s prefix_pages=%s prefix_rows=%s dense_rows=%s",
"start next_layer=%s prefix_pages=%s prefix_rows=%s dense_rows=%s "
"reduce_enqueued=True",
next_layer_id,
self.prefix_pages,
prefix_rows.stop - prefix_rows.start,
@@ -776,10 +863,13 @@ class CpSharedKVIndexPrefetcher:
)
return None
prefetch_stream = stream if stream is not None else torch.cuda.Stream()
try:
first_layer_id = int(getattr(token_to_kv_pool, "start_layer", 0))
page_buffer = token_to_kv_pool.get_index_k_with_scale_buffer(
layer_id=first_layer_id
page_buffer = _prefetch_pool_get_index_buffer(
token_to_kv_pool=token_to_kv_pool,
layer_id=first_layer_id,
stream=prefetch_stream,
)
remap = get_or_build_shared_paged_buffer_slot_remap(
forward_batch,
@@ -822,7 +912,7 @@ class CpSharedKVIndexPrefetcher:
dense_num_pages=remap.dense_num_pages,
owned_prefix_pages=owned_prefix_pages,
owned_total_pages=owned_total_pages,
stream=stream,
stream=prefetch_stream,
)
def _layer_in_pool(self, token_to_kv_pool: Any, layer_id: int) -> bool:
@@ -990,8 +1080,10 @@ class CpSharedKVIndexPrefetcher:
return
try:
page_buffer = token_to_kv_pool.get_index_k_with_scale_buffer(
layer_id=next_layer_id
page_buffer = _prefetch_pool_get_index_buffer(
token_to_kv_pool=token_to_kv_pool,
layer_id=next_layer_id,
stream=self.stream,
)
except Exception:
logger.exception(
@@ -1006,26 +1098,56 @@ class CpSharedKVIndexPrefetcher:
return
try:
dense_page_buffer = page_buffer.new_zeros(
(self.dense_num_pages, *page_buffer.shape[1:])
)
self._log_next_layer(
next_layer_id,
"index_start_prefix_begin next_layer=%s start_slot=0 "
"end_slot=%s dense_pages=%s",
next_layer_id,
self.prefix_pages,
int(dense_page_buffer.shape[0]),
)
materialize_local_paged_buffer_page_slots_into(
page_buffer=page_buffer,
dense_page_buffer=dense_page_buffer,
slot_logical_pages=self.slot_logical_pages,
layout=self.layout,
start_slot=0,
end_slot=self.prefix_pages,
)
current_stream = torch.cuda.current_stream()
self.stream.wait_stream(current_stream)
prefix_rows = slot_range_to_page_slice(0, self.prefix_pages)
with torch.cuda.stream(self.stream):
dense_page_buffer = page_buffer.new_zeros(
(self.dense_num_pages, *page_buffer.shape[1:])
)
self._log_next_layer(
next_layer_id,
"index_start_prefix_begin next_layer=%s start_slot=0 "
"end_slot=%s dense_pages=%s",
next_layer_id,
self.prefix_pages,
int(dense_page_buffer.shape[0]),
)
materialize_local_paged_buffer_page_slots_into(
page_buffer=page_buffer,
dense_page_buffer=dense_page_buffer,
slot_logical_pages=self.slot_logical_pages,
layout=self.layout,
start_slot=0,
end_slot=self.prefix_pages,
)
event = _all_reduce_materialized_buffer_async(
dense_page_buffer[prefix_rows],
cp_size=self.layout.cp_size,
stream=self.stream,
nvtx_source="index.prefetch_prefix",
nvtx_layer_id=next_layer_id,
nvtx_cp_rank=self.layout.cp_rank,
nvtx_rows=(prefix_rows.start, prefix_rows.stop),
)
if event is None:
self.disabled = True
_index_prefetch_fallback_log(
"async_reduce_unavailable",
"async reduce unavailable. layer_id=%s cp_rank=%s cp_size=%s "
"prefix_pages=%s total_slots=%s",
next_layer_id,
self.layout.cp_rank,
self.layout.cp_size,
self.prefix_pages,
self.total_slots,
)
self._log_next_layer(
next_layer_id,
"index_start_disable reason=async_reduce_unavailable next_layer=%s",
next_layer_id,
)
return
except Exception:
logger.exception("Failed to start CP shared KV index prefix prefetch.")
self.disabled = True
@@ -1040,12 +1162,14 @@ class CpSharedKVIndexPrefetcher:
layer_id=next_layer_id,
dense_page_buffer=dense_page_buffer,
prefix_rows=prefix_rows,
event=event,
)
self.handles[next_layer_id] = handle
self.pending_attention_handle = handle
self._log_next_layer(
next_layer_id,
"index_start next_layer=%s prefix_pages=%s dense_pages=%s",
"index_start next_layer=%s prefix_pages=%s dense_pages=%s "
"reduce_enqueued=True",
next_layer_id,
self.prefix_pages,
int(dense_page_buffer.shape[0]),

View File

@@ -61,6 +61,10 @@ class LayerLoadingEvent:
def wait(self, layer_index: int):
device_module.current_stream().wait_event(self.load_events[layer_index])
def wait_on_stream(self, layer_index: int, stream):
assert 0 <= layer_index < self._num_layers
stream.wait_event(self.load_events[layer_index])
@property
def finish_event(self):
return self.load_events[-1]
@@ -101,6 +105,12 @@ class LayerDoneCounter:
for consumer_index in self.consumer_indices:
self.events[consumer_index].wait(threshold)
def wait_until_on_stream(self, threshold: int, stream) -> None:
if not self.consumer_indices:
return
for consumer_index in self.consumer_indices:
self.events[consumer_index].wait_on_stream(threshold, stream)
def reset(self):
self.producer_index = -1
self.consumer_index = -1

View File

@@ -179,8 +179,9 @@ from sglang.srt.managers.scheduler_update_weights_mixin import (
)
from sglang.srt.managers.session_controller import SessionController
from sglang.srt.managers.utils import GenerationBatchResult, validate_input_length
from sglang.srt.mem_cache.base_prefix_cache import DecLockRefParams
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.common import release_kv_cache
from sglang.srt.mem_cache.common import KVCapacityWaitError, release_kv_cache
from sglang.srt.mem_cache.radix_cache import RadixCache
from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache
from sglang.srt.model_executor.forward_batch_info import ForwardMode, PPProxyTensors
@@ -2417,6 +2418,8 @@ class Scheduler(
if len(can_run_list) == 0:
return None
waiting_queue_before_prepare = list(self.waiting_queue)
chunked_req_before_prepare = self.chunked_req
can_run_set = set(can_run_list)
self.waiting_queue = [x for x in self.waiting_queue if x not in can_run_set]
if adder.preempt_list:
@@ -2456,7 +2459,20 @@ class Scheduler(
self.tree_cache.ready_to_load_host_cache()
)
new_batch.prepare_for_extend()
try:
new_batch.prepare_for_extend()
except KVCapacityWaitError as exc:
self._release_prefill_adder_locks(
can_run_list, skip_req=chunked_req_before_prepare
)
self.waiting_queue = waiting_queue_before_prepare
self.chunked_req = chunked_req_before_prepare
self.running_batch.batch_is_full = True
logger.warning(
"[CP_SHARED_KV_CAPACITY_WAIT] prefill allocation deferred: %s",
exc,
)
return None
# Record prefill stats for logging after forward
new_batch.prefill_stats = PrefillStats.from_adder(
@@ -2485,6 +2501,38 @@ class Scheduler(
return new_batch
def _release_prefill_adder_locks(
self, reqs: List[Any], *, skip_req: Optional[Any] = None
) -> None:
"""Release lock refs acquired by PrefillAdder for a failed prefill batch.
`PrefillAdder.add_one_req` persists one `inc_lock_ref` per staged
request after its short-lived scheduling lock context exits. If KV
allocation later reports a recoverable capacity wait, the batch never
runs, so those persistent refs must be dropped before the requests are
returned to the waiting queue.
"""
if getattr(self.tree_cache, "disable", False):
return
use_swa_params = self.tree_cache.supports_swa() and self.tree_cache.is_tree_cache()
for req in reqs:
if req is skip_req:
continue
last_node = getattr(req, "last_node", None)
if last_node is None:
continue
if use_swa_params:
params = DecLockRefParams(
swa_uuid_for_lock=getattr(req, "swa_uuid_for_lock", None)
)
self.tree_cache.dec_lock_ref(last_node, params)
if hasattr(req, "swa_uuid_for_lock"):
req.swa_uuid_for_lock = None
else:
self.tree_cache.dec_lock_ref(last_node)
def update_running_batch(self, batch: ScheduleBatch) -> Optional[ScheduleBatch]:
"""Update the current running decoding batch."""
initial_bs = batch.batch_size()

View File

@@ -78,6 +78,7 @@ class EvictParams:
num_tokens: int
swa_num_tokens: int = 0
mamba_num: int = 0
owner_lane_deficits: Optional[list[int]] = None
@dataclasses.dataclass

View File

@@ -32,6 +32,27 @@ MAMBA_STATE_PER_REQ_NO_CACHE = 1
logger = logging.getLogger(__name__)
class KVCapacityWaitError(RuntimeError):
"""Recoverable KV-capacity pressure with CP owner-lane diagnostics."""
def __init__(
self,
message: str,
*,
required_by_owner: list[int] | None = None,
available_by_owner: list[int] | None = None,
deficit_by_owner: list[int] | None = None,
):
self.required_by_owner = required_by_owner
self.available_by_owner = available_by_owner
self.deficit_by_owner = deficit_by_owner
super().__init__(
f"{message}; required_by_owner={required_by_owner} "
f"available_by_owner={available_by_owner} "
f"deficit_by_owner={deficit_by_owner}"
)
def _log_cp_shared_kv_alloc_fallback(
reason: str,
message: str,
@@ -341,7 +362,12 @@ def _evict_for_compute_owner_lanes(
before_available,
evictable_size,
)
evict_result = tree_cache.evict(EvictParams(num_tokens=evict_tokens))
evict_result = tree_cache.evict(
EvictParams(
num_tokens=evict_tokens,
owner_lane_deficits=[int(v) for v in deficits],
)
)
after_available = allocator.available_size()
evicted_tokens = getattr(evict_result, "num_tokens_evicted", 0)
logger.info(
@@ -367,7 +393,7 @@ def alloc_paged_token_slots_extend(
# Over estimate the number of tokens: assume each request needs a new page.
allocator = tree_cache.token_to_kv_pool_allocator
num_tokens = extend_num_tokens + len(seq_lens_cpu) * allocator.page_size
evict_result = evict_from_tree_cache(tree_cache, num_tokens)
evict_result = None
# logger.info(
# "[MemCache-alloc] alloc_paged_token_slots_extend: extend_num_tokens=%d batch_size=%d num_tokens=%d page_size=%d "
@@ -420,11 +446,11 @@ def alloc_paged_token_slots_extend(
"multi_batch" if len(prefix_lens_cpu) != 1 else "unknown"
)
state = None
if backup_state:
state = allocator.backup_state()
if page_compute_owners is not None:
state = None
if backup_state:
state = allocator.backup_state()
out_cache_loc = alloc_extend_compute_owner(
prefix_lens,
prefix_lens_cpu,
@@ -460,10 +486,10 @@ def alloc_paged_token_slots_extend(
required, available, deficits = compute_owner_lane_stats(
page_compute_owners
)
_log_cp_shared_kv_alloc_fallback(
"owner_lane_exhausted",
"failed to allocate pages from compute-owner lanes; "
"falling back to legacy page allocation. extend_num_tokens=%s page_size=%s "
logger.warning(
"[CP_SHARED_KV_FAIL_FAST][owner_lane_exhausted] "
"failed to allocate pages from compute-owner lanes after eviction; "
"raising recoverable capacity wait. extend_num_tokens=%s page_size=%s "
"required_by_owner=%s available_by_owner=%s deficit_by_owner=%s",
extend_num_tokens,
allocator.page_size,
@@ -471,15 +497,19 @@ def alloc_paged_token_slots_extend(
available,
deficits,
)
out_cache_loc = allocator.alloc_extend(
prefix_lens,
prefix_lens_cpu,
seq_lens,
seq_lens_cpu,
last_loc,
extend_num_tokens,
raise KVCapacityWaitError(
"owner_lane_exhausted: failed to allocate pages from CP "
"compute-owner lanes after owner-lane eviction",
required_by_owner=required,
available_by_owner=available,
deficit_by_owner=deficits,
)
else:
evict_result = evict_from_tree_cache(tree_cache, num_tokens)
state = None
if backup_state:
state = allocator.backup_state()
if alloc_extend_compute_owner is not None:
_log_cp_shared_kv_alloc_fallback(
compute_owner_unavailable_reason or "compute_owner_not_available",
@@ -568,6 +598,7 @@ def alloc_for_extend(
extend_lens_device = extend_lens_cpu.to(batch.device, non_blocking=True)
# Allocate req slots
newly_allocated_reqs = [req for req in batch.reqs if req.req_pool_idx is None]
req_pool_indices = alloc_req_slots(
batch.req_to_token_pool, batch.reqs, batch.tree_cache
)
@@ -575,23 +606,29 @@ def alloc_for_extend(
req_pool_indices_device = req_pool_indices_cpu.to(batch.device, non_blocking=True)
# Allocate KV cache (throws exception on failure)
if batch.tree_cache.page_size == 1:
out_cache_loc = alloc_token_slots(batch.tree_cache, batch.extend_num_tokens)
else:
# Paged allocation - build last_loc
last_loc = [
(t[-1:] if len(t) > 0 else torch.tensor([-1], device=batch.device))
for t in prefix_tensors
]
out_cache_loc = alloc_paged_token_slots_extend(
tree_cache=batch.tree_cache,
prefix_lens=prefix_lens_device,
prefix_lens_cpu=prefix_lens_cpu,
seq_lens=batch.seq_lens,
seq_lens_cpu=batch.seq_lens_cpu,
last_loc=torch.cat(last_loc),
extend_num_tokens=batch.extend_num_tokens,
)
try:
if batch.tree_cache.page_size == 1:
out_cache_loc = alloc_token_slots(batch.tree_cache, batch.extend_num_tokens)
else:
# Paged allocation - build last_loc
last_loc = [
(t[-1:] if len(t) > 0 else torch.tensor([-1], device=batch.device))
for t in prefix_tensors
]
out_cache_loc = alloc_paged_token_slots_extend(
tree_cache=batch.tree_cache,
prefix_lens=prefix_lens_device,
prefix_lens_cpu=prefix_lens_cpu,
seq_lens=batch.seq_lens,
seq_lens_cpu=batch.seq_lens_cpu,
last_loc=torch.cat(last_loc),
extend_num_tokens=batch.extend_num_tokens,
)
except KVCapacityWaitError:
for req in newly_allocated_reqs:
if req.req_pool_idx is not None:
batch.req_to_token_pool.free(req)
raise
# Write to req_to_token_pool
write_cache_indices(

View File

@@ -33,7 +33,6 @@ from sglang.srt.mem_cache.base_prefix_cache import (
MatchPrefixParams,
MatchResult,
)
from sglang.srt.mem_cache.common import _evict_for_compute_owner_lanes
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
from sglang.srt.mem_cache.memory_pool import (
MHATokenToKVPool,
@@ -288,6 +287,22 @@ class CpHiCacheEvictionPlan:
remaining_deficit: Tuple[int, ...]
@dataclass(frozen=True)
class CpLoadBackPlan:
page_owners: List[int]
required_by_owner: List[int]
available_by_owner: List[int]
deficit_by_owner: List[int]
host_hit_len: int
@dataclass(frozen=True)
class CpLoadBackEvictionPlan:
victims: Tuple[TreeNode, ...]
planned_freed_by_owner: Tuple[int, ...]
remaining_deficit_by_owner: Tuple[int, ...]
class HiCachePendingBackupSplit(Exception):
def __init__(self, node: TreeNode):
self.node = node
@@ -869,6 +884,249 @@ class HiRadixCache(RadixCache):
remaining_deficit=tuple(deficits),
)
def _build_cp_load_back_plan(
self, nodes_to_load: List[TreeNode], *, node_id: int
) -> CpLoadBackPlan:
allocator = self.token_to_kv_pool_allocator
compute_owner_lane_stats = getattr(allocator, "compute_owner_lane_stats", None)
if compute_owner_lane_stats is None:
raise RuntimeError(
"CP HiCache load-back requires an allocator with "
f"compute_owner_lane_stats: node_id={node_id}"
)
page_owners: List[int] = []
host_hit_len = 0
for node in nodes_to_load:
metadata = getattr(node, "cp_hicache", None)
if metadata is None:
raise RuntimeError(
"CP HiCache load-back missing cp_hicache metadata: "
f"load_node_id={getattr(node, 'id', '?')} root_node_id={node_id}"
)
node_host_len = int(self._node_host_len(node))
if node_host_len % self.page_size != 0:
raise RuntimeError(
"CP HiCache load-back requires page-aligned host lengths: "
f"load_node_id={getattr(node, 'id', '?')} "
f"host_len={node_host_len} page_size={self.page_size}"
)
if int(getattr(metadata, "logical_len", node_host_len)) != node_host_len:
raise RuntimeError(
"CP HiCache load-back metadata length mismatch: "
f"load_node_id={getattr(node, 'id', '?')} "
f"host_len={node_host_len} metadata_len={metadata.logical_len}"
)
page_owners.extend(int(owner) for owner in metadata.page_owners.tolist())
host_hit_len += node_host_len
expected_pages = host_hit_len // self.page_size
if len(page_owners) != expected_pages:
raise RuntimeError(
"CP HiCache load-back page owner count mismatch: "
f"node_id={node_id} host_hit_len={host_hit_len} "
f"page_size={self.page_size} page_owners={len(page_owners)}"
)
required, available, deficits = compute_owner_lane_stats(page_owners)
return CpLoadBackPlan(
page_owners=list(page_owners),
required_by_owner=[int(v) for v in required],
available_by_owner=[int(v) for v in available],
deficit_by_owner=[int(v) for v in deficits],
host_hit_len=host_hit_len,
)
def _refresh_cp_load_back_plan(self, plan: CpLoadBackPlan) -> CpLoadBackPlan:
required, available, deficits = (
self.token_to_kv_pool_allocator.compute_owner_lane_stats(plan.page_owners)
)
return CpLoadBackPlan(
page_owners=plan.page_owners,
required_by_owner=[int(v) for v in required],
available_by_owner=[int(v) for v in available],
deficit_by_owner=[int(v) for v in deficits],
host_hit_len=plan.host_hit_len,
)
def _cp_device_leaf_is_load_back_victim(self, node: TreeNode) -> bool:
if node == getattr(self, "root_node", None):
return False
if getattr(node, "lock_ref", 0) > 0:
return False
if getattr(node, "value", None) is None:
return False
parent = getattr(node, "parent", None)
if parent is None:
return False
try:
parent_key = self.get_child_key_fn(node.key)
except Exception:
return False
if getattr(parent, "children", {}).get(parent_key) is not node:
return False
if self._is_pinned(node):
return False
return True
def _cp_load_back_node_owner_page_counts(
self, node: TreeNode, cp_size: int
) -> Tuple[int, ...]:
value = getattr(node, "value", None)
if value is None or int(value.numel()) == 0:
return tuple(0 for _ in range(cp_size))
if int(value.numel()) % self.page_size != 0:
raise RuntimeError(
"CP HiCache load-back eviction requires page-aligned device values: "
f"node_id={getattr(node, 'id', '?')} value_len={value.numel()} "
f"page_size={self.page_size}"
)
first_locs = value[:: self.page_size]
logical_pages = torch.div(first_locs, self.page_size, rounding_mode="floor")
owners = torch.remainder(logical_pages - 1, cp_size)
return tuple(int((owners == owner).sum().item()) for owner in range(cp_size))
def _plan_cp_load_back_owner_lane_evictions(
self, plan: CpLoadBackPlan
) -> CpLoadBackEvictionPlan:
deficits = [max(0, int(v)) for v in plan.deficit_by_owner]
cp_size = len(deficits)
planned_freed = [0 for _ in range(cp_size)]
victims: List[TreeNode] = []
used_node_ids = set()
while any(v > 0 for v in deficits):
best_node = None
best_counts = None
best_score = None
for node in list(getattr(self, "evictable_leaves", set())):
node_id = getattr(node, "id", None)
if node_id in used_node_ids:
continue
if not self._cp_device_leaf_is_load_back_victim(node):
continue
counts = self._cp_load_back_node_owner_page_counts(node, cp_size)
contribution = sum(
min(int(count), int(deficit))
for count, deficit in zip(counts, deficits)
)
if contribution <= 0:
continue
score = (
-int(contribution),
self.eviction_strategy.get_priority(node),
int(node_id or 0),
)
if best_score is None or score < best_score:
best_score = score
best_node = node
best_counts = counts
if best_node is None or best_counts is None:
break
victims.append(best_node)
used_node_ids.add(getattr(best_node, "id", None))
for owner, count in enumerate(best_counts):
planned_freed[owner] += int(count)
deficits[owner] = max(0, deficits[owner] - int(count))
return CpLoadBackEvictionPlan(
victims=tuple(victims),
planned_freed_by_owner=tuple(planned_freed),
remaining_deficit_by_owner=tuple(deficits),
)
def _evict_cp_load_back_owner_lanes(
self, plan: CpLoadBackPlan, *, node_id: int
) -> CpLoadBackPlan:
eviction_plan = self._plan_cp_load_back_owner_lane_evictions(plan)
if any(v > 0 for v in eviction_plan.remaining_deficit_by_owner):
logger.warning(
"[CP_HICACHE_FALLBACK][cp_load_back_owner_lane_eviction_insufficient] "
"node_id=%d required_by_owner=%s available_by_owner=%s "
"deficit_by_owner=%s planned_freed_by_owner=%s "
"remaining_deficit_by_owner=%s evictable_size=%d protected_size=%d "
"victims=%s allocator_state=%s",
node_id,
plan.required_by_owner,
plan.available_by_owner,
plan.deficit_by_owner,
eviction_plan.planned_freed_by_owner,
eviction_plan.remaining_deficit_by_owner,
int(getattr(self, "evictable_size_", 0)),
int(getattr(self, "protected_size_", 0)),
[getattr(node, "id", None) for node in eviction_plan.victims],
self.token_to_kv_pool_allocator.allocator_state_str(),
)
return self._refresh_cp_load_back_plan(plan)
if len(eviction_plan.victims) == 0:
return plan
num_evicted = 0
write_back_nodes: List[TreeNode] = []
for victim in eviction_plan.victims:
victim_id = getattr(victim, "id", None)
if not self._cp_device_leaf_is_load_back_victim(victim):
raise RuntimeError(
"CP HiCache load-back owner-lane eviction plan diverged "
f"before reservation: node_id={node_id} victim_id={victim_id}"
)
if victim.pin_expiry > 0 and time.monotonic() > victim.pin_expiry:
self._clear_pin(victim)
if not self._node_backuped(victim):
if self.cache_controller.write_policy == "write_back":
logger.warning(
"[CP_HICACHE_FALLBACK][cp_load_back_owner_lane_write_back] "
"node_id=%d victim_id=%s reason=unbacked_write_back_victim",
node_id,
victim_id,
)
written = self.write_backup(victim, write_back=True)
if written > 0:
write_back_nodes.append(victim)
continue
num_evicted += self._evict_regular(victim)
else:
num_evicted += self._evict_backuped(victim)
for child in victim.parent.children.values():
if child in write_back_nodes:
continue
if not child.evicted:
break
else:
self._update_leaf_status(victim.parent)
if write_back_nodes:
self.writing_check(write_back=True)
for victim in write_back_nodes:
if self._node_backuped(victim):
num_evicted += self._evict_backuped(victim)
refreshed = self._refresh_cp_load_back_plan(plan)
logger.info(
"[HiCache-load] owner-lane device eviction before CP load-back: "
"node_id=%d victims=%s num_evicted=%d required_by_owner=%s "
"before_available_by_owner=%s before_deficit_by_owner=%s "
"after_available_by_owner=%s after_deficit_by_owner=%s "
"planned_freed_by_owner=%s",
node_id,
[getattr(node, "id", None) for node in eviction_plan.victims],
num_evicted,
plan.required_by_owner,
plan.available_by_owner,
plan.deficit_by_owner,
refreshed.available_by_owner,
refreshed.deficit_by_owner,
eviction_plan.planned_freed_by_owner,
)
return refreshed
def _cp_capacity_debug_enabled(self) -> bool:
return bool(
getattr(self, "_cp_hicache_capacity_debug", False)
@@ -2470,14 +2728,28 @@ class HiRadixCache(RadixCache):
result = self.inc_lock_ref(ancester_node)
delta = result.delta
# load it all or not at all
host_hit_len = sum(self._node_host_len(n) for n in nodes_to_load)
# load it all or not at all. The scalar length remains only a
# coarse upper bound; final CP admission is the owner-lane vector
# check below.
try:
load_back_plan = self._build_cp_load_back_plan(
nodes_to_load, node_id=last_hit_node.id
)
except Exception:
self.dec_lock_ref(ancester_node)
raise
host_hit_len = load_back_plan.host_hit_len
logger.info(
"[HiCache-load] load_back CP: node_id=%d nodes_to_load=%d host_hit_len=%d threshold=%d",
"[HiCache-load] load_back CP: node_id=%d nodes_to_load=%d "
"host_hit_len=%d threshold=%d required_by_owner=%s "
"available_by_owner=%s deficit_by_owner=%s",
last_hit_node.id,
len(nodes_to_load),
host_hit_len,
self.load_back_threshold,
load_back_plan.required_by_owner,
load_back_plan.available_by_owner,
load_back_plan.deficit_by_owner,
)
if host_hit_len < self.load_back_threshold or (
host_hit_len > mem_quota + delta if mem_quota is not None else False
@@ -2491,32 +2763,48 @@ class HiRadixCache(RadixCache):
)
return None
if any(v > 0 for v in load_back_plan.deficit_by_owner):
load_back_plan = self._evict_cp_load_back_owner_lanes(
load_back_plan, node_id=last_hit_node.id
)
if any(v > 0 for v in load_back_plan.deficit_by_owner):
self.dec_lock_ref(ancester_node)
logger.warning(
"[CP_HICACHE_FALLBACK][cp_load_back_owner_lane_capacity_failed] "
"node_id=%d host_hit_len=%d nodes_to_load=%d "
"required_by_owner=%s available_by_owner=%s "
"deficit_by_owner=%s evictable_size=%d protected_size=%d "
"ongoing_load_back_count=%d allocator_state=%s",
last_hit_node.id,
host_hit_len,
len(nodes_to_load),
load_back_plan.required_by_owner,
load_back_plan.available_by_owner,
load_back_plan.deficit_by_owner,
int(getattr(self, "evictable_size_", 0)),
int(getattr(self, "protected_size_", 0)),
len(getattr(self, "ongoing_load_back", {})),
self.token_to_kv_pool_allocator.allocator_state_str(),
)
return None
device_indices = self.cache_controller.load_cp(
nodes_to_load, node_id=last_hit_node.id
)
if device_indices is None:
logger.info(
"[HiCache-load] load_back CP retry with lane-aware eviction: "
"node_id=%d tokens_needed=%d",
failed_plan = self._refresh_cp_load_back_plan(load_back_plan)
logger.warning(
"[CP_HICACHE_FALLBACK][cp_load_back_preflight_mismatch] "
"node_id=%d host_hit_len=%d required_by_owner=%s "
"available_by_owner=%s deficit_by_owner=%s "
"allocator_state=%s",
last_hit_node.id,
host_hit_len,
)
# Lane-aware eviction: alloc_pages_with_owners failed because
# at least one owner lane is short of free pages. Targeted
# eviction frees pages from THOSE specific lanes, leaving
# other lanes untouched. Mirrors cold prefill's retry path
# in common.py:alloc_paged_token_slots_extend.
_retry_page_owners: List[int] = []
for _node in nodes_to_load:
# page_owners is CPU int8; tolist() returns list[int].
_retry_page_owners.extend(_node.cp_hicache.page_owners.tolist())
_evict_for_compute_owner_lanes(
tree_cache=self,
allocator=self.token_to_kv_pool_allocator,
page_compute_owners=_retry_page_owners,
)
device_indices = self.cache_controller.load_cp(
nodes_to_load, node_id=last_hit_node.id
failed_plan.required_by_owner,
failed_plan.available_by_owner,
failed_plan.deficit_by_owner,
self.token_to_kv_pool_allocator.allocator_state_str(),
)
self.dec_lock_ref(ancester_node)
if device_indices is None:

View File

@@ -746,6 +746,32 @@ class KVCache(abc.ABC):
def register_layer_transfer_counter(self, layer_transfer_counter: LayerDoneCounter):
self.layer_transfer_counter = layer_transfer_counter
def wait_layer_transfer_on_stream(self, layer_id: int, stream) -> None:
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until_on_stream(
layer_id - self.start_layer, stream
)
def get_key_buffer_for_prefetch(self, layer_id: int, stream) -> torch.Tensor:
self.wait_layer_transfer_on_stream(layer_id, stream)
raw_getter = getattr(self, "_get_key_buffer", None)
if raw_getter is None:
raise NotImplementedError(
f"{type(self).__name__} does not expose a non-blocking key-buffer getter"
)
return raw_getter(layer_id)
def get_index_k_with_scale_buffer_for_prefetch(
self, layer_id: int, stream
) -> torch.Tensor:
self.wait_layer_transfer_on_stream(layer_id, stream)
raw_getter = getattr(self, "_get_index_k_with_scale_buffer", None)
if raw_getter is None:
raise NotImplementedError(
f"{type(self).__name__} does not expose a non-blocking index-buffer getter"
)
return raw_getter(layer_id)
def register_layer_backup_notifier(self, notifier):
self.layer_backup_notifiers.append(notifier)
@@ -1379,6 +1405,10 @@ class HybridLinearKVPool(KVCache):
layer_id = self._transfer_full_attention_id(layer_id)
return self.full_kv_pool.get_key_buffer(layer_id)
def get_key_buffer_for_prefetch(self, layer_id: int, stream):
layer_id = self._transfer_full_attention_id(layer_id)
return self.full_kv_pool.get_key_buffer_for_prefetch(layer_id, stream)
def get_value_buffer(self, layer_id: int):
layer_id = self._transfer_full_attention_id(layer_id)
return self.full_kv_pool.get_value_buffer(layer_id)
@@ -1552,25 +1582,29 @@ class MLATokenToKVPool(KVCache):
return _copy_buf_infos(self._contiguous_buf_infos)
def get_key_buffer(self, layer_id: int):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
def _get_key_buffer(self, layer_id: int):
if self.store_dtype != self.dtype:
return self.kv_buffer[layer_id - self.start_layer].view(self.dtype)
return self.kv_buffer[layer_id - self.start_layer]
def get_value_buffer(self, layer_id: int):
def get_key_buffer(self, layer_id: int):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
return self._get_key_buffer(layer_id)
def _get_value_buffer(self, layer_id: int):
if self.store_dtype != self.dtype:
return self.kv_buffer[layer_id - self.start_layer][
..., : self.kv_lora_rank
].view(self.dtype)
return self.kv_buffer[layer_id - self.start_layer][..., : self.kv_lora_rank]
def get_value_buffer(self, layer_id: int):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
return self._get_value_buffer(layer_id)
def get_kv_buffer(self, layer_id: int):
return self.get_key_buffer(layer_id), self.get_value_buffer(layer_id)
@@ -1724,10 +1758,7 @@ class MLATokenToKVPoolFP4(MLATokenToKVPool):
del self.kv_scale_buffer
self._contiguous_buf_infos = None
def get_key_buffer(self, layer_id: int):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
def _get_key_buffer(self, layer_id: int):
if self.store_dtype != self.dtype:
cache_k_nope_fp4 = self.kv_buffer[layer_id - self.start_layer].view(
torch.uint8
@@ -1743,6 +1774,11 @@ class MLATokenToKVPoolFP4(MLATokenToKVPool):
return self.kv_buffer[layer_id - self.start_layer]
def get_key_buffer(self, layer_id: int):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
return self._get_key_buffer(layer_id)
def set_kv_buffer(
self,
layer: RadixAttention,
@@ -1893,10 +1929,13 @@ class NSATokenToKVPool(MLATokenToKVPool):
]
self._finalize_allocation_log(size)
def _get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor:
return self.index_k_with_scale_buffer[layer_id - self.start_layer]
def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor:
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
return self.index_k_with_scale_buffer[layer_id - self.start_layer]
return self._get_index_k_with_scale_buffer(layer_id)
def get_index_k_continuous(
self,