Team Ai
Apppublic

bigscience/petals-api

sourceHugging Faceupdated 4y agoView on Hugging Face
18likes
remote_sequential.py136 linesDownload Raw Back to client
1from __future__ import annotations2 3import contextlib4import logging5import random6 7import torch8from hivemind import DHT, P2P, get_logger, use_hivemind_log_handler9from hivemind.moe.client.remote_expert_worker import RemoteExpertWorker10from hivemind.moe.expert_uid import ExpertInfo11from torch import nn12 13import src14from src.client.remote_block import RemoteTransformerBlock15from src.client.remote_sequence_info import RemoteSequenceInfo16from src.data_structures import UID_DELIMITER17from src.dht_utils import _create_remote_modules_from_infos18 19use_hivemind_log_handler("in_root_logger")20logger = get_logger(__file__)21 22 23class RemoteSequential(nn.Module):24    """25    A sequence of transformer blocks hosted by the swarm.26    """27 28    def __init__(self, config: src.DistributedBloomConfig, dht: DHT, prefix: str, max_retries: int = 3):29        logger.warning(f"{self.__class__.__name__} is in active development; expect adventures")30        if prefix.endswith(UID_DELIMITER):31            logger.warning(32                f"dht_prefix {prefix} already ends with '{UID_DELIMITER}'."33                f"This will cause {self.__class__.__name__} to look for modules under "34                f"{prefix}{UID_DELIMITER}*. Please make sure this is what you intended."35            )36 37        super().__init__()38        self.config = config39        self.dht = dht40        self.prefix = prefix41        self.max_retries = max_retries42        self.p2p = RemoteExpertWorker.run_coroutine(dht.replicate_p2p())43 44        block_uids = tuple(f"{prefix}{UID_DELIMITER}{i}" for i in range(config.n_layer))45 46        logger.debug(f"Remote block uids: {block_uids}")47        self.remote_sequence_info = RemoteSequenceInfo(dht, block_uids)48 49    def forward(self, inputs: torch.Tensor):50        assert isinstance(inputs, torch.Tensor) and inputs.ndim == 3 and inputs.shape[-1] == self.config.n_embed51        for block_index in range(self.config.n_layer):52            for retry_index in range(self.max_retries):53                try:54                    block = self[block_index]55                    (outputs,) = block(inputs)56                    assert isinstance(outputs, torch.Tensor)57                    assert outputs.shape == inputs.shape, f"Expected {block} output {inputs.shape}, got {outputs.shape}"58                    inputs = outputs59                    break60                except Exception as e:61                    if retry_index == self.max_retries - 1:62                        raise e63                    else:64                        logging.debug(f"Caught {e} when running forward for block {block_index}", exc_info=True)65        return inputs66 67    def __getitem__(self, block_index: int):68        assert 0 <= block_index < self.config.n_layer69        (module,) = _create_remote_modules_from_infos([self.remote_sequence_info.block_infos[block_index]], self.p2p)70        return module71 72    def __iter__(self):73        for block_index in range(self.config.n_layer):74            yield self[block_index]75 76    def __len__(self):77        return len(self.remote_sequence_info)78 79    def inference_session(self) -> RemoteSequentialInferenceSession:80        self.remote_sequence_info.update_()81        return RemoteSequentialInferenceSession(self.remote_sequence_info, self.p2p)82 83 84class RemoteSequentialInferenceSession:85    """An interface to a multi-step *inference* session for a sequence of remote transformer blocks"""86 87    def __init__(self, remote_sequence_info: RemoteSequenceInfo, p2p: P2P):88        self.remote_sequence_info = remote_sequence_info89        self.p2p = p2p90        self.closed = False91        self.stack = contextlib.ExitStack()92        self.active_sessions = []93 94    def __enter__(self):95        assert not self.closed96        self.stack.__enter__()97        # TODO(yozh) replace this code with a fault-tolerant chain that can be reconstructed if some peers fail98        current_block = 099        while current_block != len(self.remote_sequence_info):100            candidate_spans = self.remote_sequence_info.spans_containing_block[current_block]101            chosen_span = random.choice(candidate_spans)  # TODO this is a temporary code102            assert chosen_span.start <= current_block < chosen_span.end103 104            # TODO begin throwaway prototype code105            remote = RemoteTransformerBlock(self.remote_sequence_info.block_infos[current_block], self.p2p)106            _ = remote.info  # TODO fix107            span_uids = self.remote_sequence_info.block_uids[current_block : chosen_span.end]108            remote._info = ExpertInfo(" ".join(span_uids), chosen_span.peer_id)109            self.active_sessions.append(remote.inference_session())110            self.stack.enter_context(self.active_sessions[-1])111            current_block = chosen_span.end112            # TODO end throwaway prototype code113 114        return self115 116    def step(self, inputs: torch.Tensor):117        assert not self.closed118        for session in self.active_sessions:119            outputs = session.step(inputs)120            assert outputs.shape == inputs.shape, f"expected {inputs.shape}, got {outputs.shape}"121            inputs = outputs122        return inputs123 124    def close(self, *exc_details):125        """Finish a given inference session, close the underlying connection"""126        if not self.closed:127            self.stack.__exit__(*exc_details or (None, None, None))128            self.active_sessions.clear()129            self.closed = True130 131    def __exit__(self, *exc_details):132        self.close(*exc_details)133 134    def __del__(self):135        self.close()136