bigscience/petals-api
18
1# Note: this code is being actively modified by justheuristic. If you want to change anything about it, please warn me.2import contextlib3from typing import AsyncIterator, Dict, Sequence4 5import torch6from hivemind import DHT, P2PContext, TensorDescriptor, deserialize_torch_tensor, nested_flatten, serialize_torch_tensor7from hivemind.moe.server.connection_handler import ConnectionHandler8from hivemind.p2p.p2p_daemon import DEFAULT_MAX_MSG_SIZE9from hivemind.proto import runtime_pb210from hivemind.utils import as_aiter11from hivemind.utils.asyncio import anext12from hivemind.utils.streaming import split_for_streaming13 14from src.data_structures import CHAIN_DELIMITER, ModuleUID15from src.server.backend import MAX_LENGTH, TransformerBackend16 17 18class TransformerConnectionHandler(ConnectionHandler):19 """Handles three request types: forward, backward and forward-incremental (inference)"""20 21 module_backends: Dict[ModuleUID, TransformerBackend]22 23 def __init__(self, dht: DHT, module_backends: Dict[str, TransformerBackend]):24 super().__init__(dht, module_backends)25 for module_backend in self.module_backends.values():26 assert isinstance(module_backend, TransformerBackend)27 28 async def rpc_inference(29 self, requests: AsyncIterator[runtime_pb2.ExpertRequest], context: P2PContext30 ) -> AsyncIterator[runtime_pb2.ExpertRequest]:31 """Compute a single step of inference using attention cache; update attention cache accordingly."""32 try:33 print("OPENED RPC_INFERENCE")34 request = await anext(requests)35 requested_uids = self._check_header(request)36 requested_backends = tuple(self.module_backends[uid] for uid in requested_uids)37 38 cache_metadata = torch.tensor([[-1, -1]], dtype=torch.int64) # [cache_handle, prefix_length]39 prefix_length = 040 41 async with self._allocate_caches(requested_backends) as cache_handles:42 assert len(cache_handles) == len(requested_backends)43 while request.tensors: # iterate while user is willing to supply tensors44 hidden_states = [deserialize_torch_tensor(tensor) for tensor in request.tensors]45 46 # run request tensors through all requested modules, update caches47 for backend, cache_handle in zip(requested_backends, cache_handles):48 cache_metadata[0, 0], cache_metadata[0, 1] = cache_handle, prefix_length49 assert (50 len(hidden_states) == 1 and hidden_states[0].ndim == 351 ), f"inputs to {type(backend)} must be a list with a single 3d tensor of hidden states"52 53 hidden_states = await backend.inference_pool.submit_task(cache_metadata, *hidden_states)54 assert isinstance(hidden_states, (list, tuple))55 assert len(hidden_states) == 1 and hidden_states[0].ndim == 356 57 # serialize and send last layer outputs58 yield runtime_pb2.ExpertResponse(59 tensors=[60 serialize_torch_tensor(result, proto.compression, allow_inplace=True)61 for result, proto in zip(62 hidden_states, nested_flatten(requested_backends[-1].outputs_schema)63 )64 ]65 )66 67 # prepare for next step68 prefix_length += hidden_states[0].shape[1]69 request = await (anext(requests))70 finally:71 print("CLOSED RPC_INFERENCE")72 73 async def rpc_forward(self, request: runtime_pb2.ExpertRequest, context: P2PContext) -> runtime_pb2.ExpertResponse:74 # Parse request and prepare backends75 hidden_states = [deserialize_torch_tensor(tensor) for tensor in request.tensors]76 requested_uids = self._check_header(request)77 requested_backends = tuple(self.module_backends[uid] for uid in requested_uids)78 79 # Run a chain of requested backends80 for backend in requested_backends:81 assert isinstance(hidden_states, (list, tuple))82 assert (83 len(hidden_states) == 1 and hidden_states[0].ndim == 384 ), f"inputs to {type(backend)} must be a list with a single 3d tensor of hidden states"85 hidden_states = await backend.forward_pool.submit_task(*hidden_states)86 87 # Serialize the overall output and respond88 assert len(hidden_states) == 1 and hidden_states[0].ndim == 389 return runtime_pb2.ExpertResponse(90 tensors=[91 serialize_torch_tensor(result, proto.compression, allow_inplace=True)92 for result, proto in zip(hidden_states, nested_flatten(requested_backends[-1].outputs_schema))93 ]94 )95 96 async def rpc_forward_stream(97 self, requests: AsyncIterator[runtime_pb2.ExpertRequest], context: P2PContext98 ) -> AsyncIterator[runtime_pb2.ExpertRequest]:99 # Parse requests and prepare backends100 uids_header, hidden_states = await self._gather_inputs(requests, context)101 requested_uids = self._check_header_str(uids_header)102 requested_backends = tuple(self.module_backends[uid] for uid in requested_uids)103 104 # Run a chain of requested backends105 for backend in requested_backends:106 assert isinstance(hidden_states, (list, tuple))107 assert (108 len(hidden_states) == 1 and hidden_states[0].ndim == 3109 ), f"inputs to {type(backend)} must be a list with a single 3d tensor of hidden states"110 hidden_states = await backend.forward_pool.submit_task(*hidden_states)111 112 # Serialize the overall output113 assert len(hidden_states) == 1 and hidden_states[0].ndim == 3114 serialized_output = [115 serialize_torch_tensor(result, proto.compression, allow_inplace=True)116 for result, proto in zip(hidden_states, nested_flatten(requested_backends[-1].outputs_schema))117 ]118 119 # Split the serialized_output for streaming and respond120 output_split = [121 part for tensor in serialized_output for part in split_for_streaming(tensor, DEFAULT_MAX_MSG_SIZE)122 ]123 async for part in as_aiter(*output_split):124 yield runtime_pb2.ExpertResponse(tensors=[part])125 126 async def rpc_backward(self, request: runtime_pb2.ExpertRequest, context: P2PContext) -> runtime_pb2.ExpertResponse:127 # Parse requests and prepare backends128 inputs, grads = [deserialize_torch_tensor(tensor) for tensor in request.tensors]129 requested_uids = self._check_header(request)130 requested_backends = tuple(self.module_backends[uid] for uid in requested_uids)131 132 # Run a forward chain to collect intermediate inputs133 # Note that we do not forward for the last module since we do not need its output134 inter_inputs = [inputs]135 for backend in requested_backends[:-1]:136 assert inputs.ndim == 3, f"inputs to {type(backend)} must be a single 3d tensor of hidden states"137 inputs = await backend.forward_pool.submit_task(inputs)138 assert isinstance(inputs, (list, tuple)) and len(inputs) == 1139 inputs = inputs[0]140 inter_inputs.append(inputs)141 142 # Run a chain of requested backends143 for inp, backend in zip(inter_inputs[::-1], requested_backends[::-1]):144 inputs_and_grads = [inp, grads]145 grads = await backend.backward_pool.submit_task(*inputs_and_grads)146 assert isinstance(grads, (list, tuple)) and len(grads) == 1147 grads = grads[0]148 149 # Serialize the overall grad_input and respond150 return runtime_pb2.ExpertResponse(151 tensors=[152 serialize_torch_tensor(result, proto.compression, allow_inplace=True)153 for result, proto in zip([grads], nested_flatten(requested_backends[0].grad_inputs_schema))154 ]155 )156 157 async def rpc_backward_stream(158 self, requests: AsyncIterator[runtime_pb2.ExpertRequest], context: P2PContext159 ) -> AsyncIterator[runtime_pb2.ExpertResponse]:160 uids_header, inputs_and_grads = await self._gather_inputs(requests, context)161 inputs, grads = inputs_and_grads162 requested_uids = self._check_header_str(uids_header)163 requested_backends = tuple(self.module_backends[uid] for uid in requested_uids)164 165 # Run a forward chain to collect intermediate inputs166 # Note that we do not forward for the last module since we do not need its outputs167 inter_inputs = [inputs]168 for backend in requested_backends[:-1]:169 assert inputs.ndim == 3, f"inputs to {type(backend)} must be a single 3d tensor of hidden states"170 inputs = await backend.forward_pool.submit_task(inputs)171 assert isinstance(inputs, (list, tuple)) and len(inputs) == 1172 inputs = inputs[0]173 inter_inputs.append(inputs)174 175 # Run a backward chain for requested backends176 for inp, backend in zip(inter_inputs[::-1], requested_backends[::-1]):177 inputs_and_grads = [inp, grads]178 grads = await backend.backward_pool.submit_task(*inputs_and_grads)179 assert isinstance(grads, (list, tuple)) and len(grads) == 1180 grads = grads[0]181 182 # Serialize the overall grad_inputs183 serialized_grad_inputs = [184 serialize_torch_tensor(result, proto.compression, allow_inplace=True)185 for result, proto in zip([grads], nested_flatten(requested_backends[0].grad_inputs_schema))186 ]187 # Split the serialized_grad_inputs for streaming and respond188 output_split = [189 part for tensor in serialized_grad_inputs for part in split_for_streaming(tensor, DEFAULT_MAX_MSG_SIZE)190 ]191 192 async for part in as_aiter(*output_split):193 yield runtime_pb2.ExpertResponse(tensors=[part])194 195 def _check_header(self, request: runtime_pb2.ExpertRequest) -> Sequence[ModuleUID]:196 """Check that the first request to rpc_inference is valid"""197 uids = (request.uid or "").split(CHAIN_DELIMITER)198 if not uids:199 raise RuntimeError("User did not provide any uids")200 for uid in uids:201 if uid not in self.module_backends:202 raise RuntimeError(f"Remote peer does not serve {uid}")203 return tuple(uids)204 205 def _check_header_str(self, header) -> Sequence[ModuleUID]:206 """Check that the first request to rpc_inference is valid"""207 uids = (header or "").split(CHAIN_DELIMITER)208 if not uids:209 raise RuntimeError("User did not provide any uids")210 for uid in uids:211 if uid not in self.module_backends:212 raise RuntimeError(f"Remote peer does not serve {uid}")213 return tuple(uids)214 215 @contextlib.asynccontextmanager216 async def _allocate_caches(self, backends: Sequence[TransformerBackend]) -> Sequence[int]:217 """Allocate memory caches for each transformer block, return cache handles"""218 async with contextlib.AsyncExitStack() as stack:219 handles = []220 for backend in backends:221 num_heads = backend.module.self_attention.num_heads222 head_dim = backend.module.self_attention.head_dim223 224 cache_descriptor = TensorDescriptor(size=(2, 1, MAX_LENGTH, num_heads, head_dim), dtype=torch.float32)225 # [key_or_value, batch_size, max_length, num_heads, head_dim]226 227 handles.append(await stack.enter_async_context(backend.memory_cache.allocate_cache(cache_descriptor)))228 229 yield handles230 