Team Ai
Apppublic

bigscience/petals-api

sourceHugging Faceupdated 4y agoView on Hugging Face
18likes
handler.py230 linesDownload Raw Back to server
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