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
47 changes: 24 additions & 23 deletions src/kvcr/local_disk.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
import os
from collections import deque
from collections.abc import Callable, Collection, Mapping
from contextlib import closing
from dataclasses import dataclass, field
from pathlib import Path
from typing import TYPE_CHECKING, Any, Literal
Expand Down Expand Up @@ -516,29 +517,29 @@ def _allocate_slot(self, protected: set[BlockKey]) -> int | None:
if self._free_slots:
return self._free_slots.popleft()
self._retry_unscored()
skipped = set(protected)
while (key := self._evictable.select(skipped)) is not None:
record = self._kvcr._block_record_map.get(key)
residency = record.g3 if record is not None else None
if record is None or residency is None or residency.claim_count:
raise RuntimeError(f"invalid G3 eviction candidate {key!r}")
if any(op_id[0] == "fetch" for op_id in record.in_flight_ops or ()):
skipped.add(key)
continue
decision = self._kvcr._policy.decide_eviction(
self._kvcr._block_meta(key, record, self._slot_size),
CacheTier.G3,
)
if decision[0] is PlacementAction.KEEP:
skipped.add(key)
continue
self._remove_evictable(key)
record.g3 = None
self._residency_observer(key, record)
self._kvcr._on_remove(self._kvcr._block_meta(key, record, self._slot_size))
self._kvcr._publish_inventory((key,), CacheTier.G3, removed=True)
self._kvcr._prune_block_record(key)
return residency.slot
with closing(self._evictable.candidates(protected)) as candidates:
for key in candidates:
record = self._kvcr._block_record_map.get(key)
residency = record.g3 if record is not None else None
if record is None or residency is None or residency.claim_count:
raise RuntimeError(f"invalid G3 eviction candidate {key!r}")
if any(op_id[0] == "fetch" for op_id in record.in_flight_ops or ()):
continue
decision = self._kvcr._policy.decide_eviction(
self._kvcr._block_meta(key, record, self._slot_size),
CacheTier.G3,
)
if decision[0] is PlacementAction.KEEP:
continue
self._remove_evictable(key)
record.g3 = None
self._residency_observer(key, record)
self._kvcr._on_remove(
self._kvcr._block_meta(key, record, self._slot_size)
)
self._kvcr._publish_inventory((key,), CacheTier.G3, removed=True)
self._kvcr._prune_block_record(key)
return residency.slot
return None

def _commit(self, reservation: _Reservation) -> None:
Expand Down
66 changes: 32 additions & 34 deletions src/kvcr/local_dram.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
import logging
from collections import Counter, deque
from collections.abc import Callable, Collection, Mapping
from contextlib import closing
from dataclasses import dataclass, field
from enum import Enum, auto
from typing import TYPE_CHECKING, cast
Expand Down Expand Up @@ -1021,7 +1022,6 @@ def _allocate_slots(
if self._capacity_eviction_key is not None:
return None, [], True
self._retry_unscored()
skipped = set(protected)
victims: list[tuple[BlockKey, "_BlockRecord", _LocalDramResidency, int]] = []
freed: Counter[str] = Counter()

Expand All @@ -1032,39 +1032,37 @@ def short() -> set[str]:
if len(self._free_slots[name]) + freed[name] < count
}

while deficient := short():
key = self._evictable.select(skipped)
if key is None:
return None, [], False
record = self._kvcr._block_record_map.get(key)
residency = record.local_dram if record is not None else None
if (
record is None
or residency is None
or residency.state is not _LocalDramState.READY
or residency.claim_count
):
raise RuntimeError(f"invalid evictable local DRAM entry {key!r}")
if not any(name in deficient for name, _ in residency.slots):
skipped.add(key)
continue
size_bytes = self._size_bytes(residency.slots)
decision, eviction_pending = self._kvcr._decide_eviction(
self._kvcr._block_meta(key, record, size_bytes),
CacheTier.LOCAL_G2,
deadline,
)
if not short():
break
if eviction_pending:
self._capacity_eviction_key = key
return None, [], True
if decision[0] is PlacementAction.KEEP:
skipped.add(key)
continue
victims.append((key, record, residency, size_bytes))
skipped.add(key)
freed.update(name for name, _ in residency.slots)
with closing(self._evictable.candidates(protected)) as candidates:
while deficient := short():
key = next(candidates, None)
if key is None:
return None, [], False
record = self._kvcr._block_record_map.get(key)
residency = record.local_dram if record is not None else None
if (
record is None
or residency is None
or residency.state is not _LocalDramState.READY
or residency.claim_count
):
raise RuntimeError(f"invalid evictable local DRAM entry {key!r}")
if not any(name in deficient for name, _ in residency.slots):
continue
size_bytes = self._size_bytes(residency.slots)
decision, eviction_pending = self._kvcr._decide_eviction(
self._kvcr._block_meta(key, record, size_bytes),
CacheTier.LOCAL_G2,
deadline,
)
if not short():
break
if eviction_pending:
self._capacity_eviction_key = key
return None, [], True
if decision[0] is PlacementAction.KEEP:
continue
victims.append((key, record, residency, size_bytes))
freed.update(name for name, _ in residency.slots)

for key, record, residency, size_bytes in victims:
self._remove_evictable(key, residency)
Expand Down
36 changes: 19 additions & 17 deletions src/kvcr/policy_runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
import heapq
import logging
import math
from collections.abc import Collection
from collections.abc import Collection, Generator
from dataclasses import dataclass

