Team Ai
Apppublic

bigscience/petals-api

sourceHugging Faceupdated 4y agoView on Hugging Face
18likes
remote_block.py136 linesDownload Raw Back to client
1# Note: this code is being actively modified by justheuristic. If you want to change anything about it, please warn me.2from __future__ import annotations3 4import asyncio5import random6from typing import Any, AsyncIterator, Dict, Optional7 8import torch9from hivemind.compression import deserialize_torch_tensor, serialize_torch_tensor10from hivemind.moe.client.expert import RemoteExpert, RemoteExpertWorker11from hivemind.moe.expert_uid import ExpertInfo12from hivemind.p2p import P2P, StubBase13from hivemind.proto import runtime_pb214from hivemind.utils import anext, get_logger, nested_flatten, use_hivemind_log_handler15 16from src.data_structures import RemoteModuleInfo17from src.dht_utils import ModuleUID18from src.server.handler import TransformerConnectionHandler19 20use_hivemind_log_handler("in_root_logger")21logger = get_logger(__file__)22 23 24class RemoteTransformerBlock(RemoteExpert):25    """A class that interacts with a remote module on a specific server for forward/backward or inference"""26 27    def __init__(self, peers_info: RemoteModuleInfo, p2p: P2P):28        peer_info = ExpertInfo(peers_info.uid, random.choice(list(peers_info.peer_ids)))  # TODO replace this29        super().__init__(peer_info, p2p)30 31    @property32    def stub(self) -> StubBase:33        return TransformerConnectionHandler.get_stub(self.p2p, self.peer_id)34 35    def forward(self, inputs: torch.Tensor, **kwargs):36        for k, v in kwargs.items():37            assert v is None or v is False, f"Extra keyword arguments are not yet supported (got {k} = {v})"38        return super().forward(inputs)39 40    def inference_session(self) -> RemoteTransformerBlockInferenceSession:41        """Initialize a new inference session with the specified remote server"""42        _ = self.info  # create _info manually since the built-in property will not work inside RemoteExpertWorker43        return RemoteExpertWorker.run_coroutine(RemoteTransformerBlockInferenceSession._create(self))44 45    def begin_inference_session(self):46        logger.warning("beging_inference_session was renamed to just inference_session")47        return self.inference_session()48 49 50class RemoteTransformerBlockInferenceSession:51    """An interface to a single multi-step *inference* session for a specific remote module with a specific server"""52 53    def __init__(self, uid: ModuleUID, info: Dict[str, Any], inputs_queue: asyncio.Queue, outputs_aiter: AsyncIterator):54        self.uid, self.info = uid, info55        # warning: this code manages async objects that are only usable inside RemoteExpertWorker's background thread;56        # using them in any other EventLoop may cause side-effects including, headaches, diarrhea, and loss of sleep57        self._inputs_queue: asyncio.Queue[runtime_pb2.ExpertRequest] = inputs_queue58        self._outputs_stream: AsyncIterator[runtime_pb2.ExpertResponse] = outputs_aiter59        self.stepped = False60        self.closed = False61 62    @classmethod63    async def _create(64        cls, remote_module: RemoteTransformerBlock, timeout: Optional[float] = None65    ) -> RemoteTransformerBlockInferenceSession:66        """Create a new session for a given remote module. This code is meant to be run inside RemoteExpertWorker"""67        inputs_queue = asyncio.Queue()68        outputs_stream = await remote_module.stub.rpc_inference(69            cls._read_inputs_from_queue(inputs_queue, timeout), timeout=timeout70        )71        return cls(remote_module.uid, remote_module.info, inputs_queue, outputs_stream)72 73    @staticmethod74    async def _read_inputs_from_queue(queue: asyncio.Queue, timeout: Optional[float]) -> AsyncIterator:75        while True:76            next_input_message = await asyncio.wait_for(queue.get(), timeout)77            yield next_input_message78            if not next_input_message.uid and not next_input_message.tensors:79                break  # this message means "done sending"80 81    def step(self, new_hidden_states: torch.Tensor):82        """Inference step: send a chunk of input tensors and receive a chunk of outputs"""83        if self.closed:84            raise Exception("Session is closed, cannot perform step")85        # serialize inputs and put them into the queue86        inputs = (new_hidden_states,)87        outputs_serialized = RemoteExpertWorker.run_coroutine(88            self._step(89                runtime_pb2.ExpertRequest(90                    uid=self.uid,91                    tensors=[92                        serialize_torch_tensor(tensor, proto.compression)93                        for tensor, proto in zip(inputs, nested_flatten(self.info["forward_schema"]))94                    ],95                )96            )97        )98        outputs = list(map(deserialize_torch_tensor, outputs_serialized.tensors))99        assert outputs[0].shape == inputs[0].shape, f"expected outputs[0] to be hidden states but got {outputs[0]}"100        return outputs[0]101 102    async def _step(self, inputs_serialized: runtime_pb2.ExpertRequest) -> runtime_pb2.ExpertResponse:103        """Inference step on serialized data. This code is meant to be run inside RemoteExpertWorker"""104        await self._inputs_queue.put(inputs_serialized)105        self.stepped = True106        return await anext(self._outputs_stream)107 108    def close(self):109        """Finish a given inference session, close the underlying connection"""110        if self._outputs_stream is None:111            return  # already closed112        RemoteExpertWorker.run_coroutine(self._aclose_stream())113        self._outputs_stream = self._inputs_queue = None114        self.closed = True115 116    async def _aclose_stream(self):117        """Close the inference session. This code is meant to be run inside RemoteExpertWorker"""118        if self._outputs_stream is None:119            return  # already closed120        if self.stepped:121            await self._inputs_queue.put(runtime_pb2.ExpertRequest())  # empty request will trigger end of session122            try:123                await anext(self._outputs_stream)124            except StopAsyncIteration:125                pass126 127    def __del__(self):128        self.close()129 130    def __enter__(self):131        assert not self.closed132        return self133 134    def __exit__(self, *exc_details):135        self.close()136