CoolFace
Apppublic

Mike0021/zonos2

sourceHugging Faceupdated 4mo agoView on Hugging Face
3likes
cache.py87 linesDownload Raw Back to scheduler
1from __future__ import annotations2 3from typing import TYPE_CHECKING4 5import torch6from zonos2.kvcache import BaseCacheHandle, create_cache_manager7 8if TYPE_CHECKING:9    from .utils import PendingReq10 11 12class CacheManager:13    def __init__(self, device: torch.device, num_pages: int, type: str):14        # TODO: support page_size > 115        self._free_slots = torch.arange(num_pages, dtype=torch.int32, device=device)16        self.device = device17        self.manager = create_cache_manager(device=device, type=type)18        self.num_pages = num_pages19 20    def _free(self, indices: torch.Tensor) -> None:21        if len(indices) > 0:22            self._free_slots = torch.cat([self._free_slots, indices])23 24    def match_req(self, req: PendingReq):25        input_len = req.input_len26        assert input_len > 0, "Input length must be greater than 0."27        return self.manager.match_prefix(req.input_ids[: input_len - 1])28 29    def allocate_new_handle(self) -> BaseCacheHandle:30        """Allocate a new empty cache handle (for TTS which doesn't use prefix caching)."""31        # Use match_prefix with empty tensor to get an empty handle32        handle, _ = self.manager.match_prefix(torch.empty(0, dtype=torch.int32, device=self.device))33        return handle34 35    def free_handle(self, handle: BaseCacheHandle) -> None:36        """Free a cache handle (for TTS which doesn't use prefix caching)."""37        # For TTS, we just unlock the handle - no prefix insertion needed38        self.unlock(handle)39 40    def free_slots(self, indices: torch.Tensor) -> None:41        """Free cache slots directly (for TTS which doesn't use prefix caching)."""42        self._free(indices)43 44    @property45    def available_size(self) -> int:46        return self.manager.size_info.evictable_size + len(self._free_slots)47 48    def lock(self, handle: BaseCacheHandle) -> None:49        self.manager.lock_handle(handle, unlock=False)50 51    def unlock(self, handle: BaseCacheHandle) -> None:52        self.manager.lock_handle(handle, unlock=True)53 54    def allocate(self, needed_len: int) -> torch.Tensor:55        if needed_len <= (free_len := len(self._free_slots)):56            allocated = self._free_slots[:needed_len]57            self._free_slots = self._free_slots[needed_len:]58            return allocated59 60        # NOTE: len(evicted) + free_len >= needed_len61        evicted = self.manager.evict(needed_len - free_len)62        merged = torch.cat([self._free_slots, evicted])63        assert len(merged) >= needed_len, "Eviction did not free enough space."64 65        allocated = merged[:needed_len]66        self._free_slots = merged[needed_len:]67        return allocated68 69    def free_and_cache_finished_req(70        self,71        old_handle: BaseCacheHandle,72        input_ids: torch.Tensor,73        indices: torch.Tensor,74    ) -> None:75        in_cache_len = self.manager.insert_prefix(input_ids, indices)76        self._free(indices[old_handle.cached_len : in_cache_len])77        self.unlock(old_handle)78 79    def check_integrity(self) -> None:80        self.manager.check_integrity()81        if len(self._free_slots) + self.manager.size_info.total_size != self.num_pages:82            raise RuntimeError(83                "CacheManager integrity check failed:"84                f" free_slots({len(self._free_slots)}) +"85                f" total_size({self.manager.size_info.total_size}) != num_pages({self.num_pages})"86            )87