Files
2026-07-13 12:24:33 +08:00

1869 lines
79 KiB
Python

# SPDX-License-Identifier: Apache-2.0
"""Blend V3: paged-aware CacheBlend as an EngineModule.
Plugs into the unified MPCacheServer; standard REGISTER_KV_CACHE +
CB_REGISTER_ROPE_V3 for setup; STORE wrapper registers fingerprints;
retrieve scatters into the request's paged blocks.
"""
# Standard
from dataclasses import dataclass
from queue import Empty as QueueEmpty
from queue import Queue
from typing import TYPE_CHECKING, Any
import threading
import time
if TYPE_CHECKING:
# First Party
from lmcache.v1.mp_coordinator.blend_client import (
BlendCoordinatorClient,
RemoteMatch,
)
# Third Party
import numpy as np
import torch
# First Party
from lmcache import torch_dev, torch_device_type
from lmcache.logging import init_logger
from lmcache.utils import check_interprocess_event_support
from lmcache.v1.distributed.api import (
MemoryLayoutDesc,
TrimPolicy,
ipc_key_to_object_keys,
)
from lmcache.v1.distributed.storage_manager import PrefetchHandle
from lmcache.v1.gpu_connector.gpu_ops import lmcache_memcpy_async_h2d
from lmcache.v1.mp_coordinator.blend_client import PENDING
from lmcache.v1.mp_observability.event import Event, EventType
from lmcache.v1.multiprocess.custom_types import (
CBMatchResult,
CBUnifiedLookupResult,
DeviceIPCWrapper,
IPCCacheServerKey,
)
from lmcache.v1.multiprocess.engine_context import MPCacheServerContext
from lmcache.v1.multiprocess.engine_module import (
HandlerSpec,
InstanceLivenessTarget,
ThreadPoolType,
)
from lmcache.v1.multiprocess.modules.lmcache_driven_transfer import (
LMCacheDrivenTransferModule,
)
from lmcache.v1.multiprocess.modules.lookup import compute_extra_count
from lmcache.v1.multiprocess.protocol import RequestType
from lmcache.v1.multiprocess.token_hasher import (
TokenHasher,
chunk_hash_windows_numba,
rolling_hash_windows_numba,
update_table_id_numba,
)
from lmcache.v1.platform.base_cache_context import BaseCacheContext
import lmcache.c_ops as lmc_ops
logger = init_logger(__name__)
@dataclass
class _CBRopeState:
"""Per-instance RoPE state IPC-shared from vLLM; dangles on reallocate.
Models with per-layer-type RoPE (distinct local/global theta)
register one cache per distinct rope and a per-layer index into
``cos_sin_caches``.
"""
head_size: int
is_neox_style: bool # NeoX = contiguous halves; else GPT-J.
cos_sin_caches: list[torch.Tensor]
group_to_cache: list[int] # engine group idx -> cache idx; empty = cache 0
def cache_for_group(self, engine_group_idx: int) -> torch.Tensor:
"""The cos/sin cache for one engine group.
Engine groups partition layers by attention type, and rope follows
attention type (sliding=local theta, full=global theta),
so each engine group has exactly one cache.
Args:
engine_group_idx: The kernel group's engine group index.
Returns:
The group's cos/sin cache tensor.
Raises:
RuntimeError: If ``engine_group_idx`` is outside the map.
"""
if not self.group_to_cache:
return self.cos_sin_caches[0]
if engine_group_idx >= len(self.group_to_cache):
raise RuntimeError(
f"CB re-RoPE: engine group {engine_group_idx} has no rope "
f"cache mapping (map covers {len(self.group_to_cache)} groups)."
)
return self.cos_sin_caches[self.group_to_cache[engine_group_idx]]
@dataclass
class _CBUnifiedJob:
"""Per-request poll state for non-blocking cb_unified_lookup.
Stashed across polls because the underlying status/found polls are
consume-once.
"""
matches: list[CBMatchResult]
num_tokens: int = 0
# Prefix leg (blend_v3-owned submit/poll). ``prefix_handle`` is None when
# there is no GPU context / no full chunk (poll reports 0 coverage).
prefix_handle: PrefetchHandle | None = None
prefix_world_size: int = 1
prefix_chunks: int | None = None # stashed when the prefix poll completes
retained_chunks: list[int] | None = None # SEGMENTED_PREFIX: full gapped set
sparse_started: bool = False # prefix done -> sparse leg submitted/skipped
handle: PrefetchHandle | None = None # sparse handle, None if no sparse leg
non_prefix: list[CBMatchResult] | None = None
per_hash_obj_keys: dict | None = None
expanded_uidx: list[int] | None = None
found_uidx: set[int] | None = None # stashed when the sparse poll completes
l2_keys: int = 0 # sparse keys needing an L2 load (0 => no L2 read, span skipped)
coord_submitted: bool = False # coordinator match query was issued
coord_deadline: float = 0.0 # time.monotonic() wall-clock cutoff for the leg
class BlendTokenRangeMatcherV3:
"""V3 matcher: token-level probe (any offset) + full-hash collision
rejection. Self-contained (does not inherit a base matcher)."""
_TABLE_BITS: int = 20 # 2^20 ~ 1 M entries
_TABLE_SIZE: int = 1 << _TABLE_BITS
_BASE: np.uint64 = np.uint64(0x9E3779B97F4A7C15) # Fibonacci-hashing const
def __init__(self, chunk_size: int = 256):
"""Initialize the V3 matcher.
Args:
chunk_size (int): Tokens per non-overlapping fingerprint chunk.
"""
self.chunk_size = chunk_size
# poly_chunk_hash -> compact_chunk_id; -1 = empty
self._table_id = np.full(self._TABLE_SIZE, -1, dtype=np.int64)
self._mask = np.uint64(self._TABLE_SIZE - 1)
# compact_chunk_id -> caller token_hash (full bytes); None once evicted
self._chunk_token_hash: list[bytes | None] = []
# token_hash -> start position in its registered sequence
self._token_hash_to_start: dict[bytes, int] = {}
# compact_chunk_id -> table slot (reverse lookup for eviction)
self._compact_id_to_slot = np.full(self._TABLE_SIZE, -1, dtype=np.int64)
# token_hash -> compact_chunk_id (for eviction lookup)
self._token_hash_to_compact_id: dict[bytes, int] = {}
self._lock = threading.Lock()
# V3 addition: compact_chunk_id -> full poly hash, for collision reject.
self._chunk_poly_hash: list[int] = []
def on_new_token_hashes(
self,
token_ids: list[int],
token_hashes: list[bytes],
start_chunk_idx: int = 0,
position_offset: int = 0,
) -> None:
"""Index a stored sequence's non-overlapping chunks into the matcher.
Records each new chunk's poly hash + start position so a later
match_sub_sequence can find it. Thread-safe (holds the matcher lock).
Args:
token_ids (list[int]): The stored sequence's token IDs.
token_hashes (list[bytes]): Per-chunk content hashes (one per
chunk), used as the dedup/eviction key.
start_chunk_idx (int): First chunk to index; 1 skips chunk 0 (the
prefix lookup leg owns it).
position_offset (int): Added to each recorded start position (for
indexing a tail-slice of a larger sequence).
Returns:
None.
"""
arr = np.array(token_ids, dtype=np.uint64)
chunk_hashes = chunk_hash_windows_numba(arr, self.chunk_size, self._BASE)
n = int(chunk_hashes.shape[0])
if n == 0 or start_chunk_idx >= n:
return
with self._lock:
new_idxs = [
i
for i in range(start_chunk_idx, n)
if token_hashes[i] not in self._token_hash_to_compact_id
]
if not new_idxs:
return
n_new = len(new_idxs)
new_chunk_hashes = chunk_hashes[new_idxs]
base_id = len(self._chunk_token_hash)
if base_id + n_new > self._TABLE_SIZE:
logger.error(
"BlendTokenRangeMatcherV3 compact-ID overflow: %d chunks "
"registered, cannot add %d more (limit %d). Skipping.",
base_id,
n_new,
self._TABLE_SIZE,
)
return
if base_id + n_new > int(self._TABLE_SIZE * 0.8):
logger.warning(
"BlendTokenRangeMatcherV3 nearing capacity: %d/%d "
"compact IDs used. Hash collision rate is rising; "
"hit rate will degrade.",
base_id + n_new,
self._TABLE_SIZE,
)
compact_ids = np.arange(base_id, base_id + n_new, dtype=np.int64)
update_table_id_numba(new_chunk_hashes, self._table_id, compact_ids)
for k, orig_i in enumerate(new_idxs):
th = token_hashes[orig_i]
cid = int(compact_ids[k])
poly_hash = int(new_chunk_hashes[k])
slot = poly_hash & int(self._mask)
self._chunk_token_hash.append(th)
self._chunk_poly_hash.append(poly_hash)
self._token_hash_to_start[th] = (
position_offset + orig_i * self.chunk_size
)
self._compact_id_to_slot[cid] = slot
self._token_hash_to_compact_id[th] = cid
def match_sub_sequence(
self,
token_ids: list[int],
) -> list[CBMatchResult]:
"""Find every registered chunk reused anywhere in a query sequence.
Vectorized direct-address probe over all token positions, then a small
verify loop over the surviving hits (a full poly-hash check rejects
bucket collisions; evicted/unknown chunks are skipped). Thread-safe.
Args:
token_ids (list[int]): The query sequence's token IDs.
Returns:
list[CBMatchResult]: One result per unique reused chunk (cur_st
= its first query position, old_st = its stored position).
Empty if the query is shorter than one chunk or nothing matched.
"""
if len(token_ids) < self.chunk_size:
return []
arr = np.array(token_ids, dtype=np.uint64)
rolling = rolling_hash_windows_numba(arr, self.chunk_size, self._BASE)
with self._lock:
if not self._chunk_token_hash:
return []
# Vectorized direct-address probe over all positions. The table is
# sparse (TABLE_SIZE >> registered chunks), so only true matches and
# a few bucket collisions reach the Python verify loop below.
cids_at_pos = self._table_id[rolling & self._mask]
hit_positions = np.nonzero(cids_at_pos >= 0)[0]
seen_cids: set[int] = set()
results: list[CBMatchResult] = []
for pos in hit_positions:
pos = int(pos)
cid = int(cids_at_pos[pos])
if cid in seen_cids:
continue
if int(rolling[pos]) != self._chunk_poly_hash[cid]:
continue # bucket-only collision
th = self._chunk_token_hash[cid]
if th is None:
continue # evicted
old_st = self._token_hash_to_start.get(th)
if old_st is None:
continue
seen_cids.add(cid)
results.append(
CBMatchResult(
old_st=old_st,
old_ed=old_st + self.chunk_size,
cur_st=pos,
cur_ed=pos + self.chunk_size,
hash=th,
)
)
logger.info(
"[match_probe] n_tok=%d table_hits=%d matches=%d",
len(token_ids),
len(hit_positions),
len(results),
)
return results
def remove_chunks(self, token_hashes: list[bytes]) -> None:
"""Evict the given chunks from the matcher.
Clears each chunk's table slot + poly hash so later probes cannot match
it. Thread-safe.
Args:
token_hashes (list[bytes]): Content hashes of the chunks to evict.
"""
with self._lock:
for th in token_hashes:
cid = self._token_hash_to_compact_id.get(th)
if cid is None:
continue
slot = int(self._compact_id_to_slot[cid])
if slot < 0:
logger.warning(
"compact_id %d has no valid table slot; "
"entry may have been evicted twice",
cid,
)
continue
self._table_id[slot] = -1
self._compact_id_to_slot[cid] = -1
self._chunk_token_hash[cid] = None
self._chunk_poly_hash[cid] = 0
self._token_hash_to_start.pop(th, None)
del self._token_hash_to_compact_id[th]
def _unique_token_coverage(results: list[CBMatchResult]) -> int:
"""Total token coverage, merging overlapping ranges (sliding-window probe
can return overlaps; naive sum would double-count)."""
if not results:
return 0
intervals = sorted((r.cur_st, r.cur_ed) for r in results)
coverage = 0
cur_end = -1
for st, ed in intervals:
if st >= cur_end:
coverage += ed - st
elif ed > cur_end:
coverage += ed - cur_end
cur_end = max(cur_end, ed)
return coverage
class BlendV3Module(InstanceLivenessTarget):
"""Paged-aware V3 CacheBlend. Wraps LMCacheDrivenTransfer STORE to register
fingerprints; serves CB rope/lookup/retrieve RPCs; reads cross-module
GPU state via :class:`LMCacheDrivenTransferModule.cache_contexts`."""
def __init__(
self,
ctx: MPCacheServerContext,
lmcache_driven_transfer: LMCacheDrivenTransferModule,
coordinator: "BlendCoordinatorClient | None" = None,
enable_segmented_prefix: bool = False,
):
self._ctx = ctx
self._transfer_module = lmcache_driven_transfer
# Server config (--enable-segmented-prefix): retain the gapped prefix on
# a mid-prefix L2 retrieve failure instead of truncating at the gap.
self._segmented_prefix = enable_segmented_prefix
# Optional bridge to the fleet-wide fingerprint directory. ``None`` =>
# purely local matching (publish/query paths skipped).
self._coordinator = coordinator
self._token_range_matcher = BlendTokenRangeMatcherV3(ctx.chunk_size)
self._event_bus = ctx.event_bus
self._cb_rope_state: dict[int, _CBRopeState] = {}
# L2 opt: cache TP-expanded obj_keys at lookup, pop at retrieve.
self._lookup_obj_keys_cache: dict[str, dict[bytes, list]] = {}
self._lookup_obj_keys_lock = threading.Lock()
# Non-blocking cb_unified_lookup poll state (submit-once, poll-on-recall)
# so the handler never holds a worker thread across the L2->L1 loads.
self._cb_jobs: dict[str, _CBUnifiedJob] = {}
self._cb_jobs_lock = threading.Lock()
# Async fingerprint registration: store enqueues, worker drains.
_FpJob = tuple[list[int], list[bytes], int, int]
self._fingerprint_queue: "Queue[_FpJob]" = Queue()
self._fingerprint_stop = threading.Event()
self._fingerprint_worker = threading.Thread(
target=self._drain_fingerprint_queue,
name="cb-fingerprint-worker",
daemon=True,
)
self._fingerprint_worker.start()
# In-flight fingerprint hashes; storage_gate keeps these from eviction.
self._pending_fp_hashes: set[bytes] = set()
self._pending_fp_lock = threading.Lock()
# Lazy eviction strikes; evict only at threshold so async re-store
# can refresh the bucket first.
self._stale_strike: dict[bytes, int] = {}
self._STALE_STRIKE_THRESHOLD = 2
# ------------------------------------------------------------------
# EngineModule protocol
# ------------------------------------------------------------------
@property
def context(self) -> MPCacheServerContext:
return self._ctx
def get_handlers(self) -> list[HandlerSpec]:
# STORE shadows LMCacheDrivenTransfer's; compositor registers V3 last.
return [
HandlerSpec(RequestType.STORE, self.store, ThreadPoolType.AFFINITY),
HandlerSpec(
RequestType.CB_REGISTER_ROPE_V3,
self.cb_register_rope,
ThreadPoolType.SYNC,
),
HandlerSpec(
RequestType.CB_UNREGISTER_ROPE_V3,
self.cb_unregister_rope,
ThreadPoolType.SYNC,
),
HandlerSpec(
RequestType.CB_UNIFIED_LOOKUP,
self.cb_unified_lookup,
ThreadPoolType.NORMAL,
),
HandlerSpec(
RequestType.CB_RETRIEVE_PRE_COMPUTED_V3,
self.cb_retrieve_pre_computed,
ThreadPoolType.AFFINITY,
),
]
def report_status(self) -> dict:
# Meta is derived live from MP server gpu_transfe
cache_contexts = self._transfer_module.context_entries_snapshot()
def _meta(iid: int) -> "tuple[str, int] | None":
entry = cache_contexts.get(iid)
return (entry.model_name, entry.world_size) if entry is not None else None
return {
"registered_cb_rope_instances": list(self._cb_rope_state.keys()),
"cb_rope_meta": {str(iid): _meta(iid) for iid in self._cb_rope_state},
"active_cb_lookups": len(self._cb_jobs),
}
def close(self) -> None:
self._fingerprint_stop.set()
if self._coordinator is not None:
# Joins the client's daemon thread and closes its httpx.Client;
# otherwise the coordinator leg leaks both on server shutdown.
self._coordinator.close()
self._cb_rope_state.clear()
# ------------------------------------------------------------------
# V3 RPCs
# ------------------------------------------------------------------
def cb_register_rope(
self,
instance_id: int,
cos_sin_caches_ipc: list[DeviceIPCWrapper],
head_size: int,
is_neox_style: bool,
group_to_cache: list[int],
) -> None:
"""Bolt CB re-RoPE state onto an already-registered KV-cache instance.
Idempotent; ``REGISTER_KV_CACHE`` must precede this. Strips any
YaRN/longrope mscale baked into each rope cache so re-RoPE stays a
pure rotation.
Args:
instance_id (int): KV-cache instance to attach rope state to.
cos_sin_caches_ipc (list[DeviceIPCWrapper]): IPC handles to vLLM's
cos/sin rope cache(s) — one per distinct rope (dual-RoPE
models send local/global); single-rope models send one.
head_size (int): Rotary head dimension.
is_neox_style (bool): True for NeoX (contiguous halves), else GPT-J.
group_to_cache (list[int]): Per-engine-group index into the
caches list; empty means every group uses cache 0.
Raises:
ValueError: If ``instance_id`` has no registered KV cache, the
cache list is empty, or ``group_to_cache`` references a
missing cache or does not cover every engine group of the
registered model.
"""
entry = self._transfer_module.get_and_touch_context_entry(instance_id)
if entry is None:
raise ValueError(
f"Instance {instance_id} has no paged KV cache registered; "
"send REGISTER_KV_CACHE before CB_REGISTER_ROPE_V3."
)
if not cos_sin_caches_ipc:
raise ValueError("CB_REGISTER_ROPE_V3 requires >=1 cos/sin cache.")
if group_to_cache:
if min(group_to_cache) < 0 or max(group_to_cache) >= len(
cos_sin_caches_ipc
):
raise ValueError(
f"group_to_cache {group_to_cache} contains indices outside "
f"[0, {len(cos_sin_caches_ipc)}) for the sent cache(s)."
)
# Fail at registration, not mid-retrieve: every engine group of
# the registered model must have a cache mapping.
max_eg_idx = max(
(
g.engine_group_idx
for g in entry.cache_context.kv_layer_groups_manager.kernel_groups
),
default=-1,
)
if len(group_to_cache) <= max_eg_idx:
raise ValueError(
f"group_to_cache covers {len(group_to_cache)} engine "
f"group(s) but the registered model has engine groups up "
f"to index {max_eg_idx}."
)
cos_sin_caches: list[torch.Tensor] = []
for cache_idx, cache_ipc in enumerate(cos_sin_caches_ipc):
cos_sin_cache = cache_ipc.to_tensor()
# YaRN/longrope bake an mscale m into the rope cache
# (cos²+sin²=m²≠1). vLLM already folds m into stored K, but CB
# re-RoPE assumes a pure rotation, so an un-normalized m injects
# an m² error per K element.
_c32 = cos_sin_cache.to(torch.float32)
_half = _c32.shape[1] // 2
_m = float((_c32[:, :_half] ** 2 + _c32[:, _half:] ** 2).mean().sqrt())
if abs(_m - 1.0) >= 1e-3:
logger.info(
"CB re-RoPE: cache %d: stripping rope-cache mscale=%.4f "
"(m²=%.4f → K inflation if uncorrected) → unit magnitude",
cache_idx,
_m,
_m * _m,
)
cos_sin_cache = (_c32 / _m).to(cos_sin_cache.dtype)
cos_sin_caches.append(cos_sin_cache)
self._cb_rope_state[instance_id] = _CBRopeState(
head_size=head_size,
is_neox_style=is_neox_style,
cos_sin_caches=cos_sin_caches,
group_to_cache=list(group_to_cache),
)
logger.info(
"Registered CB rope state for instance %d "
"(%d cache(s), shapes=%s dtype=%s, head_size=%d, is_neox=%s, "
"group_map=%s)",
instance_id,
len(cos_sin_caches),
[tuple(c.shape) for c in cos_sin_caches],
cos_sin_caches[0].dtype,
head_size,
is_neox_style,
"uniform" if not group_to_cache else str(group_to_cache),
)
def cb_unregister_rope(self, instance_id: int) -> None:
"""Drop the instance's CB rope state; the paged KV cache is left intact.
Args:
instance_id (int): Instance whose rope state to remove (use
``UNREGISTER_KV_CACHE`` to free the KV cache itself).
"""
self._cb_rope_state.pop(instance_id, None)
if self._transfer_module.get_and_touch_context_entry(instance_id) is None:
logger.warning(
"cb_unregister_rope: instance %d not registered", instance_id
)
return
logger.info("Unregistered CB rope state for instance %d", instance_id)
def drop_instance_state(self, instance_id: int) -> None:
"""Drop blend state for a reaped instance (InstanceLivenessTarget hook).
Only the CB rope state is held per instance; the GPU cache context is
owned by ``LMCacheDrivenTransferModule`` (no mirror here), so reaping
the GPU entry frees it directly. A no-op if no rope state is held.
Args:
instance_id: The reaped worker's instance ID.
"""
if self._cb_rope_state.pop(instance_id, None) is not None:
logger.info("Dropped CB rope state for reaped instance %d", instance_id)
# ------------------------------------------------------------------
# Unified lookup (CB_UNIFIED_LOOKUP) + shared helpers
# ------------------------------------------------------------------
def _drain_fingerprints_sync(self) -> None:
"""Sync-drain pending fingerprint registrations (the async drainer
races at low max_tokens)."""
while True:
try:
job = self._fingerprint_queue.get_nowait()
except QueueEmpty:
break
tokens_in_range, chunk_hashes, start_chunk_idx, position_offset = job
try:
self._token_range_matcher.on_new_token_hashes(
tokens_in_range,
chunk_hashes,
start_chunk_idx=start_chunk_idx,
position_offset=position_offset,
)
except Exception:
logger.exception("CB fingerprint registration failed (sync drain)")
def _match_fingerprints(self, key: IPCCacheServerKey) -> list[CBMatchResult]:
"""Drain pending registrations and fingerprint-match sub-sequences.
Returns the raw matches (any order, possibly overlapping); the caller
applies the prefix filter + overlap dedup once via
:meth:`_non_overlapping_after_prefix`.
"""
self._drain_fingerprints_sync()
return self._token_range_matcher.match_sub_sequence(list(key.token_ids))
@staticmethod
def _non_overlapping_after_prefix(
matches: list[CBMatchResult], prefix_tokens: int
) -> list[CBMatchResult]:
"""Matches outside the prefix coverage, leftmost-greedy overlap-deduped.
Drops matches the prefix leg already covers (``cur_st < prefix_tokens``),
then keeps a left-to-right non-overlapping subset -- two matches over the
same request range can't both scatter. Filtering precedes the dedup so a
prefix-covered match cannot suppress a usable one in the greedy pass.
Args:
matches: Candidate matches in any order; ``cur_st``/``cur_ed`` are
request token positions.
prefix_tokens: Contiguous prefix coverage in tokens; matches starting
before it are dropped. Pass ``0`` to keep all (dedup only).
Returns:
Non-overlapping matches in ascending ``cur_st`` order.
"""
kept: list[CBMatchResult] = []
covered_end = -1
for r in sorted(
(r for r in matches if r.cur_st >= prefix_tokens),
key=lambda r: r.cur_st,
):
if r.cur_st >= covered_end:
kept.append(r)
covered_end = r.cur_ed
return kept
def _resolve_cb_layout_desc(
self, model_name: str, world_size: int
) -> "MemoryLayoutDesc | None":
"""Find the CB KV buffer layout for ``(model_name, world_size)``.
Reads the thread-safe ``layout_desc_registry`` (populated by
``lmcache_driven_transfer`` on KV-cache registration) rather than
iterating ``cache_contexts``: iteration races concurrent
register/unregister, and the registry holds the complete multi-group
descriptor instead of a single-group manual reconstruction.
Args:
model_name (str): Model name to match.
world_size (int): Tensor-parallel world size to match.
Returns:
MemoryLayoutDesc | None: The matching layout, or None if no
registered CB context matches.
"""
return self._ctx.layout_desc_registry.find(model_name, world_size)
def _sparse_prefetch_submit(
self,
key: IPCCacheServerKey,
layout_desc: "MemoryLayoutDesc",
matches: list[CBMatchResult],
) -> "tuple[PrefetchHandle, dict[bytes, list], list[int]]":
"""Coalesce all matches into one sparse L2->L1 prefetch and submit it.
Non-blocking. Dedups object keys before submit (sparse keeps one read
lock per loaded key, so a duplicate would leak). The caller polls
``query_prefetch_status(handle)`` then calls :meth:`_sparse_classify`
with the found set.
Args:
key (IPCCacheServerKey): The request key.
layout_desc (MemoryLayoutDesc): CB KV buffer layout for L1 alloc.
matches (list[CBMatchResult]): Non-prefix matches to prefetch.
Returns:
tuple[PrefetchHandle, dict[bytes, list], list[int]]: the prefetch
handle, per-hash TP-expanded object keys, and each expanded
position's deduped-key index (maps the per-key found set back to
every chunk).
"""
world_size = key.world_size
per_hash_obj_keys: dict[bytes, list] = {}
all_hashes = [r.hash for r in matches]
all_obj_keys = ipc_key_to_object_keys(key, all_hashes, [0])[0]
for i, h in enumerate(all_hashes):
per_hash_obj_keys[h] = all_obj_keys[i * world_size : (i + 1) * world_size]
# Dedup keys before submit (sparse keeps one read lock per loaded key;
# a duplicate would leak). Map each expanded position to its deduped
# index so the per-key found set resolves back to every chunk.
uniq_keys: list = []
key_to_uidx: dict = {}
expanded_uidx: list[int] = []
for k in all_obj_keys:
uidx = key_to_uidx.get(k)
if uidx is None:
uidx = len(uniq_keys)
key_to_uidx[k] = uidx
uniq_keys.append(k)
expanded_uidx.append(uidx)
handle: PrefetchHandle = self._ctx.storage_manager.submit_prefetch_task(
uniq_keys,
layout_desc,
external_request_id=key.request_id,
policy=TrimPolicy.SPARSE,
)
return handle, per_hash_obj_keys, expanded_uidx
def _sparse_classify(
self,
key: IPCCacheServerKey,
matches: list[CBMatchResult],
found_uidx: set[int],
per_hash_obj_keys: dict[bytes, list],
expanded_uidx: list[int],
) -> list[CBMatchResult]:
"""Classify each prefetched chunk as found or stale, and finalize state.
A chunk is found only if every TP rank's key loaded; stale chunks take
an eviction strike (evicted at threshold, kept while still in-flight).
Stashes the found chunks' obj_keys for the retrieve path.
Args:
key (IPCCacheServerKey): The request key.
matches (list[CBMatchResult]): The submitted non-prefix matches.
found_uidx (set[int]): Deduped-key indices that loaded.
per_hash_obj_keys (dict[bytes, list]): Per-hash TP-expanded keys.
expanded_uidx (list[int]): Each expanded position's deduped index.
Returns:
list[CBMatchResult]: The found subset, in cur_st order.
"""
world_size = key.world_size
found_cb_match_result: list[CBMatchResult] = []
stale_hashes: list[bytes] = []
for j, r in enumerate(matches):
base = j * world_size
if all(expanded_uidx[base + t] in found_uidx for t in range(world_size)):
found_cb_match_result.append(r)
else:
stale_hashes.append(r.hash)
# Reset strikes for confirmed hashes.
if found_cb_match_result:
with self._pending_fp_lock:
for r in found_cb_match_result:
self._stale_strike.pop(r.hash, None)
# Stale: in-flight keep; >= threshold strikes -> evict.
if stale_hashes:
with self._pending_fp_lock:
truly_evict: list[bytes] = []
for h in stale_hashes:
if h in self._pending_fp_hashes:
continue
n = self._stale_strike.get(h, 0) + 1
if n >= self._STALE_STRIKE_THRESHOLD:
truly_evict.append(h)
self._stale_strike.pop(h, None)
else:
self._stale_strike[h] = n
if truly_evict:
self._token_range_matcher.remove_chunks(truly_evict)
self._event_bus.publish(
Event(
event_type=EventType.CB_CHUNKS_EVICTED,
metadata={"num_chunks": len(stale_hashes)},
)
)
# Stash per-hash obj_keys for retrieve (L2 opt).
if found_cb_match_result:
cache_entry = {
r.hash: per_hash_obj_keys[r.hash]
for r in found_cb_match_result
if r.hash in per_hash_obj_keys
}
with self._lookup_obj_keys_lock:
self._lookup_obj_keys_cache[key.request_id] = cache_entry
return found_cb_match_result
def _submit_prefix_leg(
self,
key: IPCCacheServerKey,
tp_size: int,
policy: TrimPolicy,
) -> "tuple[PrefetchHandle | None, int]":
"""Submit the CB prefix prefetch (non-blocking).
Opens the ``cb.prefix_lookup`` span (CB namespace — CB requests no longer
feed the MP request / mp.lookup_prefetch spans or the MP hit-rate
aggregate; the CB hit-rate metric carries prefix tokens via
CB_LOOKUP_END) and writes the shared session (``set_tokens`` +
``lookup_ipc_key``) so ``end_session``'s L1 keep-alive touch still
resolves the request's keys.
Args:
key (IPCCacheServerKey): Request key (token IDs, request_id, model,
world_size).
tp_size (int): Tensor-parallel size for MLA multi-reader locking.
policy (TrimPolicy): ``PREFIX`` or ``SEGMENTED_PREFIX``.
Returns:
tuple: ``(handle, world_size)``. ``handle`` is None when there is no
GPU context or no full chunk (the poll then reports 0 coverage).
"""
rid = key.request_id
model_name, world_size = key.model_name, key.world_size
self._event_bus.publish(
Event(event_type=EventType.CB_PREFIX_LOOKUP_START, session_id=rid)
)
layout_desc = self._resolve_cb_layout_desc(model_name, world_size)
if layout_desc is None:
logger.error(
"No CB GPU context for model %s ws %d during prefix lookup!",
model_name,
world_size,
)
return None, world_size
chunk_hashes = self._ctx.token_hasher.compute_chunk_hashes(list(key.token_ids))
if not chunk_hashes:
return None, world_size
# Lookup-hash logger (chunk hashes, for debug); guarded so the metadata
# dict is built only when a subscriber is listening.
if self._event_bus.has_subscribers(EventType.MP_LOOKUP):
self._event_bus.publish(
Event(
event_type=EventType.MP_LOOKUP,
session_id=rid,
metadata={
"request_id": rid,
"chunk_hashes": chunk_hashes,
"model_name": model_name,
"chunk_size": self._ctx.chunk_size,
"seq_len": len(key.token_ids),
"dtypes": [str(d) for d in layout_desc.dtypes],
"shapes": [list(s) for s in layout_desc.shapes],
},
)
)
# Shared session: end_session reads lookup_ipc_key + the rolling hashes
# to keep the request's KV alive in L1.
session = self._ctx.session_manager.get_or_create(rid)
session.set_tokens(list(key.token_ids))
session.lookup_ipc_key = key
extra_count = compute_extra_count(tp_size, world_size)
obj_keys = ipc_key_to_object_keys(key, chunk_hashes, [0])[0]
handle = self._ctx.storage_manager.submit_prefetch_task(
obj_keys,
layout_desc,
extra_count=extra_count,
external_request_id=rid,
policy=policy,
)
return handle, world_size
def _poll_prefix_leg(
self, job: "_CBUnifiedJob", rid: str, segmented: bool
) -> "tuple[int, list[int] | None] | None":
"""Poll the CB prefix handle; on completion close the cb.prefix_lookup span.
Consume-once: publishes CB_PREFIX_LOOKUP_END exactly once when the
prefetch lands. The prefix hit tokens are accounted on the CB hit-rate
metric at CB_LOOKUP_END, not here. For SEGMENTED_PREFIX also surfaces the
gapped retained chunk set.
Args:
job (_CBUnifiedJob): Poll state holding the prefix handle + world size.
rid (str): Request ID (event session_id).
segmented (bool): SEGMENTED_PREFIX active -> also surface the gapped
retained chunk set.
Returns:
tuple | None: ``(leading_chunks, retained_or_None)`` when resident;
``None`` while still loading. ``retained`` is the full gapped chunk
set for SEGMENTED_PREFIX, else None.
"""
if job.prefix_handle is not None:
bm = self._ctx.storage_manager.query_prefetch_status(job.prefix_handle)
if bm is None:
return None # still loading
ws = job.prefix_world_size
# NOTE(Kuntai): assumes uniform world size and prefix-ordered keys
# that break at the first miss.
leading = bm.count_leading_ones() // ws
retained = (
sorted({ki // ws for ki in bm.get_indices_list()})
if segmented
else None
)
else:
# No GPU context / no full chunk: nothing loaded.
leading, retained = 0, ([] if segmented else None)
self._event_bus.publish(
Event(
event_type=EventType.CB_PREFIX_LOOKUP_END,
session_id=rid,
metadata={"prefix_chunks": leading},
)
)
return leading, retained
def cb_unified_lookup(
self, key: IPCCacheServerKey, tp_size: int
) -> CBUnifiedLookupResult | None:
"""Non-blocking single-RPC CB lookup (submit-once, poll-on-recall).
First call submits the prefix lookup + fingerprint match; later calls
poll both legs, returning ``None`` until the prefix and the sparse
non-prefix complement are both resident in L1 (so a worker thread never
blocks on the L2->L1 loads). The prefix job's L1 read locks persist for
the retrieve.
Args:
key (IPCCacheServerKey): Request key (token IDs, request_id, model,
world_size).
tp_size (int): Tensor-parallel size for the prefix lookup.
Returns:
CBUnifiedLookupResult | None: ``None`` while either leg is still
loading (the caller re-issues to poll); on completion, the prefix
coverage in tokens plus the found non-prefix segments.
"""
rid = key.request_id
chunk_size = self._ctx.chunk_size
with self._cb_jobs_lock:
job = self._cb_jobs.get(rid)
if job is None:
# First invocation: start events + submit prefix + fingerprint match.
self._event_bus.publish(
Event(event_type=EventType.CB_REQUEST_START, session_id=rid)
)
self._event_bus.publish(
Event(
event_type=EventType.CB_LOOKUP_START,
session_id=rid,
metadata={"num_tokens": len(key.token_ids)},
)
)
# SEGMENTED_PREFIX: request the contiguous prefix with gap-tolerant
# retention so a mid-prefix L2 retrieve failure leaves the post-gap
# chunks L1-resident (picked up by the sparse leg as L1 hits, hole
# recomputed) instead of truncating the prefix at the gap.
prefix_policy = (
TrimPolicy.SEGMENTED_PREFIX
if self._segmented_prefix
else TrimPolicy.PREFIX
)
# Prefix leg: blend_v3 owns the submit + the cb.prefix_lookup span
# (under cb.lookup); prefix hit tokens land on the CB hit-rate
# metric via CB_LOOKUP_END below.
prefix_handle, prefix_ws = self._submit_prefix_leg(
key, tp_size, prefix_policy
)
# Local and coordinator matching are mutually exclusive: with a
# coordinator the fleet directory is the only source, so skip the
# local matcher (and its span). The coordinator leg is async
# (submitted below, resolved at poll) and is timed by cb.lookup.
matches: list[CBMatchResult]
if self._coordinator is not None:
matches = []
else:
# Local fingerprint match: CPU-bound, tight span.
self._event_bus.publish(
Event(
event_type=EventType.CB_FINGERPRINT_MATCH_START,
session_id=rid,
)
)
matches = self._match_fingerprints(key)
self._event_bus.publish(
Event(
event_type=EventType.CB_FINGERPRINT_MATCH_END,
session_id=rid,
metadata={"matches": len(matches)},
)
)
job = _CBUnifiedJob(
matches=matches,
num_tokens=len(key.token_ids),
prefix_handle=prefix_handle,
prefix_world_size=prefix_ws,
)
job.coord_submitted = self._submit_coordinator_match(key)
if job.coord_submitted and self._coordinator is not None:
job.coord_deadline = time.monotonic() + self._coordinator.match_budget_s
# Coordinator match leg: async span, ended on the resolving poll
# (or deadline) in _poll_coordinator_match.
self._event_bus.publish(
Event(
event_type=EventType.CB_COORDINATOR_MATCH_START,
session_id=rid,
)
)
with self._cb_jobs_lock:
self._cb_jobs[rid] = job
segmented = self._segmented_prefix
# --- Prefix leg: poll (consume-once) until the L1+L2 prefix lands. ---
if job.prefix_chunks is None:
res = self._poll_prefix_leg(job, rid, segmented)
if res is None:
return None # prefix still loading -> defer
job.prefix_chunks, prefix_retained = res
if segmented:
job.retained_chunks = prefix_retained
# Poll above set it (or we returned); narrow for the arithmetic below.
assert job.prefix_chunks is not None
prefix_chunks: int = job.prefix_chunks
# Prefix done: reconcile the complement outside the prefix coverage and
# submit one sparse prefetch for it (once). Prefix-covered chunks never
# enter the sparse prefetch, so they cannot leak a read lock.
if not job.sparse_started:
prefix_tokens = prefix_chunks * chunk_size
if self._coordinator is not None:
candidates = self._poll_coordinator_match(job, rid)
if candidates is None:
return None # coordinator still in flight (bounded by deadline)
else:
candidates = job.matches
# Under SEGMENTED_PREFIX, a same-position match the prefix leg already
# retained rides the segmented tail (prefix-class: pure load, no CHECK)
# -- drop it here so it is not scattered twice. A same-position match
# the tail does NOT cover is a genuine cross-context hit: keep it as
# non-prefix (re-RoPE no-ops at delta 0, but it still needs CHECK).
# Shifted (cur != old) matches are always kept. Then the single
# prefix-filter + overlap dedup over the rest.
if segmented:
retained = set(job.retained_chunks or [])
candidates = [
c
for c in candidates
if c.old_st != c.cur_st or (c.cur_st // chunk_size) not in retained
]
job.non_prefix = self._non_overlapping_after_prefix(
candidates, prefix_tokens
)
if job.non_prefix:
layout_desc = self._resolve_cb_layout_desc(
key.model_name, key.world_size
)
if layout_desc is not None:
(
job.handle,
job.per_hash_obj_keys,
job.expanded_uidx,
) = self._sparse_prefetch_submit(key, layout_desc, job.non_prefix)
# Only trace the span when the prefetch actually reads L2;
# all-L1-resident matches do no L2 work worth a span.
job.l2_keys = len(job.handle.l2_orig_indices)
if job.l2_keys > 0:
self._event_bus.publish(
Event(
event_type=EventType.CB_SPARSE_PREFETCH_START,
session_id=rid,
metadata={
"n_chunks": len(job.non_prefix),
"world_size": key.world_size,
"n_keys": len(job.non_prefix) * key.world_size,
"l2_keys": job.l2_keys,
},
)
)
else:
logger.error(
"No CB GPU context for model %s ws %d during cb_unified_lookup",
key.model_name,
key.world_size,
)
job.non_prefix = []
job.sparse_started = True
# --- Sparse leg: poll (consume-once) until the scattered chunks land. ---
if job.handle is not None and job.found_uidx is None:
bm = self._ctx.storage_manager.query_prefetch_status(job.handle)
if bm is None:
return None # sparse still loading -> defer
job.found_uidx = set(bm.get_indices_list())
if job.l2_keys > 0:
self._event_bus.publish(
Event(
event_type=EventType.CB_SPARSE_PREFETCH_END,
session_id=rid,
metadata={
"found_keys": len(job.found_uidx),
"l2_keys": job.l2_keys,
},
)
)
# --- BOTH legs ready: classify the complement + finalize. ---
if job.handle is not None:
found = self._sparse_classify(
key,
job.non_prefix or [],
job.found_uidx or set(),
job.per_hash_obj_keys or {},
job.expanded_uidx or [],
)
else:
found = []
prefix_tokens = prefix_chunks * chunk_size
num_tokens = job.num_tokens
# Segmented tail: post-gap chunks the SEGMENTED_PREFIX prefix leg kept
# resident (retained index > the leading run). Delivered at their
# original positions (old_st == cur_st) so the connector tags them
# ``prefix`` (pure load, no recompute); only the gap is recomputed. The
# storage key is the same prefix-chained chunk hash the prefix leg used,
# so no fingerprint match is needed to retrieve them.
segmented_tail: list[CBMatchResult] = []
if segmented and job.retained_chunks:
chunk_hashes = self._ctx.token_hasher.compute_chunk_hashes(
list(key.token_ids)
)
for i in job.retained_chunks:
if i < prefix_chunks or i >= len(chunk_hashes):
continue # leading run (already prefix) / sub-chunk tail
st = i * chunk_size
segmented_tail.append(
CBMatchResult(
old_st=st,
old_ed=st + chunk_size,
cur_st=st,
cur_ed=st + chunk_size,
hash=chunk_hashes[i],
)
)
seg_tail_tokens = _unique_token_coverage(segmented_tail)
non_prefix_hit_tokens = _unique_token_coverage(found)
self._event_bus.publish(
Event(
event_type=EventType.CB_LOOKUP_END,
session_id=rid,
metadata={
"num_tokens": num_tokens,
"fingerprint_hits": len(found),
"prefix_hits": job.prefix_chunks,
"prefix_chunks": job.prefix_chunks,
"storage_hits": len(found),
"stale_chunks": len(job.non_prefix or []) - len(found),
"no_gpu_context": False,
"prefix_hit_tokens": prefix_tokens,
"segmented_prefix_hit_tokens": seg_tail_tokens,
"non_prefix_hit_tokens": non_prefix_hit_tokens,
"hit_tokens": prefix_tokens
+ _unique_token_coverage(found + segmented_tail),
"requested_tokens": (num_tokens // chunk_size) * chunk_size,
},
)
)
with self._cb_jobs_lock:
self._cb_jobs.pop(rid, None)
return CBUnifiedLookupResult(
prefix_coverage_tokens=prefix_tokens,
non_prefix_segments=found,
segmented_prefix_segments=segmented_tail,
)
def store(
self,
key: IPCCacheServerKey,
instance_id: int,
gpu_block_ids: list[list[int]],
event_ipc_handle: bytes,
) -> tuple[bytes, bool]:
"""Paged store, then register the stored chunks as match fingerprints.
Delegates the KV write to ``LMCacheDrivenTransfer.store``, then (worker 0 only)
enqueues the chunk hashes for async fingerprint registration ordered
after the L1 commit. Chunk 0 of a position-0 store is skipped (owned by
the standard prefix path). Fingerprint failures are logged, never
raised — they do not affect store correctness.
Args:
key (IPCCacheServerKey): Store key (token IDs + ``[start, end)``).
instance_id (int): Target KV-cache instance.
gpu_block_ids (list[list[int]]): Per-layer-group paged block IDs.
event_ipc_handle (bytes): IPC handle to the producer's CUDA event.
Returns:
tuple[bytes, bool]: The underlying ``LMCacheDrivenTransfer.store`` result
(event handle, success).
"""
result = self._transfer_module.store(
key, instance_id, gpu_block_ids, event_ipc_handle
)
# The matcher is engine-shared; only worker 0 registers.
if key.worker_id not in (0, None):
return result
# Enqueue on cupy_stream so CUDA FIFO ordering puts registration
# after the L1-commit callback; otherwise lookups see the chunk as
# not-yet-committed and drop the whole group as stale.
try:
session = self._ctx.session_manager.get_or_create(key.request_id)
chunk_hashes = [
TokenHasher.hash_to_bytes(h)
for h in session.get_hashes(key.start, key.end)
]
if not chunk_hashes:
return result
tokens_in_range = list(key.token_ids)[key.start : key.end]
start_chunk_idx = 1 if key.start == 0 else 0
job = (tokens_in_range, chunk_hashes, start_chunk_idx, key.start)
with self._pending_fp_lock:
self._pending_fp_hashes.update(chunk_hashes[start_chunk_idx:])
entry = self._transfer_module.get_and_touch_context_entry(instance_id)
gpu_ctx = entry.cache_context if entry is not None else None
if gpu_ctx is not None and gpu_ctx.cupy_stream is not None:
gpu_ctx.cupy_stream.launch_host_func(
self._fingerprint_queue.put_nowait, job
)
else:
self._fingerprint_queue.put_nowait(job)
except Exception:
logger.exception(
"CB fingerprint enqueue failed for request %s "
"(does not affect store correctness)",
key.request_id,
)
if self._coordinator is not None:
self._publish_fingerprints(key, chunk_hashes, tokens_in_range)
return result
def _publish_fingerprints(
self,
key: IPCCacheServerKey,
chunk_hashes: list[bytes],
tokens_in_range: list[int],
) -> None:
"""Publish this stored range's chunk fingerprints to the coordinator.
Best-effort and fire-and-forget (enqueue only): one wire
``ChunkFingerprint`` per stored chunk -- its content poly-hash (the same
``chunk_hash_windows_numba`` the match probes, with the fleet base), its
shared-L2 ``object_key`` (the chunk storage key ``th``), and its token
position. Never raises into the store path.
Args:
key: The store request key (model/scope/positions).
chunk_hashes: Per-chunk storage keys (``th``) for the range.
tokens_in_range: The stored tokens ``token_ids[start:end]``.
"""
coordinator = self._coordinator
if coordinator is None or not chunk_hashes:
return
try:
model_scope = key.model_name
store_range = {
"model_scope": model_scope,
"tokens": list(tokens_in_range),
"object_keys": [h.hex() for h in chunk_hashes],
"old_st_base": key.start,
}
coordinator.enqueue_register([store_range])
except Exception:
logger.warning(
"CB coordinator publish build failed for request %s "
"(does not affect store correctness)",
key.request_id,
)
def _submit_coordinator_match(self, key: IPCCacheServerKey) -> bool:
"""Issue a fleet directory match query for this request (best-effort).
Args:
key: The lookup request key.
Returns:
``True`` if a query was submitted (so the finalize step should poll
for it), ``False`` when there is no coordinator or submission failed.
"""
coordinator = self._coordinator
if coordinator is None:
return False
try:
tokens = list(key.token_ids)
if len(tokens) < self._ctx.chunk_size:
return False
coordinator.submit_match(key.request_id, key.model_name, tokens)
return True
except Exception:
logger.warning(
"CB coordinator match submit failed for request %s", key.request_id
)
return False
def _poll_coordinator_match(
self, job: "_CBUnifiedJob", rid: str
) -> "list[CBMatchResult] | None":
"""Poll the coordinator match result, deferring until it resolves.
Mirrors the prefix/sparse legs: ``return None`` to defer while pending.
A per-lookup wall-clock deadline (``job.coord_deadline``) bounds the
total wait, including queue/pool time. Past the deadline the leg is
abandoned and the lookup proceeds local-only (the client's later fill,
if any, is dropped via ``take_match``).
Args:
job: The per-request poll state.
rid: Request id.
Returns:
The global segments (possibly empty) once resolved or timed out, or
``None`` to defer (still in flight and within the deadline).
"""
coordinator = self._coordinator
if coordinator is None or not job.coord_submitted:
return []
poll = coordinator.poll_match(rid)
if poll is PENDING:
if time.monotonic() < job.coord_deadline:
return None # defer; bounded by job.coord_deadline
coordinator.take_match(rid)
logger.warning(
"CB coordinator match deadline exceeded for %s; local-only", rid
)
self._event_bus.publish(
Event(
event_type=EventType.CB_COORDINATOR_MATCH_END,
session_id=rid,
metadata={"matches": 0, "timed_out": True},
)
)
return []
coordinator.take_match(rid)
segments = self._build_global_segments(poll) if isinstance(poll, list) else []
self._event_bus.publish(
Event(
event_type=EventType.CB_COORDINATOR_MATCH_END,
session_id=rid,
metadata={"matches": len(segments), "timed_out": False},
)
)
return segments
def _build_global_segments(
self, matches: "list[RemoteMatch]"
) -> list[CBMatchResult]:
"""Convert coordinator matches into chunk-granular retrievable segments.
Each coordinator ``object_key`` is the hex of the chunk's content hash
(the same ``th`` a local ``CBMatchResult.hash`` holds), so the matches
are returned as ``CBMatchResult`` directly: the retrieve path then
expands ``hash`` to per-rank shared-L2 object keys via
``ipc_key_to_object_keys``, identical to local matches.
Args:
matches: Matched chunks returned by the coordinator client.
Returns:
One :class:`CBMatchResult` per matched chunk (request order).
"""
chunk_size = self._ctx.chunk_size
return [
CBMatchResult(
old_st=m.old_st,
old_ed=m.old_st + chunk_size,
cur_st=m.cur_st,
cur_ed=m.cur_st + chunk_size,
hash=bytes.fromhex(m.object_key),
)
for m in matches
]
def _drain_fingerprint_queue(self) -> None:
"""Best-effort background drainer for _fingerprint_queue."""
while not self._fingerprint_stop.is_set():
try:
job = self._fingerprint_queue.get(timeout=0.1)
except QueueEmpty:
continue
tokens_in_range, chunk_hashes, start_chunk_idx, position_offset = job
try:
self._token_range_matcher.on_new_token_hashes(
tokens_in_range,
chunk_hashes,
start_chunk_idx=start_chunk_idx,
position_offset=position_offset,
)
except Exception:
logger.exception("CB fingerprint registration failed (async)")
finally:
with self._pending_fp_lock:
self._pending_fp_hashes.difference_update(
chunk_hashes[start_chunk_idx:]
)
def _apply_cb_rope_batched(
self,
gpu_context: BaseCacheContext,
rope_state: _CBRopeState,
batch_len: int,
slots_to_rope: list[tuple[int, int, int]],
) -> None:
"""Re-RoPE the given tmp-pool slots in place (K-only, per kernel group).
Args:
gpu_context (GPUCacheContext): The instance's GPU cache context.
rope_state (_CBRopeState): Cached cos/sin + head layout.
batch_len (int): Number of tmp slots staged for this batch.
slots_to_rope (list[tuple[int, int, int]]): ``(slot_idx, old_st,
cur_st)`` per shifted slot — re-RoPE K from stored position
``old_st`` to new position ``cur_st``.
Raises:
RuntimeError: On a compressed (compress_ratio != 1) or MLA
(kv_size != 2) layout, or a head_size/hidden_dim mismatch.
"""
if not slots_to_rope:
return
num_groups = gpu_context.kv_layer_groups_manager.num_kernel_groups
for group_idx in range(num_groups):
group = gpu_context.kv_layer_groups_manager.kernel_groups[group_idx]
if group.tokens_per_block != group.slots_per_block:
raise RuntimeError(
f"CB v3: group {group_idx} is compressed "
f"(tokens_per_block={group.tokens_per_block}, "
f"slots_per_block={group.slots_per_block}); "
f"compressed layouts unsupported."
)
all_slots = [
gpu_context.get_temp_kernel_group_buffer(slot_idx, group_idx)
for slot_idx in range(batch_len)
]
if all_slots[0].shape[0] != 2:
raise RuntimeError(
f"CB v3: group {group_idx} has kv_size={all_slots[0].shape[0]}; "
"MLA layouts unsupported."
)
num_layers, slots, hidden_dim = all_slots[0].shape[1:]
n_heads = hidden_dim // rope_state.head_size
if n_heads * rope_state.head_size != hidden_dim:
raise RuntimeError(
f"CB rope: group {group_idx} hidden_dim ({hidden_dim}) "
f"not a multiple of head_size ({rope_state.head_size})."
)
# Per-group rope cache: dual-RoPE models rotate each
# kernel group with its own theta's cos/sin.
group_cos_sin = rope_state.cache_for_group(group.engine_group_idx)
# slot ramp tiled across layers is invariant per (num_layers,
# slots) — cache it; each shifted slot then just adds its offset.
device = all_slots[0].device
sp_key = (str(device), num_layers, slots)
sp_cache = getattr(self, "_cb_sp_rep_cache", None)
if sp_cache is None:
sp_cache = {}
self._cb_sp_rep_cache = sp_cache
slot_positions_rep = sp_cache.get(sp_key)
if slot_positions_rep is None:
slot_positions_rep = torch.arange(
slots, device=device, dtype=torch.long
).repeat(num_layers)
sp_cache[sp_key] = slot_positions_rep
for slot_idx, old_st, cur_st in slots_to_rope:
# reshape returns an in-place view (tmp slots are contiguous).
k_view = all_slots[slot_idx][0].reshape(
num_layers * slots, n_heads, rope_state.head_size
)
lmc_ops.rotary_embedding_k_fused(
old_st + slot_positions_rep,
cur_st + slot_positions_rep,
k_view,
rope_state.head_size,
group_cos_sin,
rope_state.is_neox_style,
)
def cb_retrieve_pre_computed(
self,
key: IPCCacheServerKey,
cb_match_result: list[CBMatchResult],
gpu_block_ids: list[list[int]],
instance_id: int,
event_ipc_handle: bytes,
) -> tuple[bytes, bool]:
"""Scatter every matched token range into the request's paged KV.
Reuses the lookup's prefetched chunks: fills tmp slots, K-only re-RoPEs
the shifted (non-prefix) subset, then writes per-token via the slot
kernel — so non-block-aligned matches and partial vLLM blocks shared
with recomputed tokens are written correctly (no block-alignment trim).
Only matches past the currently allocated slots are dropped (vLLM may
call this twice: partial- then full-block alloc).
Args:
key (IPCCacheServerKey): The request key.
cb_match_result (list[CBMatchResult]): Matched ranges to scatter
(prefix-hit and shifted), any order.
gpu_block_ids (list[list[int]]): This request's paged block table
per engine (kernel) group; single-group models pass [[...]].
Mirrors the engine RETRIEVE/STORE per-group block-id contract.
instance_id (int): Target KV-cache instance.
event_ipc_handle (bytes): IPC handle to the forward's CUDA event.
Returns:
tuple[bytes, bool]: The scatter-complete event handle and whether
the scatter ran (False if the prefetched objects were unavailable).
Raises:
ValueError: If the instance has no registered KV cache or rope
state. MLA layouts are unsupported (raised during re-RoPE).
"""
entry = self._transfer_module.get_and_touch_context_entry(instance_id)
if entry is None:
raise ValueError(
f"Instance {instance_id} not registered for paged KV cache"
)
if instance_id not in self._cb_rope_state:
raise ValueError(
f"Instance {instance_id} has no CB rope state; "
"send CB_REGISTER_ROPE_V3 before CB_RETRIEVE_PRE_COMPUTED_V3."
)
gpu_context = entry.cache_context
rope_state = self._cb_rope_state[instance_id]
chunk_size = self._ctx.chunk_size
_retrieve_t0 = time.perf_counter()
cb_match_result = sorted(cb_match_result, key=lambda r: r.cur_st)
# L2 opt: reuse lookup's obj_keys cache; fall back to re-resolve.
with self._lookup_obj_keys_lock:
cached = self._lookup_obj_keys_cache.pop(key.request_id, None)
if cached is not None and all(r.hash in cached for r in cb_match_result):
# The lookup cached all-ranks obj keys (world_size per hash). This
# retrieve is per-worker, so select THIS rank's key -> M objects, not
# M*world_size (else the zip below silently truncates and mispairs
# ranks at TP>1). Mirrors the non-cached path's per-worker resolve.
if key.worker_id is not None and key.world_size > 1:
all_obj_keys = [cached[r.hash][key.worker_id] for r in cb_match_result]
else:
all_obj_keys = [k for r in cb_match_result for k in cached[r.hash]]
else:
all_obj_keys = ipc_key_to_object_keys(
key, [r.hash for r in cb_match_result], [0]
)[0]
# Lookup read-locked the full found set, but the connector may have
# dropped some matches (parent-covered / misaligned) before retrieve,
# leaking their per-key read locks. Release those orphans now (disjoint
# from all_obj_keys, which retrieve still consumes; needs the key cache).
if cached is not None:
retrieved_hashes = {r.hash for r in cb_match_result}
orphan_keys = [
k for h, ks in cached.items() if h not in retrieved_hashes for k in ks
]
if orphan_keys:
self._ctx.storage_manager.finish_read_prefetched(orphan_keys)
logger.debug(
"CB V3 released %d prefetched-but-unretrieved keys (req=%s)",
len(orphan_keys),
key.request_id,
)
# Non-prefix sparse hits split by re-rope need (not prefix coverage).
n_non_shifted = sum(1 for r in cb_match_result if r.old_st == r.cur_st)
n_shifted = len(cb_match_result) - n_non_shifted
if not all_obj_keys:
self._event_bus.publish(
Event(
event_type=EventType.CB_REQUEST_END,
session_id=key.request_id,
)
)
return event_ipc_handle, True
logger.debug("CB V3 retrieving object keys: %s", all_obj_keys)
# CB v3 only supports uncompressed single-block-id-space layouts
# (enforced per group in ``_apply_cb_rope_batched``), so the first
# kernel group's chunk geometry is representative.
tokens_per_block = gpu_context.kv_layer_groups_manager.kernel_groups[
0
].tokens_per_block
if chunk_size % tokens_per_block != 0:
raise ValueError(
f"chunk_size {chunk_size} must be a multiple of "
f"tokens_per_block {tokens_per_block}"
)
num_groups = gpu_context.kv_layer_groups_manager.num_kernel_groups
with (
torch_dev.device(gpu_context.device),
torch_dev.stream(gpu_context.stream),
):
check_interprocess_event_support()
event = torch_dev.Event(interprocess=True)
# One staged block-id tensor per engine group, indexed by group.
block_ids_per_group_gpu = gpu_context.stage_block_ids(gpu_block_ids)
# Resolve each kernel group's block table + block size once. Select
# by engine_group_idx (kernel groups may share one, e.g. MiniMax-M3).
kgm = gpu_context.kv_layer_groups_manager
resolved_groups: list[tuple[torch.Tensor, int]] = []
for group_idx in range(num_groups):
eg_idx = kgm.kernel_groups[group_idx].engine_group_idx
if eg_idx >= len(block_ids_per_group_gpu):
# Engine groups have independent block tables under HMA;
# substituting another group's table would scatter KV into
# the wrong physical blocks (silent corruption).
raise ValueError(
f"CB retrieve: kernel group {group_idx} maps to engine "
f"group {eg_idx}, but only "
f"{len(block_ids_per_group_gpu)} block table(s) were "
"provided."
)
resolved_groups.append(
(
block_ids_per_group_gpu[eg_idx],
kgm.kernel_groups[group_idx].tokens_per_block,
)
)
self._event_bus.publish_on_stream(
gpu_context.cupy_stream,
Event(
event_type=EventType.CB_RETRIEVE_START,
session_id=key.request_id,
metadata={
"num_chunks": len(cb_match_result),
"model_name": key.model_name,
},
),
)
if not hasattr(torch_dev.Event, "from_ipc_handle"):
raise RuntimeError(
f"Backend '{torch_device_type}' does not support IPC "
"event handles (Event.from_ipc_handle not available). "
"Multiprocess IPC requires CUDA."
)
vllm_event = torch_dev.Event.from_ipc_handle(
gpu_context.device, event_ipc_handle
)
vllm_event.wait(stream=gpu_context.stream)
try:
with self._ctx.storage_manager.read_prefetched_results(
all_obj_keys
) as memory_objs:
if memory_objs is None:
return event_ipc_handle, False
# Per-token scatter handles any cur_st; just bound the
# matched range to the allocated slots.
pairs: list[tuple[CBMatchResult, Any]] = []
# Bound by the smallest group: under HMA the sliding group
# has fewer blocks than the full group, so [0] isn't safe.
num_slots = min(
int(block_ids.numel()) * group_bs
for block_ids, group_bs in resolved_groups
)
for r, memory_obj in zip(cb_match_result, memory_objs, strict=True):
if r.cur_ed > num_slots:
logger.warning(
"Dropping CB match cur_st=%d cur_ed=%d: exceeds "
"%d slots. Request %s.",
r.cur_st,
r.cur_ed,
num_slots,
key.request_id,
)
continue
pairs.append((r, memory_obj))
# cb.scatter span (GPU): the L1->paged write of every
# applied match. Re-RoPE is folded in (n_shifted) — it is
# interleaved per-batch, so not a separate span.
self._event_bus.publish_on_stream(
gpu_context.cupy_stream,
Event(
event_type=EventType.CB_SCATTER_START,
session_id=key.request_id,
metadata={
"scattered_tokens": sum(
r.cur_ed - r.cur_st for r, _ in pairs
),
"n_prefix": sum(
1 for r, _ in pairs if r.old_st == r.cur_st
),
"n_shifted": sum(
1 for r, _ in pairs if r.old_st != r.cur_st
),
"dropped": len(cb_match_result) - len(pairs),
},
),
)
# Consecutive matches → one batched scatter per group.
runs: list[list[tuple[CBMatchResult, Any]]] = []
for r_obj in pairs:
r = r_obj[0]
if runs and runs[-1][-1][0].cur_ed == r.cur_st:
runs[-1].append(r_obj)
else:
runs.append([r_obj])
max_batch = gpu_context.max_batch_size
for run in runs:
for batch_start in range(0, len(run), max_batch):
batch = run[batch_start : batch_start + max_batch]
batch_len = len(batch)
# (a) H2D fill into per-chunk tmp slots.
for slot_idx, (_, memory_obj) in enumerate(batch):
# Single object group => object_group_idx=0.
flat_slot = gpu_context.get_temp_object_group_buffer(
slot_idx, 0
)
lmcache_memcpy_async_h2d(memory_obj, flat_slot)
# (b) Re-RoPE shifted (non-prefix) slots in place.
slots_to_rope = [
(slot_idx, r.old_st, r.cur_st)
for slot_idx, (r, _) in enumerate(batch)
if r.old_st != r.cur_st
]
self._apply_cb_rope_batched(
gpu_context, rope_state, batch_len, slots_to_rope
)
# (c) Per-token slot scatter: partial vLLM blocks
# shared with recomputed tokens stay disjoint.
pos = torch.cat(
[
torch.arange(
r.cur_st,
r.cur_ed,
device=gpu_context.device,
dtype=torch.long,
)
for (r, _) in batch
]
)
for group_idx in range(num_groups):
# This group's block table + size (resolved above).
group_block_ids, group_bs = resolved_groups[group_idx]
# Per-group block count: under HMA the sliding
# group has fewer blocks than the full group, so
# gpu_context.num_blocks (group 0's) would
# truncate the other groups' bounds check.
page_buffer_size = (
kgm.kernel_groups[group_idx].shape_desc.nb
* group_bs
)
slot_mapping = group_block_ids[
pos // group_bs
] * group_bs + (pos % group_bs)
tmp_buffers = [
gpu_context.get_temp_kernel_group_buffer(
slot_idx, group_idx
)
for slot_idx in range(batch_len)
]
key_value = torch.cat(tmp_buffers, dim=2)
lmc_ops.multi_layer_kv_transfer(
key_value,
gpu_context.get_kernel_group_kv_pointers(group_idx),
slot_mapping,
gpu_context.device,
page_buffer_size,
lmc_ops.TransferDirection.H2D,
gpu_context.get_engine_kv_format(group_idx),
block_size=group_bs,
head_size=rope_state.head_size,
)
self._event_bus.publish_on_stream(
gpu_context.cupy_stream,
Event(
event_type=EventType.CB_SCATTER_END,
session_id=key.request_id,
),
)
except Exception:
logger.exception("Error during retrieving prefetched results")
self._event_bus.publish_on_stream(
gpu_context.cupy_stream,
Event(
event_type=EventType.CB_RETRIEVE_END,
session_id=key.request_id,
metadata={"success": False},
),
)
self._event_bus.publish_on_stream(
gpu_context.cupy_stream,
Event(
event_type=EventType.CB_REQUEST_END,
session_id=key.request_id,
),
)
return event_ipc_handle, False
event.record()
self._event_bus.publish_on_stream(
gpu_context.cupy_stream,
Event(
event_type=EventType.CB_RETRIEVE_END,
session_id=key.request_id,
metadata={"success": True},
),
)
_scatter_ms = (time.perf_counter() - _retrieve_t0) * 1000
logger.info(
"Retrieved pre-computed for %d match results into request %s "
"paged blocks (scatter_ms=%.2f, non_shifted=%d shifted=%d)",
len(cb_match_result),
key.request_id,
_scatter_ms,
n_non_shifted,
n_shifted,
)
self._event_bus.publish_on_stream(
gpu_context.cupy_stream,
Event(
event_type=EventType.CB_REQUEST_END,
session_id=key.request_id,
),
)
return event.ipc_handle(), True