Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 21 additions & 0 deletions .lil/changes/lmcache-restore-admission.json
Original file line number Diff line number Diff line change
@@ -0,0 +1,21 @@
{
"schema": "local-inference-release-change/v1",
"id": "lmcache-restore-admission",
"category": "fix",
"summary": "Recurrent checkpoint restores wait for GPU and copy capacity instead of admitting requests to recompute cached prefixes.",
"models": [
"all"
],
"compatibility": "Requires the paired vLLM restore-admission change; deploy vLLM before enabling this LMCache change.",
"details": [
"Capacity waits retain the answered checkpoint lookup and do not consume the timeout for an unanswered directory request.",
"A local checkpoint evicted before admission triggers a fresh external lookup, while genuine misses and exhausted transfer retries retain recomputation fallback.",
"This change does not repair missing checkpoint pages in cache storage."
],
"requires": [
"vllm-restore-admission"
],
"pull_requests": [
100
]
}
106 changes: 79 additions & 27 deletions lmcache/integration/vllm/checkpoint_scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,9 +76,11 @@ class _Lookup:
done: bool = False
task_id: str | None = None
checkpoint_id: int | None = None
local_checkpoint_id: int | None = None
attempts: int = 1
started: float = field(default_factory=time.monotonic)
capacity_refusal_logged: bool = False
task_refusal_logged: bool = False


class CheckpointSchedulerBridge:
Expand All @@ -91,10 +93,10 @@ class CheckpointSchedulerBridge:
Worker layout is added after all ranks report identical descriptors.
world_size: Number of engine ranks contributing to one atomic generation.
max_tasks: Admission limit for collective stores and restores.
lookup_timeout: Seconds a request may wait for directory replies,
including shorter-checkpoint retries, and a store for its begin
reply. A request whose lookup has not been answered by then is
admitted to recompute its prompt; an unanswered store is aborted.
lookup_timeout: Seconds allowed for each directory reply (including
retries) and for a store's begin reply. Capacity waiting does not
consume this deadline. An unanswered lookup permits prompt
recomputation; an unanswered store is aborted.
A started worker copy is never abandoned.

The scheduler must keep issuing connector-only steps while has_pending is
Expand Down Expand Up @@ -175,9 +177,8 @@ def poll_prefix(self, request: "Request") -> bool:

