Team Ai
Apppublic

bigscience/petals-api

sourceHugging Faceupdated 4y agoView on Hugging Face
18likes
cache.py128 linesDownload Raw Back to server
1"""2A pytorch memory cache that can be allocated by ConnectionHandler (on cpu) and used over multiple calls to Runtime.3 4For now, the only purpose of this code is to ensure that allocated memory will be deleted properly.5 6"""7import contextlib8import ctypes9import multiprocessing as mp10import os11from typing import AsyncContextManager, Dict, Optional, Union12 13import hivemind14import torch15from hivemind import use_hivemind_log_handler16from hivemind.utils import TensorDescriptor, get_logger17 18use_hivemind_log_handler("in_root_logger")19logger = get_logger(__file__)20 21Handle = int22 23 24class MemoryCache:25    """A shared cache for storing tensors that persist across calls. Main use case: storing past attention KVs"""26 27    def __init__(self, device: Union[str, torch.device], max_size_bytes: Optional[int]):28        self.max_size_bytes = max_size_bytes if max_size_bytes is not None else (2**64 - 1)29        self.device = device30        self.lock_metadata, self.size_decreased_event = mp.Lock(), mp.Event()31        self._current_size = mp.Value(ctypes.c_int64, 0, lock=False)32        self._handle_counter = mp.Value(ctypes.c_int64, 0, lock=False)33        self._active_handles: Optional[Dict[Handle, TensorDescriptor]] = None34        self._allocated_tensors: Optional[Dict[Handle, torch.Tensor]] = None35        self.runtime_pid = os.getpid()36 37        self._pipe_recv, self._pipe_send = mp.Pipe(duplex=False)  # any ConnectionHandler -> runtime38        self._pending_messages = mp.Value(ctypes.c_int64, 0, lock=False)39 40    @property41    def current_size_bytes(self) -> int:42        return self._current_size.value43 44    @current_size_bytes.setter45    def current_size_bytes(self, value: int):46        self._current_size.value = value47 48    @property49    def handle_counter(self) -> int:50        return self._handle_counter.value51 52    @handle_counter.setter53    def handle_counter(self, value: int):54        self._handle_counter.value = value55 56    @contextlib.asynccontextmanager57    async def allocate_cache(self, descr: TensorDescriptor) -> AsyncContextManager[Handle]:58        """59        Create a handle that is associated with buffers on unique device. If cache full, raises AllocationFailed.60 61        :param descr: allocate a tensor of this size, dtype, etc62 63        :note: This function should be called by connection handlers, it can be called concurrently from multiple processes.64        Furthermore, it can be called concurrently with at most one use_cache call in runtime.65        """66        assert os.getpid() != self.runtime_pid, "must be called by a ConnectionHandler, not runtime"67        assert descr.device is None and descr68        allocated_handle = None69        allocated_size_bytes = descr.numel() * torch.finfo(descr.dtype).bits // 870        try:71            async with hivemind.utils.enter_asynchronously(self.lock_metadata):72                if self.current_size_bytes + allocated_size_bytes > self.max_size_bytes:73                    raise AllocationFailed(74                        f"Could not allocate {allocated_size_bytes} bytes in cache; cache size = "75                        f"{self.max_size_bytes} bytes; {self.current_size_bytes} already allocated."76                    )77 78                allocated_handle = int(self.handle_counter)79                self.current_size_bytes += allocated_size_bytes80                self.handle_counter += 1  # note: this will eventually overflow and it is okay81                self._pending_messages.value += 182                self._pipe_send.send((allocated_handle, descr))83 84            yield allocated_handle85        finally:86            if allocated_handle is not None:87                async with hivemind.utils.enter_asynchronously(self.lock_metadata):88                    self._pending_messages.value += 189                    self._pipe_send.send((allocated_handle, None))  # signal runtime to free that handle90                    self.current_size_bytes -= allocated_size_bytes91 92    @contextlib.contextmanager93    def use_cache(self, handle: Handle) -> torch.Tensor:94        """95        Return a tensor that was previously allocated with try_allocate_cache,96 97        :note: This method is called by ExpertBackend in runtime: a single process with NO process parallelism.98        However, runtime may call use_cache concurrently with one or more connection handlers calling allocate_cache99        """100        assert os.getpid() == self.runtime_pid101        # note: this specific function is not concurrent, so you can safely allocate/offload/defragment data here102 103        with self.lock_metadata:104            if self._allocated_tensors is None:105                self._allocated_tensors = {}106 107            # read creation/deletion requests from connection handlers108            for i in range(int(self._pending_messages.value)):109                recv_handle, recv_data = self._pipe_recv.recv()110                self._pending_messages.value -= 1111                if isinstance(recv_data, TensorDescriptor):112                    self._allocated_tensors[recv_handle] = recv_data.make_zeros(device=self.device)113                elif recv_data is None:114                    if recv_handle not in self._allocated_tensors:115                        logger.warning(116                            f"Sanity check failed: asked to delete handle {recv_handle}, but there is no such handle"117                        )118                    self._allocated_tensors.pop(recv_handle, None)119                else:120                    logger.error(f"MemoryCache pipe received unexpected message: {recv_data}")121 122        assert handle in self._allocated_tensors, f"Sanity check failed: no such handle ({handle})"123        yield self._allocated_tensors[handle]124 125 126class AllocationFailed(Exception):127    pass128