bigscience/petals-api
18
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 