diff --git a/.lil/changes/lmcache-105.json b/.lil/changes/lmcache-105.json new file mode 100644 index 00000000000..86979be1488 --- /dev/null +++ b/.lil/changes/lmcache-105.json @@ -0,0 +1,23 @@ +{ + "schema": "local-inference-release-change/v1", + "id": "lmcache-105", + "category": "fix", + "summary": "A checkpoint restore that cannot get RAM no longer removes an intact disk checkpoint from the directory", + "models": [ + "GLM-5.3-Flash", + "Qwen3.8-Flash-Next" + ], + "compatibility": "No user action required.", + "details": [ + "A restore from disk loads a checkpoint's pages only into RAM reserved for all of them. When that reservation failed, the restore was repeated once, and if the repeat failed too while RAM looked free by the time the cache checked (RAM held briefly by other restores or writes, or free RAM split into pieces smaller than a page), the checkpoint was removed from the directory although every page was intact on disk. Later turns of that conversation then recomputed their prompts instead of restoring.", + "The cache now decides from the allocation result itself: a restore whose pages could not be reserved is repeated, 20 ms apart, until the store admission timeout (default 8 s), and then misses with its checkpoint still listed, so a later request restores it. Only a restore that got its RAM and still found pages unreadable retires the checkpoint, also when the first lookup answers after the timeout. A checkpoint larger than the whole RAM tier also stays listed, and the room a waiting restore asks eviction for now includes each page's 4 KiB alignment." + ], + "pull_requests": [ + 105 + ], + "requires": [], + "authors": [ + "ktsaou", + "voipmonitor" + ] +} diff --git a/lmcache/v1/distributed/storage_controllers/prefetch_controller.py b/lmcache/v1/distributed/storage_controllers/prefetch_controller.py index 54451854ca2..0e6f8e25355 100644 --- a/lmcache/v1/distributed/storage_controllers/prefetch_controller.py +++ b/lmcache/v1/distributed/storage_controllers/prefetch_controller.py @@ -193,6 +193,18 @@ class PrefetchPhase(enum.Enum): PLAN_AND_LOAD = enum.auto() +@dataclass(frozen=True) +class PrefetchResult: + """Retained keys and whether reserving their load buffers failed. + + Allocation failure or contention does not prove stored pages are absent. + ``found`` indexes the keys submitted to the prefetch request. + """ + + found: Bitmap + reservation_failed: bool = False + + @dataclass class InFlightPrefetchRequest: """Tracks a single prefetch request across its lifecycle phases.""" @@ -245,6 +257,8 @@ class InFlightPrefetchRequest: """Maps object_group_id to that group's layout (one ``MemoryLayoutDesc`` describes a single group's MemoryObj). Covers every object group.""" + reservation_failed: bool = False + def all_lookups_done(self) -> bool: return len(self.pending_lookup_tasks) == 0 @@ -338,6 +352,7 @@ def __init__( self._prefetch_results_lock = threading.Lock() self._prefetch_results_cv = threading.Condition(self._prefetch_results_lock) self._completed_results: dict[PrefetchRequestId, Bitmap] = {} + self._completed_reservation_failures: set[PrefetchRequestId] = set() # Map eventfds to adapter indices for quick lookup in poll. # Relies on the L2AdapterInterface contract that every adapter @@ -488,12 +503,26 @@ def query_prefetch_result(self, request_id: PrefetchRequestId) -> Bitmap | None: query_lookup_result after calling this function, otherwise it will get None forever. """ + result = self.query_prefetch_result_detailed(request_id) + return result.found if result is not None else None + + def query_prefetch_result_detailed( + self, request_id: PrefetchRequestId + ) -> PrefetchResult | None: + """Consume hits and allocation evidence for ``request_id`` atomically. + + Returns None while pending or already consumed. Bitmap-only queries + consume the same result and release the same lookup bookkeeping. + """ with self._prefetch_results_lock: - result = self._completed_results.pop(request_id, None) - if result is not None: - with self._lookup_results_lock: - self._completed_lookups.pop(request_id, None) - return result + found = self._completed_results.pop(request_id, None) + if found is None: + return None + failed = request_id in self._completed_reservation_failures + self._completed_reservation_failures.discard(request_id) + with self._lookup_results_lock: + self._completed_lookups.pop(request_id, None) + return PrefetchResult(found, failed) def wait_prefetch_result( self, request_id: PrefetchRequestId, timeout: float @@ -1102,6 +1131,7 @@ def _reserve_load_buffers( request.write_reserved_objs[key] = mem_obj reserved.add(key) continue + request.reservation_failed = True if err == L1Error.OUT_OF_MEMORY: oom_keys.append(key) elif err == L1Error.KEY_NOT_WRITABLE: @@ -1600,6 +1630,9 @@ def _release_l2_locks( def _complete_request(self, request_id: PrefetchRequestId, result: Bitmap) -> None: """Store the retained-key bitmap and remove from in-flight tracking.""" with self._prefetch_results_lock: + request = self._in_flight_requests.get(request_id) + if request is not None and request.reservation_failed: + self._completed_reservation_failures.add(request_id) self._completed_results[request_id] = result # Wake any WAIT_PREFETCH_STATUS handler blocked on this result. self._prefetch_results_cv.notify_all() diff --git a/lmcache/v1/distributed/storage_manager.py b/lmcache/v1/distributed/storage_manager.py index e8dba946e1f..72a943ac218 100644 --- a/lmcache/v1/distributed/storage_manager.py +++ b/lmcache/v1/distributed/storage_manager.py @@ -56,6 +56,9 @@ PrefetchController, StoreController, ) +from lmcache.v1.distributed.storage_controllers.prefetch_controller import ( + PrefetchResult, +) from lmcache.v1.distributed.storage_controllers.prefetch_policy import ( create_prefetch_policy, ) @@ -1030,13 +1033,27 @@ def query_prefetch_status( done, None if it's still in progress. Derive the prefix hit count via ``count_leading_ones``. """ + result = self.query_prefetch_status_detailed(handle) + return result.found if result is not None else None + + def query_prefetch_status_detailed( + self, handle: PrefetchHandle + ) -> PrefetchResult | None: + """Consume hits and load-allocation evidence for ``handle``. + + Returns None while pending or already consumed. Found bits index the + original requested keys. Bitmap-only queries consume the same result. + """ l2_r: Bitmap | None = None + reservation_failed = False if handle.prefetch_request_id != -1: - l2_r = self._prefetch_controller.query_prefetch_result( + result = self._prefetch_controller.query_prefetch_result_detailed( handle.prefetch_request_id ) - if l2_r is None: + if result is None: return None + l2_r = result.found + reservation_failed = result.reservation_failed found = self._combine_found(handle, l2_r) # popcount (not count_leading_ones) so the log is accurate for @@ -1060,7 +1077,7 @@ def query_prefetch_status( handle.external_request_id, handle.prefetch_request_id, ) - return found + return PrefetchResult(found, reservation_failed) def touch_l1_keys(self, keys: list[ObjectKey]): """ diff --git a/lmcache/v1/multiprocess/checkpoint_storage.py b/lmcache/v1/multiprocess/checkpoint_storage.py index fd1fb6dbd1e..9c8ef9477e2 100644 --- a/lmcache/v1/multiprocess/checkpoint_storage.py +++ b/lmcache/v1/multiprocess/checkpoint_storage.py @@ -49,6 +49,8 @@ # seconds, as store admission does: a pass frees nothing while its victims are # still pinned by other restores or L2 writes. _ROOM_REQUEST_INTERVAL_SECONDS = 0.5 +# Pause before repeating a lookup whose load buffers could not be reserved. +_RESERVATION_RETRY_SECONDS = 0.02 # Released cancelled lookups remembered so that a later poll from their worker # still gets the miss; a poll of an older one raises KeyError. _MAX_RELEASED_CANCELLED_LOOKUPS = 4096 @@ -237,6 +239,7 @@ class _RetrieveLease: readable: int = 0 unread_bytes: int = 0 eviction_requested: float = 0.0 + resubmit_after: float = 0.0 class CheckpointPayloadStore: @@ -624,8 +627,8 @@ def poll_retrieve(self, lease_id: str) -> CheckpointSlots | bool | None: The storage manager loads pages from L2 only into RAM it reserved for all of them, so a full RAM leaves stored pages unread. A lookup with unread pages is therefore repeated once RAM has room for them, - and only a repeated lookup that had that room invalidates the - generation. The lease stays pending while eviction makes room, for + and only a repeated lookup without allocation failure invalidates + the generation. The lease stays pending while eviction makes room, for at most the storage admission timeout; a lease that never gets room misses without invalidating its generation. @@ -680,13 +683,16 @@ def _poll(self, lease_id: str) -> CheckpointSlots | bool | None: return lease.slots if lease.handle is None: return self._repeat_with_room(lease_id, lease) - found = self._storage.query_prefetch_status(lease.handle) - if found is None: + result = self._storage.query_prefetch_status_detailed(lease.handle) + if result is None: return None + found = result.found readable_keys = [key for i, key in enumerate(lease.keys) if found.test(i)] if len(readable_keys) != len(lease.keys) or lease.cancelled: self._storage.finish_read_prefetched(readable_keys) - return self._repeat_or_miss(lease_id, lease, readable_keys) + return self._repeat_or_miss( + lease_id, lease, readable_keys, result.reservation_failed + ) try: keys, objects = self._storage.unsafe_read(lease.keys) if keys != lease.keys or len(objects) != len(keys): @@ -711,30 +717,51 @@ def _poll(self, lease_id: str) -> CheckpointSlots | bool | None: raise def _repeat_or_miss( - self, lease_id: str, lease: _RetrieveLease, readable_keys: list[ObjectKey] + self, + lease_id: str, + lease: _RetrieveLease, + readable_keys: list[ObjectKey], + reservation_failed: bool, ) -> bool | None: """Repeat or end a lookup whose pages were not all pinned. Called with the lookup's read locks already released. Returns None when the lookup will be repeated, otherwise False after forgetting the - lease. Only a repeated lookup that had room in RAM invalidates. + lease. Only a repeated lookup without allocation failure invalidates. """ lease.readable = len(readable_keys) if lease.cancelled: return self._miss(lease_id, lease, "was cancelled by the engine") groups = checkpoint_page_groups(lease.manifest) readable = set(readable_keys) + alignment = self._storage.l1_memory_desc.align_bytes lease.unread_bytes = sum( - groups[key.object_group_id].page_bytes + ((groups[key.object_group_id].page_bytes + alignment - 1) // alignment) + * alignment for key in lease.keys if key not in readable ) used, total = self._storage.get_l1_usage() + if reservation_failed and lease.unread_bytes > total: + return self._miss( + lease_id, lease, "cannot fit in RAM; its checkpoint stays listed" + ) + if reservation_failed: + now = time.monotonic() + if now - lease.started >= self._storage.store_admission_timeout_seconds: + return self._miss( + lease_id, lease, "found no room in RAM; its checkpoint stays listed" + ) + lease.resubmit_after = now + _RESERVATION_RETRY_SECONDS # Without L2, or with more pages than RAM can hold, no repeat loads them. if ( self._storage.l2_adapters() and lease.unread_bytes <= total - and (lease.lookups == 1 or total - used < lease.unread_bytes) + and ( + reservation_failed + or lease.lookups == 1 + or total - used < lease.unread_bytes + ) ): lease.handle = None return self._repeat_with_room(lease_id, lease) @@ -751,12 +778,12 @@ def _repeat_with_room(self, lease_id: str, lease: _RetrieveLease) -> bool | None """ if lease.cancelled: return self._miss(lease_id, lease, "was cancelled by the engine") + now = time.monotonic() used, total = self._storage.get_l1_usage() - if total - used >= lease.unread_bytes: + if total - used >= lease.unread_bytes and now >= lease.resubmit_after: lease.handle = self._lookup(lease_id, lease.manifest, lease.keys) lease.lookups += 1 return None - now = time.monotonic() if now - lease.started >= self._storage.store_admission_timeout_seconds: return self._miss( lease_id, lease, "found no room in RAM; its checkpoint stays listed" diff --git a/tests/v1/multiprocess/test_checkpoint_cancel_drain.py b/tests/v1/multiprocess/test_checkpoint_cancel_drain.py index df8e4e305c9..5c0c5aa3fb5 100644 --- a/tests/v1/multiprocess/test_checkpoint_cancel_drain.py +++ b/tests/v1/multiprocess/test_checkpoint_cancel_drain.py @@ -141,10 +141,10 @@ def test_failed_poll_of_a_pending_lookup_is_not_fatal() -> None: def _stall_lookups(monkeypatch: pytest.MonkeyPatch, storage: Any) -> threading.Event: """Keep every prefetch pending until the returned event is set.""" answer = threading.Event() - query = storage.query_prefetch_status + query = storage.query_prefetch_status_detailed monkeypatch.setattr( storage, - "query_prefetch_status", + "query_prefetch_status_detailed", lambda handle: query(handle) if answer.is_set() else None, ) return answer diff --git a/tests/v1/multiprocess/test_checkpoint_restore_capacity.py b/tests/v1/multiprocess/test_checkpoint_restore_capacity.py new file mode 100644 index 00000000000..d57c163266d --- /dev/null +++ b/tests/v1/multiprocess/test_checkpoint_restore_capacity.py @@ -0,0 +1,165 @@ +# SPDX-License-Identifier: Apache-2.0 +"""RAM allocation failures must not delist intact disk checkpoints.""" + +# Standard +from dataclasses import replace +from pathlib import Path +from typing import Any +from unittest.mock import patch +import hashlib + +# Third Party +import pytest +import torch + +# First Party +from lmcache.v1.distributed.api import ( + MemoryLayoutDesc, + ObjectKey, + PrefetchHandle, + PrefetchRequestSpec, + TrimPolicy, +) +from lmcache.v1.distributed.storage_controllers.prefetch_controller import ( + PrefetchResult, +) +from lmcache.v1.multiprocess.checkpoint_storage import ( + checkpoint_object_keys, + checkpoint_page_groups, +) +from tests.v1.multiprocess.test_checkpoint_storage import ( + make_manifest, + open_store, + poll, + publish_before_restart, + restores, +) + + +@pytest.mark.parametrize("native", [False, True]) +def test_restore_accounts_for_page_alignment(tmp_path: Path, native: bool) -> None: + entry = replace(make_manifest(), world_size=1) + publish_before_restart(tmp_path, native, [entry]) + with open_store(tmp_path, native, admission_timeout_seconds=0.15) as ( + service, + index, + storage, + mapping, + ): + filler = ObjectKey(hashlib.sha256(b"alignment-filler").digest(), "filler", 0, 0) + _, total = storage.get_l1_usage() + alignment = storage.l1_memory_desc.align_bytes + reserved = storage.reserve_write_detailed( + [filler], + MemoryLayoutDesc([torch.Size([total - alignment])], [torch.uint8]), + "new", + ) + assert reserved[filler][1] is not None + try: + # Raw payload bytes fit; seven separate aligned allocations do not. + lease_id = service.begin_retrieve(entry, 0) + assert lease_id is not None + assert poll(service, lease_id) is False + assert index.get(entry.generation) == entry + finally: + storage.abort_write([filler]) + assert restores(service, mapping, entry) + + +@pytest.mark.parametrize("native", [False, True]) +def test_restore_uses_allocation_result_before_capacity_returns( + tmp_path: Path, native: bool +) -> None: + entry = replace(make_manifest(), world_size=1) + publish_before_restart(tmp_path, native, [entry]) + with open_store(tmp_path, native, admission_timeout_seconds=0.15) as ( + service, + index, + storage, + mapping, + ): + filler = ObjectKey( + hashlib.sha256(b"concurrent-writer").digest(), "filler", 0, 0 + ) + _, total = storage.get_l1_usage() + submit = storage.submit_prefetch_task + query = storage.query_prefetch_status_detailed + attempts = 0 + + def reserve_then_submit(*args: Any, **kwargs: Any) -> PrefetchHandle: + nonlocal attempts + attempts += 1 + result = storage.reserve_write_detailed( + [filler], MemoryLayoutDesc([torch.Size([total])], [torch.uint8]), "new" + ) + assert result[filler][1] is not None + return submit(*args, **kwargs) + + def release_before_consumption(handle: PrefetchHandle) -> PrefetchResult | None: + result = query(handle) + if result is not None: + # Enforce a concurrent writer finishing after allocation failed. + storage.abort_write([filler]) + return result + + try: + with ( + patch.object( + storage, "submit_prefetch_task", side_effect=reserve_then_submit + ), + patch.object( + storage, + "query_prefetch_status_detailed", + side_effect=release_before_consumption, + ), + ): + lease_id = service.begin_retrieve(entry, 0) + assert lease_id is not None + assert poll(service, lease_id) is False + assert attempts >= 2 + assert index.get(entry.generation) == entry + finally: + storage.abort_write([filler]) + assert restores(service, mapping, entry) + + +@pytest.mark.parametrize("native", [False, True]) +@pytest.mark.parametrize("detailed", [False, True]) +def test_prefetch_queries_consume_capacity_failure_once( + tmp_path: Path, native: bool, detailed: bool +) -> None: + entry = replace(make_manifest(), world_size=1) + publish_before_restart(tmp_path, native, [entry]) + with open_store(tmp_path, native) as (service, index, storage, mapping): + filler = ObjectKey(hashlib.sha256(b"query-filler").digest(), "filler", 0, 0) + _, total = storage.get_l1_usage() + reserved = storage.reserve_write_detailed( + [filler], MemoryLayoutDesc([torch.Size([total])], [torch.uint8]), "new" + ) + assert reserved[filler][1] is not None + keys = [key for group in checkpoint_object_keys(entry, 0) for key in group] + layouts = { + i: MemoryLayoutDesc([torch.Size([group.page_bytes])], [torch.uint8]) + for i, group in enumerate(checkpoint_page_groups(entry)) + } + try: + handle = storage.submit_prefetch_task( + PrefetchRequestSpec(keys, layouts, policy=TrimPolicy.SPARSE) + ) + assert storage.wait_prefetch_status(handle, timeout=5) + if detailed: + result = storage.query_prefetch_status_detailed(handle) + assert result is not None and result.reservation_failed + found = result.found + else: + legacy_found = storage.query_prefetch_status(handle) + assert legacy_found is not None + found = legacy_found + assert found is not None and found.popcount() == 0 + assert storage.query_prefetch_status_detailed(handle) is None + assert storage.query_prefetch_status(handle) is None + assert storage.query_prefetch_lookup_hits(handle) is None + assert index.get(entry.generation) == entry + finally: + storage.abort_write([filler]) + assert restores(service, mapping, entry) diff --git a/tests/v1/multiprocess/test_checkpoint_restore_ram_failures.py b/tests/v1/multiprocess/test_checkpoint_restore_ram_failures.py new file mode 100644 index 00000000000..deb76b020f5 --- /dev/null +++ b/tests/v1/multiprocess/test_checkpoint_restore_ram_failures.py @@ -0,0 +1,180 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Checkpoint restores whose RAM reservation fails keep intact checkpoints listed. + +A restore from disk reserves RAM for all of a checkpoint's pages. When that +reservation fails, the checkpoint must stay listed so a later request restores +it; only pages that are really missing retire it. +""" + +# Standard +from dataclasses import replace +from pathlib import Path +from typing import Any +from unittest.mock import patch +import hashlib +import time + +# Third Party +import pytest +import torch + +# First Party +from lmcache.v1.distributed.api import MemoryLayoutDesc, ObjectKey +from tests.v1.multiprocess.test_checkpoint_storage import ( + entry_keys, + large_manifest, + make_manifest, + open_store, + poll, + publish_before_restart, + restores, +) + + +@pytest.mark.parametrize("native", [False, True]) +def test_ram_released_before_the_poll_keeps_checkpoint_listed( + tmp_path: Path, native: bool +) -> None: + """A writer holds all RAM while the load buffers are reserved and frees it + before the store reads the result: the checkpoint stays listed.""" + entry = replace(make_manifest(), world_size=1) + publish_before_restart(tmp_path, native, [entry]) + with open_store(tmp_path, native, admission_timeout_seconds=0.15) as ( + service, + index, + storage, + mapping, + ): + filler = ObjectKey(hashlib.sha256(b"probe-writer").digest(), "filler", 0, 0) + _, total = storage.get_l1_usage() + submit = storage.submit_prefetch_task + query = storage.query_prefetch_status_detailed + attempts = 0 + + def reserve_then_submit(*args: Any, **kwargs: Any) -> Any: + nonlocal attempts + attempts += 1 + result = storage.reserve_write_detailed( + [filler], MemoryLayoutDesc([torch.Size([total])], [torch.uint8]), "new" + ) + assert result[filler][1] is not None + return submit(*args, **kwargs) + + def release_before_consumption(handle: Any) -> Any: + result = query(handle) + if result is not None: + storage.abort_write([filler]) + return result + + try: + with ( + patch.object( + storage, "submit_prefetch_task", side_effect=reserve_then_submit + ), + patch.object( + storage, + "query_prefetch_status_detailed", + side_effect=release_before_consumption, + ), + ): + lease_id = service.begin_retrieve(entry, 0) + assert lease_id is not None + outcome = poll(service, lease_id) + listed = index.get(entry.generation) == entry + assert outcome is False + assert listed, "intact disk checkpoint was delisted" + finally: + storage.abort_write([filler]) + assert restores(service, mapping, entry) + + +@pytest.mark.parametrize("native", [False, True]) +def test_fragmented_ram_keeps_checkpoint_listed_without_spinning( + tmp_path: Path, native: bool +) -> None: + """Free RAM exceeds the unread pages, but no free block holds one page. + + No concurrency: 64 KiB holes between pinned 64 KiB blocks, 128 KiB + pages. The reservation fails every time while get_l1_usage shows room, so + the lookup is repeated with a pause until the admission timeout. + """ + entry = large_manifest(500) + publish_before_restart(tmp_path, native, [entry]) + with open_store(tmp_path, native, admission_timeout_seconds=1.0) as ( + service, + index, + storage, + mapping, + ): + used, total = storage.get_l1_usage() + assert used == 0 + block = 64 * 1024 + fillers = [ + ObjectKey(hashlib.sha256(b"frag%d" % i).digest(), "filler", 0, 0) + for i in range(total // block) + ] + reserved = storage.reserve_write_detailed( + fillers, MemoryLayoutDesc([torch.Size([block])], [torch.uint8]), "new" + ) + assert all(obj is not None for _err, obj in reserved.values()) + storage.abort_write(fillers[::2]) + kept = fillers[1::2] + used, total = storage.get_l1_usage() + unread = len(entry_keys(entry)) * 128 * 1024 + assert total - used >= unread + submit = storage.submit_prefetch_task + attempts = 0 + + def count(*args: Any, **kwargs: Any) -> Any: + nonlocal attempts + attempts += 1 + return submit(*args, **kwargs) + + try: + with patch.object(storage, "submit_prefetch_task", side_effect=count): + start = time.monotonic() + lease_id = service.begin_retrieve(entry, 0) + assert lease_id is not None + outcome = poll(service, lease_id) + elapsed = time.monotonic() - start + listed = index.get(entry.generation) == entry + assert outcome is False + assert listed, "intact disk checkpoint was delisted" + # Repeats after a RAM failure pause 20 ms: about 50 in one second. + assert attempts <= 80, f"{attempts} lookups in {elapsed:.2f} s" + finally: + storage.abort_write(kept) + assert restores(service, mapping, entry) + + +@pytest.mark.parametrize( + "timeout,delay", + [(0.2, 0.0), (0.2, 0.4), (0.0, 0.0)], + ids=["fast-first-lookup", "first-lookup-slower-than-timeout", "timeout-0"], +) +def test_lost_page_retires_checkpoint( + tmp_path: Path, timeout: float, delay: float +) -> None: + """A page file deleted from disk (not a RAM problem) retires its checkpoint. + + Also when the first lookup is answered after the admission timeout, or the + timeout is zero: the timeout bounds only repeats after a RAM failure. + """ + entry = replace(make_manifest(), world_size=1) + publish_before_restart(tmp_path, False, [entry]) + files = sorted((tmp_path / "payloads").rglob("*.data")) + assert len(files) == len(entry_keys(entry)) + files[0].unlink() + with open_store(tmp_path, False, admission_timeout_seconds=timeout) as ( + service, + index, + storage, + _mapping, + ): + lease_id = service.begin_retrieve(entry, 0) + assert lease_id is not None + time.sleep(delay) + outcome = poll(service, lease_id) + listed = index.get(entry.generation) == entry + assert outcome is False + assert not listed, "checkpoint with a lost page stays listed"