Returns:
False while layout negotiation, lookup or H2D is pending. True when
ordinary GPU admission may proceed, including when an external
import cannot reserve capacity. An unadmitted request may poll
again to retry that reservation without another directory lookup.
ordinary GPU admission may proceed. Capacity and copy-task pressure
keep an available checkpoint waiting without another lookup.
Reusing a finished request's ID waits for its admitted copies to
drain, so their completion cannot cancel or erase another lookup.
"""
Expand All @@ -188,11 +189,15 @@ def poll_prefix(self, request: "Request") -> bool:
):
return True
local = self._cache.find(request, request.num_tokens)
if local is not None and local.num_tokens == request.num_tokens:
state = self._lookups.get(request.request_id)
if (
state is None
and local is not None
and local.num_tokens == request.num_tokens
):
return True
if self._layout is None:
return False
state = self._lookups.get(request.request_id)
if state is None:
roots = self._roots(request)
state = _Lookup(
Expand All @@ -202,6 +207,27 @@ def poll_prefix(self, request: "Request") -> bool:
self._lookups[request.request_id] = state
return False
if state.done:
# Local reuse can also disappear while ordinary admission waits.
selected_id = (
state.local_checkpoint_id
if state.local_checkpoint_id is not None
else state.checkpoint_id
)
if (
selected_id is not None
and not (
self._manager.external_boundary_admission_ready(request.request_id)
)
and (local is None or local.checkpoint_id != selected_id)
):
logger.info(
"Recurrent checkpoint invalidated before admission for "
"request %s; looking up its prefix again",
request.request_id,
)
self._manager.release_external_boundary_admission(request.request_id)
del self._lookups[request.request_id]
return self.poll_prefix(request)
return True
if state.task_id is not None:
return False
Expand All @@ -227,21 +253,25 @@ def poll_prefix(self, request: "Request") -> bool:
)
state.done = True
return True
if manifest is None or (
local is not None and manifest.prefix.num_tokens <= local.num_tokens
):
if manifest is None:
state.done = True
return True
if len(self._tasks) >= self._max_tasks:
logger.info(
"Recurrent checkpoint restore of %d tokens skipped for request %s: "
"%d checkpoint copies in flight; recomputing its prompt",
manifest.prefix.num_tokens,
request.request_id,
len(self._tasks),
)
if local is not None and manifest.prefix.num_tokens <= local.num_tokens:
# Keep local selection separate from external cache-hit accounting.
state.local_checkpoint_id = local.checkpoint_id
state.done = True
return True
if len(self._tasks) >= self._max_tasks:
if not state.task_refusal_logged:
state.task_refusal_logged = True
logger.info(
"Recurrent checkpoint restore of %d tokens waiting for "
"copy capacity for request %s (%d copies in flight)",
manifest.prefix.num_tokens,
request.request_id,
len(self._tasks),
)
return False
try:
payload, positions = self._validate_manifest(manifest, state.roots)
checkpoint = self._manager.reserve_external_boundary_checkpoint(
Expand All @@ -251,6 +281,7 @@ def poll_prefix(self, request: "Request") -> bool:
draft_prefix_len=payload["draft_prefix_len"],
kind=payload["kind"],
num_ranks=self._world_size,
reserve_admission=True,
)
except (ValueError, KeyError, TypeError):
logger.warning(
Expand All @@ -263,21 +294,25 @@ def poll_prefix(self, request: "Request") -> bool:
state.done = True
return True
if checkpoint is None:
# Insufficient GPU capacity is not an external-cache miss. Permit
# ordinary admission, but retain the manifest so an unadmitted
# request can retry its import after other owners release pages.
# Resource pressure is not a cache miss. Keep the answered manifest
# while runnable owners release capacity.
if not state.capacity_refusal_logged:
state.capacity_refusal_logged = True
logger.info(
"Recurrent checkpoint restore of %d tokens deferred for "
"request %s: not enough free GPU blocks (%d free); it is "
"admitted without the restore unless capacity returns first",
"request %s: waiting for restore and execution capacity "
"(%d free GPU blocks)",
manifest.prefix.num_tokens,
request.request_id,
self._manager.block_pool.get_num_free_blocks(),
)
return True
return False
task = self._make_task(manifest, checkpoint, "RETRIEVE")
logger.debug(
"Recurrent checkpoint restore of %d tokens copying for request %s",
manifest.prefix.num_tokens,
request.request_id,
)
self._tasks[task.task_id] = _PendingTask(task, checkpoint, request.request_id)
state.task_id = task.task_id
return False
Expand Down Expand Up @@ -523,6 +558,22 @@ def complete(self, results: dict[str, dict[int, bool]]) -> None:
)
if published and state is not None:
state.checkpoint_id = pending.checkpoint.checkpoint_id
logger.debug(
"Recurrent checkpoint restore of %d tokens ready for "
"request %s; ownership retained until admission",
pending.checkpoint.num_tokens,
pending.request_id,
)
elif state is not None:
retry = state.attempts < _MAX_LOOKUP_ATTEMPTS
logger.info(
"Recurrent checkpoint publication invalidated for "
"request %s%s",
pending.request_id,
"; retrying lookup"
if retry
else "; recomputing its prompt",
)
elapsed = time.monotonic() - pending.created
if elapsed > _SLOW_RESTORE_SECONDS:
logger.info(
Expand All @@ -540,7 +591,6 @@ def complete(self, results: dict[str, dict[int, bool]]) -> None:
state is not None
and pending.request_id not in self._cancelled
and state.attempts < _MAX_LOOKUP_ATTEMPTS
and time.monotonic() - state.started < self._lookup_timeout
)
logger.info(
"Recurrent checkpoint restore of %d tokens failed for "
Expand All @@ -567,6 +617,7 @@ def complete(self, results: dict[str, dict[int, bool]]) -> None:
)
state.task_id = None
state.attempts += 1
state.started = time.monotonic()
state.future = self._client.submit_request(
RequestType.CHECKPOINT_FIND, [state.roots.roots]
)
Expand All @@ -578,6 +629,7 @@ def complete(self, results: dict[str, dict[int, bool]]) -> None:

def finish_request(self, request_id: str) -> None:
"""Forget lookup state, retaining copy pins until every admitted task drains."""
self._manager.release_external_boundary_admission(request_id)
self._lookups.pop(request_id, None)
if any(task.request_id == request_id for task in self._tasks.values()):
self._cancelled.add(request_id)
Expand Down
Loading