from .policy import KVCachePolicy
Expand Down Expand Up @@ -152,7 +152,7 @@ class _Entry:

class _EvictionQueue:
# TODO: Bound stale heap growth from DRAM/G3 claim/release cycles.
# Removal only invalidates _live; select() removes stale heap entries as
# Removal only invalidates _live; candidates() removes stale heap entries as
# it encounters them, so repeated cache use can grow _heap without eviction.
def __init__(self) -> None:
self._heap: list[tuple[float, int, BlockKey]] = []
Expand All @@ -171,19 +171,21 @@ def insert(self, key: BlockKey, score: float) -> None:
def remove(self, key: BlockKey) -> bool:
return self._live.pop(key, None) is not None

def select(self, excluded: set[BlockKey]) -> BlockKey | None:
def candidates(self, excluded: set[BlockKey]) -> Generator[BlockKey, None, None]:
"""Visit each key once in score order; close to restore live entries."""
excluded = set(excluded)
skipped: list[tuple[float, int, BlockKey]] = []
selected: BlockKey | None = None
while self._heap:
score, sequence, key = self._heap[0]
entry = self._live.get(key)
if entry != _Entry(score, sequence):
heapq.heappop(self._heap)
continue
if key not in excluded:
selected = key
break
skipped.append(heapq.heappop(self._heap))
for item in skipped:
heapq.heappush(self._heap, item)
return selected
try:
while self._heap:
item = heapq.heappop(self._heap)
score, sequence, key = item
if self._live.get(key) != _Entry(score, sequence):
continue
skipped.append(item)
if key not in excluded:
excluded.add(key)
yield key
finally:
for score, sequence, key in skipped:
if self._live.get(key) == _Entry(score, sequence):
heapq.heappush(self._heap, (score, sequence, key))
48 changes: 37 additions & 11 deletions tests/unit/test_kvcr_local_dram.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
"""KVCR local-DRAM, capacity, and policy tests."""

import ctypes
import heapq
import logging
from unittest.mock import Mock

Expand Down Expand Up @@ -227,37 +228,62 @@ def test_failed_group_reservation_does_not_evict_a_partial_group() -> None:
(QueryStatus.HIT, CacheTier.LOCAL_G2),
(QueryStatus.HIT, CacheTier.LOCAL_G2),
]
retry = kvcr.deposit({grouped: [descriptors[0]]})
assert dict(_poll_until(kvcr, lambda done: retry in dict(done)))[retry][
grouped
].success


def test_group_allocation_evicts_enough_whole_keys() -> None:
pools = [ctypes.create_string_buffer(16), ctypes.create_string_buffer(16)]
def test_group_allocation_evicts_enough_whole_keys(monkeypatch) -> None:
full = tuple(BlockKey(f"full{index}".encode()) for index in range(32))
pools = [
ctypes.create_string_buffer((len(full) + 1) * 8),
ctypes.create_string_buffer(16),
]
source = ctypes.create_string_buffer(24)
agent = FakeNixlAgent()
agent.state = "DONE"
kvcr = _two_pool_kvcr(agent, pools)
full, swa0, swa1, grouped = (
BlockKey(name) for name in (b"full", b"swa0", b"swa1", b"grouped")
)
swa0, swa1, grouped = (BlockKey(name) for name in (b"swa0", b"swa1", b"grouped"))
descriptors = [
_mem_descriptor(ctypes.addressof(source), 8, info="full"),
_mem_descriptor(ctypes.addressof(source) + 8, 8, info="swa"),
_mem_descriptor(ctypes.addressof(source) + 16, 8, info="swa"),
]

for key, descriptor in zip((full, swa0, swa1), descriptors, strict=True):
operation = kvcr.deposit({key: [descriptor]})
_poll_until(kvcr, lambda done: operation in dict(done))
operation = kvcr.deposit(
{
**{key: [descriptors[0]] for key in full},
swa0: [descriptors[1]],
swa1: [descriptors[2]],
}
)
_poll_until(kvcr, lambda done: operation in dict(done))
pops = Mock(wraps=heapq.heappop)
monkeypatch.setattr(heapq, "heappop", pops)
operation = kvcr.deposit({grouped: descriptors})
result = dict(_poll_until(kvcr, lambda done: operation in dict(done)))[operation]

assert result[grouped].success
assert kvcr.query((full, swa0, swa1, grouped)) == [
(QueryStatus.HIT, CacheTier.LOCAL_G2),
# The older full-only rows must be scanned once while making room in swa.
assert pops.call_count <= len(full) + 2
assert kvcr.query(full) == [(QueryStatus.HIT, CacheTier.LOCAL_G2)] * len(full)
assert kvcr.query((swa0, swa1, grouped)) == [
(QueryStatus.MISS, None),
(QueryStatus.MISS, None),
(QueryStatus.HIT, CacheTier.LOCAL_G2),
]
assert kvcr._core._local_dram.telemetry_state()["local_g2_evictable_slots"] == 4
assert kvcr._core._local_dram.telemetry_state()["local_g2_evictable_slots"] == (
len(full) + 3
)
retry = kvcr.deposit({swa0: [descriptors[0]]})
assert dict(_poll_until(kvcr, lambda done: retry in dict(done)))[retry][
swa0
].success
assert kvcr.query(full[:2]) == [
(QueryStatus.MISS, None),
(QueryStatus.HIT, CacheTier.LOCAL_G2),
]


@pytest.mark.parametrize(
Expand Down
Loading