Skip to content
Open
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
6 changes: 3 additions & 3 deletions daser/connector/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -1885,8 +1885,8 @@ def _submit_finished_save(self, save: _DeferredFinishedSave) -> Any | None:
save: Deferred store plan built during ``wait_for_save``.

Returns:
Future that completes after all store batches are committed, or
``None`` when no store batch could be staged.
``None`` once request KV has been copied into staging and background
store/commit futures have been tracked.

Async/thread-safety:
Called by vLLM's worker thread from ``get_finished`` while vLLM is
Expand Down Expand Up @@ -1917,7 +1917,7 @@ def _submit_finished_save(self, save: _DeferredFinishedSave) -> Any | None:
self._commit_after_store_futures(batch_futures, sorted(save.commit_keys)),
)
self._track_save_future(commit_future, 0, None)
return commit_future
return None

async def _commit_after_store_futures(
self,
Expand Down
7 changes: 6 additions & 1 deletion daser/replacement/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,5 +2,10 @@

from daser.replacement.base import ReplacementPolicy
from daser.replacement.lru import LRUReplacementPolicy
from daser.replacement.prefix_aware_lru import PrefixAwareLRUReplacementPolicy

__all__ = ["LRUReplacementPolicy", "ReplacementPolicy"]
__all__ = [
"LRUReplacementPolicy",
"PrefixAwareLRUReplacementPolicy",
"ReplacementPolicy",
]
166 changes: 166 additions & 0 deletions daser/replacement/prefix_aware_lru.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,166 @@
# SPDX-License-Identifier: Apache-2.0

"""Prefix-aware LRU replacement policy."""

# Standard
from collections.abc import Callable, Hashable, Iterable
from dataclasses import dataclass
from heapq import heappop, heappush
from typing import Generic, TypeVar

# First Party
from daser.replacement.base import ReplacementPolicy

K = TypeVar("K")


@dataclass(frozen=True, order=True)
class _Order:
epoch: int
prefix_rank: int
sequence: int


class PrefixAwareLRUReplacementPolicy(ReplacementPolicy[K], Generic[K]):
"""LRU policy that ages request suffix keys before prefix keys.

Async/thread-safety:
Not thread-safe by itself. Owners should call it from one event loop or
guard it with their own lock.
"""

def __init__(self) -> None:
self._clock = 0
self._sequence = 0
self._orders: dict[K, _Order] = {}
self._heap: list[tuple[int, int, int, K]] = []
self._group_epochs: dict[Hashable, int] = {}
self._key_groups: dict[K, Hashable] = {}
self._group_counts: dict[Hashable, int] = {}

def insert(self, key: K) -> None:
"""Insert or refresh a standalone key.

Args:
key: Cache key to track.
"""
self._set_order(key, self._next_epoch(), 0, None)

def insert_prefix(self, key: K, group: Hashable, prefix_index: int) -> None:
"""Insert or refresh one key using prefix-aware group ordering.

Args:
key: Cache key to track.
group: Stable request/chunk identity shared by sibling prefix keys.
prefix_index: Zero-based slot index inside the request prefix.
"""
epoch = self._group_epochs.get(group)
if epoch is None:
epoch = self._next_epoch()
self._group_epochs[group] = epoch
self._set_order(key, epoch, -int(prefix_index), group)

def access(self, key: K) -> None:
"""Mark an existing key as recently used.

Args:
key: Cache key that was read or written.
"""
if key in self._orders:
self._set_order(key, self._next_epoch(), 0, None)

def access_prefix(self, keys: Iterable[tuple[K, int]]) -> None:
"""Refresh sibling prefix keys with suffix-before-prefix aging.

Args:
keys: ``(key, prefix_index)`` pairs in any order.
"""
present = [
(key, int(prefix_index))
for key, prefix_index in keys
if key in self._orders
]
if not present:
return
epoch = self._next_epoch()
for key, prefix_index in present:
self._set_order(key, epoch, -prefix_index, None)

def remove(self, key: K) -> None:
"""Remove a key from replacement tracking.

Args:
key: Cache key to forget.
"""
if key not in self._orders:
return
self._orders.pop(key, None)
self._drop_key_group(key)

def evict(self) -> K | None:
"""Return and remove the least-recently-used key.

Returns:
Victim key, or None when empty.
"""
while self._heap:
epoch, prefix_rank, sequence, key = heappop(self._heap)
order = self._orders.get(key)
if order != _Order(epoch, prefix_rank, sequence):
continue
self.remove(key)
return key
return None

def evict_matching(self, predicate: Callable[[K], bool]) -> K | None:
"""Return and remove the oldest key accepted by ``predicate``.

Args:
predicate: Function returning True for eligible victim keys.

Returns:
Victim key, or None when no tracked key matches.
"""
victims = [
(order, key) for key, order in self._orders.items() if predicate(key)
]
victim = min(victims) if victims else None
if victim is None:
return None
_order, key = victim
self.remove(key)
return key

def _next_epoch(self) -> int:
"""Return the next recency epoch."""
self._clock += 1
return self._clock

def _set_order(
self,
key: K,
epoch: int,
prefix_rank: int,
group: Hashable | None,
) -> None:
"""Install a key order and push a lazy heap record."""
self._drop_key_group(key)
self._sequence += 1
order = _Order(epoch, prefix_rank, self._sequence)
self._orders[key] = order
if group is not None:
self._key_groups[key] = group
self._group_counts[group] = self._group_counts.get(group, 0) + 1
heappush(self._heap, (order.epoch, order.prefix_rank, order.sequence, key))

def _drop_key_group(self, key: K) -> None:
"""Remove a key from its prefix group accounting."""
group = self._key_groups.pop(key, None)
if group is None:
return
count = self._group_counts.get(group, 0) - 1
if count > 0:
self._group_counts[group] = count
return
self._group_counts.pop(group, None)
self._group_epochs.pop(group, None)
109 changes: 90 additions & 19 deletions daser/transfer/iouring/l1_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@
from dataclasses import dataclass

# First Party
from daser.replacement import LRUReplacementPolicy
from daser.replacement import PrefixAwareLRUReplacementPolicy
from daser.transfer.iouring.pinned_pool import PinnedMemoryPool, PinnedMemorySlice


Expand Down Expand Up @@ -59,7 +59,8 @@ def __init__(
self._starts: list[int] = []
self._by_start: dict[int, tuple[int, int]] = {}
self._used = 0
self._policy = LRUReplacementPolicy[tuple[int, int]]()
self._policy = PrefixAwareLRUReplacementPolicy[tuple[int, int]]()
self._headroom_keys: set[tuple[int, int]] = set()
self._pool_waiters: list[object] = []
self._is_pinned = pinned_predicate

Expand Down Expand Up @@ -171,8 +172,14 @@ def record_hits(self, hits: list[L1RangeHit]) -> None:
Args:
hits: slices returned by ``resolve_subranges``.
"""
for hit in hits:
self._policy.access(hit.key)
self._policy.access_prefix(
(hit.key, hit.key[0] + hit.source_offset) for hit in hits
)
for hit in sorted(
hits,
key=lambda item: item.key[0] + item.source_offset,
reverse=True,
):
self._entries.move_to_end(hit.key)

def touch(self, key: tuple[int, int]) -> None:
Expand All @@ -190,13 +197,25 @@ def put(self, key: tuple[int, int], data: PinnedMemorySlice) -> None:
self.drop_overlapping(key[0], key[1])
self._insert_entry(key, data)

def put_headroom(self, key: tuple[int, int], data: PinnedMemorySlice) -> None:
"""Insert promoted bytes into the reserved L1 headroom.

Args:
key: ``(file_offset, nbytes)`` range key.
data: pinned slice holding the range's bytes.
"""
self.drop_overlapping(key[0], key[1])
self._insert_entry(key, data, headroom=True)

def reserve(
self,
key: tuple[int, int],
nbytes: int,
*,
drop_overlaps: bool = True,
preserve_overlaps: bool = False,
target_used_bytes: int | None = None,
prefer_headroom_victims: bool = False,
) -> PinnedMemorySlice | None:
"""Try to reserve pinned space for a store or promoted load.

Expand All @@ -206,6 +225,12 @@ def reserve(
drop_overlaps: drop resident ranges overlapping ``key`` first.
preserve_overlaps: keep the non-overlapping remainder of dropped
ranges when ``drop_overlaps`` is set.
target_used_bytes: optional resident-byte target after insertion.
When set, old residents are evicted before allocation until
``bytes_used + nbytes`` fits the target or no victim is
available. The full L1 capacity remains the hard limit.
prefer_headroom_victims: evict promoted headroom entries before
falling back to the global replacement order.

Returns:
A pinned slice, or None when the pool is exhausted and no further
Expand All @@ -221,16 +246,16 @@ def reserve(
)
if drop_overlaps:
self.drop_overlapping(key[0], key[1], preserve_remainder=preserve_overlaps)
target = self._l1_bytes
if target_used_bytes is not None:
target = min(self._l1_bytes, max(nbytes, target_used_bytes))
while self._used + nbytes > target:
if not self._evict_next(prefer_headroom=prefer_headroom_victims):
return None
data = self._pool.allocate(nbytes)
while data is None:
victim = self._policy.evict()
if victim is None:
if not self._evict_next(prefer_headroom=prefer_headroom_victims):
return None
removed = self._entries.pop(victim, None)
self._remove_index(victim)
if removed is not None:
self._used -= len(removed)
self.release(victim, removed)
data = self._pool.allocate(nbytes)
return data

Expand Down Expand Up @@ -274,6 +299,28 @@ def release(self, key: tuple[int, int], data: PinnedMemorySlice) -> None:
data.close()
self.notify_pool_waiters()

def trim_to(
self,
target_used_bytes: int,
*,
prefer_headroom_victims: bool = False,
) -> bool:
"""Evict residents until usage fits ``target_used_bytes``.

Args:
target_used_bytes: Desired resident-byte ceiling.
prefer_headroom_victims: Evict promoted headroom entries before
falling back to the global replacement order.

Returns:
True when usage fits the target, False when no victim was available.
"""
target = min(self._l1_bytes, max(0, target_used_bytes))
while self._used > target:
if not self._evict_next(prefer_headroom=prefer_headroom_victims):
return False
return True

def register_pool_waiter(self, waiter: object) -> None:
"""Register a future to wake when pool space or metadata changes."""
self._pool_waiters.append(waiter)
Expand Down Expand Up @@ -311,6 +358,7 @@ def drop_overlapping(
removed = self._entries.pop(victim, None)
self._remove_index(victim)
self._policy.remove(victim)
self._headroom_keys.discard(victim)
preserved = (
self._preserve_non_overlapping(victim, removed, file_offset, end)
if preserve_remainder and removed is not None
Expand All @@ -322,25 +370,48 @@ def drop_overlapping(
for preserved_key, payload in preserved:
self._put_preserved_fragment(preserved_key, payload)

def _insert_entry(self, key: tuple[int, int], data: PinnedMemorySlice) -> None:
def _insert_entry(
self,
key: tuple[int, int],
data: PinnedMemorySlice,
*,
headroom: bool = False,
) -> None:
"""Insert one non-overlapping entry and enforce capacity."""
if len(data) > self._l1_bytes:
return
self._entries[key] = data
self._insert_index(key)
self._entries.move_to_end(key)
self._policy.insert(key)
if headroom:
self._headroom_keys.add(key)
else:
self._headroom_keys.discard(key)
self._used += len(data)
self.notify_pool_waiters()
while self._used > self._l1_bytes:
victim = self._policy.evict()
if victim is None:
if not self._evict_next():
break
removed = self._entries.pop(victim, None)
self._remove_index(victim)
if removed is not None:
self._used -= len(removed)
self.release(victim, removed)

def _evict_next(self, *, prefer_headroom: bool = False) -> bool:
"""Evict one policy-selected resident entry."""
victim = (
self._policy.evict_matching(lambda key: key in self._headroom_keys)
if prefer_headroom
else None
)
if victim is None:
victim = self._policy.evict()
if victim is None:
return False
removed = self._entries.pop(victim, None)
self._remove_index(victim)
self._headroom_keys.discard(victim)
if removed is not None:
self._used -= len(removed)
self.release(victim, removed)
return True

def _preserve_non_overlapping(
self,
Expand Down
Loading
Loading