Files
AstrAI/astrai/inference/cache/strategy.py
T
ViperEkura 0c1b7664c1 refactor: split infer core into subpackages by concern
- Eliminate core/ directory into cache/, runtime/, network/ subpackages plus flat modules
- Split cache.py (647 lines) into cache/{buffer,strategy,pool}.py by layer
- Add explicit ContiguousStrategy, make AllocationStrategy a real ABC
- Move TaskCacheState to cache/strategy.py, drop string forward references
- Rename api/ to network/, server.py to app.py
- Move sample.py into runtime/ alongside executor and graph
- Simplify TaskCacheManager.__init__ to single pool param
- Expose pool.strategy and pool.req_pool as public properties
- Fix KVCache import in attention_backend.py (TYPE_CHECKING guard)
- Fix steady-state decode reading uninitialized position_ids on first step
2026-08-08 23:43:05 +08:00

321 lines
10 KiB
Python

"""KV cache allocation layer.
Encapsulates the physical slot allocation policy, isolated from GPU buffers
and task lifecycle management.
- ``TaskCacheState``: data contract between strategy and manager (per-task slot state)
- ``Allocator``: bitmask-based page allocator with LRU eviction
- ``RadixCache``: page-granular prefix index (exact token match)
- ``AllocationStrategy``: ABC for physical slot allocation
- ``ContiguousStrategy``: statically partitioned, no dynamic allocation
- ``PagedStrategy``: dynamic paged allocation from a shared pool
"""
import threading
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Callable, Dict, List, Optional, OrderedDict
import torch
from astrai.inference.cache.buffer import ReqToTokenPool
# ---- data contract: per-task slot state ----
@dataclass
class TaskCacheState:
"""Per-task cache allocation state.
Co-locates all task-owned cache metadata so the alloc/free/extend
lifecycle is atomic. Owned by ``TaskCacheManager``, consumed by
every ``AllocationStrategy`` method.
"""
req_idx: int
length: int = 0
cached: int = 0
pages: List[int] = field(default_factory=list)
# ---- allocation primitives ----
class Allocator:
"""Bitmask-based page allocator with ref-counting and LRU eviction."""
def __init__(self, n_pages: int):
self._free_mask = (1 << n_pages) - 1
self._refs: List[int] = [0] * n_pages
self._lru: OrderedDict[int, None] = OrderedDict()
self.on_evict: Optional[Callable[[int], None]] = None
self._lock = threading.Lock()
def alloc(self) -> int:
with self._lock:
if self._free_mask:
lsb = self._free_mask & -self._free_mask
idx = lsb.bit_length() - 1
self._free_mask ^= lsb
self._refs[idx] = 1
return idx
if self._lru:
idx, _ = self._lru.popitem(last=False)
if self.on_evict:
self.on_evict(idx)
self._refs[idx] = 1
self._free_mask &= ~(1 << idx)
return idx
return -1
def free(self, idx: int, keep_cached: bool = False):
with self._lock:
self._refs[idx] -= 1
if self._refs[idx] == 0:
if keep_cached:
self._lru[idx] = None
else:
self._free_mask |= 1 << idx
def inc_ref(self, idx: int):
with self._lock:
self._refs[idx] += 1
self._lru.pop(idx, None)
def ref_count(self, idx: int) -> int:
with self._lock:
return self._refs[idx]
def touch(self, idx: int):
with self._lock:
if idx in self._lru:
self._lru.move_to_end(idx)
class RadixNode:
"""A page-aligned edge in the CPU-side prefix radix trie."""
__slots__ = ("parent", "children", "page_idx", "tokens", "lock_ref")
def __init__(self, parent=None, tokens=(), page_idx=None):
self.parent = parent
self.children: Dict[tuple, "RadixNode"] = {}
self.page_idx = page_idx
self.tokens = tuple(tokens)
self.lock_ref = 0
class RadixCache:
"""Page-granular radix prefix index with exact token matching."""
def __init__(self, page_size: int):
self._page_size = page_size
self._root = RadixNode()
self._page_to_node: Dict[int, RadixNode] = {}
self._lock = threading.Lock()
def evict(self, idx: int):
with self._lock:
node = self._page_to_node.pop(idx, None)
if node is None:
return
node.page_idx = None
parent = node.parent
if parent is not None:
parent.children.pop(node.tokens, None)
def has_page(self, idx: int) -> bool:
with self._lock:
return idx in self._page_to_node
def lookup(self, token_ids: List[int]) -> List[int]:
with self._lock:
full_pages = len(token_ids) // self._page_size
hits: List[int] = []
node = self._root
for i in range(full_pages):
start = i * self._page_size
page_tokens = tuple(token_ids[start : start + self._page_size])
child = node.children.get(page_tokens)
if child is None or child.page_idx is None:
break
hits.append(child.page_idx)
node = child
return hits
def record(self, page_idx: int, token_ids: List[int], logical_page_idx: int):
with self._lock:
full_pages = len(token_ids) // self._page_size
if logical_page_idx >= full_pages:
return
old = self._page_to_node.pop(page_idx, None)
if old is not None and old.parent is not None:
old.parent.children.pop(old.tokens, None)
node = self._root
for i in range(logical_page_idx + 1):
start = i * self._page_size
page_tokens = tuple(token_ids[start : start + self._page_size])
child = node.children.get(page_tokens)
if child is None:
child = RadixNode(node, page_tokens)
node.children[page_tokens] = child
node = child
if node.page_idx is not None and node.page_idx != page_idx:
replaced = node.page_idx
self._page_to_node.pop(replaced, None)
node.page_idx = page_idx
self._page_to_node[page_idx] = node
def release(self, pages: List[int]) -> None:
with self._lock:
for page_idx in pages:
node = self._page_to_node.get(page_idx)
if node is not None and node.lock_ref:
node.lock_ref -= 1
class AllocationStrategy(ABC):
"""Physical slot allocation policy.
Subclasses implement the actual allocation semantics. This ABC declares
the contract; there are no default implementations.
"""
@abstractmethod
def alloc(self, state: TaskCacheState, prompt_ids: List[int]) -> bool: ...
@abstractmethod
def free(self, state: TaskCacheState) -> None: ...
@abstractmethod
def extend(self, state: TaskCacheState, pos: int) -> bool: ...
@abstractmethod
def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None: ...
@abstractmethod
def record_hashes(
self,
state: TaskCacheState,
prompt_ids: List[int],
start: int,
) -> None: ...
class ContiguousStrategy(AllocationStrategy):
"""Static contiguous allocation: slots are pre-assigned at pool init.
No dynamic allocation or prefix caching. All operations are no-ops
because ``ReqToTokenPool`` is pre-filled with contiguous ranges.
"""
def alloc(self, state: TaskCacheState, prompt_ids: List[int]) -> bool:
return True
def free(self, state: TaskCacheState) -> None:
pass
def extend(self, state: TaskCacheState, pos: int) -> bool:
return True
def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None:
pass
def record_hashes(
self,
state: TaskCacheState,
prompt_ids: List[int],
start: int,
) -> None:
pass
class PagedStrategy(AllocationStrategy):
"""Dynamic paged allocation from a shared bitmask pool.
``page_size`` is a parameter, not a separate strategy: at ``page_size=1``
each allocated page *is* one token slot (``page * 1 + 0``), and prefix
caching is simply disabled (``prefix=None``). The unified page formula
``pages[page_idx] * page_size + offset`` holds for both.
"""
def __init__(
self,
alloc: Allocator,
prefix: Optional[RadixCache],
page_size: int,
req_pool: ReqToTokenPool,
device,
):
self._alloc = alloc
self._prefix = prefix
self._page_size = page_size
self._req_pool = req_pool
self._device = device
def alloc(self, state: TaskCacheState, prompt_ids: List[int]) -> bool:
if self._prefix is not None:
hits = self._prefix.lookup(prompt_ids)
state.cached = len(hits) * self._page_size
for p in hits:
self._alloc.inc_ref(p)
state.pages = list(hits)
remaining = len(prompt_ids) - state.cached
if remaining <= 0:
return True
n_new = (remaining + self._page_size - 1) // self._page_size
for _ in range(n_new):
p = self._alloc.alloc()
if p < 0:
return False
state.pages.append(p)
return True
def free(self, state: TaskCacheState) -> None:
if self._prefix is not None:
for p in state.pages:
keep = self._prefix.has_page(p)
self._alloc.free(p, keep_cached=keep)
if not keep:
self._prefix.evict(p)
else:
for p in state.pages:
self._alloc.free(p)
def extend(self, state: TaskCacheState, pos: int) -> bool:
page_idx = pos // self._page_size
if page_idx >= len(state.pages):
p = self._alloc.alloc()
if p < 0:
return False
state.pages.append(p)
offset = pos % self._page_size
self._req_pool.req_to_token[state.req_idx, pos] = (
state.pages[page_idx] * self._page_size + offset
)
return True
def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None:
total = len(prompt_ids)
for pos in range(state.cached, total):
page_idx = pos // self._page_size
offset = pos % self._page_size
if page_idx < len(state.pages):
self._req_pool.req_to_token[state.req_idx, pos] = (
state.pages[page_idx] * self._page_size + offset
)
def record_hashes(
self,
state: TaskCacheState,
prompt_ids: List[int],
start: int,
) -> None:
if self._prefix is None:
return
full = len(prompt_ids) // self._page_size
for i in range(start, min(full, len(state.pages))):
self._prefix.record(state.pages[i], prompt_ids, i)