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