Team Ai
Apppublic

bigscience/petals-api

sourceHugging Faceupdated 4y agoView on Hugging Face
18likes
server.py255 linesDownload Raw Back to server
1from __future__ import annotations2 3import multiprocessing as mp4import threading5from typing import Dict, Optional, Sequence, Union6 7import torch8from hivemind import DHT, MAX_DHT_TIME_DISCREPANCY_SECONDS, BatchTensorDescriptor, get_dht_time9from hivemind.moe.server.layers import add_custom_models_from_file10from hivemind.moe.server.runtime import Runtime11from hivemind.proto.runtime_pb2 import CompressionType12from hivemind.utils.logging import get_logger, use_hivemind_log_handler13 14from src import declare_active_modules, BloomConfig15from src.bloom.from_pretrained import DTYPE_MAP, load_pretrained_block16from src.data_structures import CHAIN_DELIMITER, UID_DELIMITER17from src.server.backend import TransformerBackend18from src.server.cache import MemoryCache19from src.server.handler import TransformerConnectionHandler20 21use_hivemind_log_handler("in_root_logger")22logger = get_logger(__file__)23 24 25class Server(threading.Thread):26    """Serves one or more bloom layers for inference, forward and backward; announces oneself to the DHT"""27 28    def __init__(29        self,30        dht: DHT,31        module_backends: Dict[str, TransformerBackend],32        *,33        device: torch.device,34        num_connection_handlers: int = 8,35        update_period: float = 30,36        expiration: Optional[float] = None,37        start: bool,38        **kwargs,39    ):40        threading.Thread.__init__(self)41        self.dht, self.module_backends, self.update_period = dht, module_backends, update_period42        self.conn_handlers = [43            TransformerConnectionHandler(dht, self.module_backends) for _ in range(num_connection_handlers)44        ]45        self.runtime = Runtime(self.module_backends, device=device, **kwargs)46        self.dht_handler_thread = ModuleAnnouncerThread(47            self.module_backends, dht, update_period, expiration, daemon=True48        )49        self.checkpoint_saver = None  # no need to save checkpoints since we do not change model state50 51        if start:52            self.run_in_background(await_ready=True)53 54    def run(self):55        """56        Starts Server in the current thread. Initializes dht if necessary, starts connection handlers,57        runs Runtime (self.runtime) to process incoming requests.58        """59        logger.info(f"Serving {len(self.module_backends)} blocks:")60        for expert_name, backend in self.module_backends.items():61            num_parameters = sum(p.numel() for p in backend.module.parameters() if p.requires_grad)62            logger.info(f"{expert_name}: {backend.module.__class__.__name__}, {num_parameters} parameters")63 64        if not self.dht.is_alive():65            self.dht.run_in_background(await_ready=True)66 67        if self.module_backends:68            self.dht_handler_thread.start()69 70        if self.checkpoint_saver is not None:71            self.checkpoint_saver.start()72 73        for process in self.conn_handlers:74            if not process.is_alive():75                process.start()76            process.ready.result()77 78        try:79            self.runtime.run()80        finally:81            self.shutdown()82 83    # noinspection PyMethodOverriding84    @classmethod85    def create(86        cls,87        prefix: Optional[str],88        converted_model_name_or_path: str,89        num_blocks: Optional[int] = None,90        block_indices: Optional[str] = None,91        num_handlers: Optional[int] = None,92        min_batch_size: int = 1,93        max_batch_size: int = 4096,94        torch_dtype: str = "auto",95        cache_size_bytes: Optional[int] = None,96        device: Union[str, torch.device] = None,97        initial_peers: Sequence[str] = (),98        compression=CompressionType.NONE,99        stats_report_interval: Optional[int] = None,100        custom_module_path=None,101        update_period: float = 30,102        expiration: Optional[float] = None,103        use_auth_token: Optional[str] = None,104        *,105        start: bool,106        **kwargs,107    ) -> Server:108        """Create a server with one or more bloom blocks. See run_server.py for documentation."""109        if custom_module_path is not None:110            add_custom_models_from_file(custom_module_path)111        if prefix is None:112            prefix = converted_model_name_or_path113            assert UID_DELIMITER not in prefix and CHAIN_DELIMITER not in prefix, (114                f"Cannot use model name as prefix (contains '{UID_DELIMITER}' or '{CHAIN_DELIMITER}'); "115                f"Please specify --prefix manually when starting a server"116            )117            logger.info(f"Automatic dht prefix: {prefix}")118        assert (block_indices is None) != (num_blocks is None), "please specify num_blocks or block_indices, not both"119        dht = DHT(initial_peers=initial_peers, start=True, **kwargs)120        visible_maddrs_str = [str(a) for a in dht.get_visible_maddrs()]121        logger.info(f"Running DHT node on {visible_maddrs_str}, initial peers = {initial_peers}")122 123        device = device or ("cuda" if torch.cuda.is_available() else "cpu")124        memory_cache = MemoryCache(device, cache_size_bytes)125 126        if isinstance(torch_dtype, str):127            torch_dtype = DTYPE_MAP[torch_dtype]128        assert torch_dtype in DTYPE_MAP.values(), f"torch_dtype must be one of {list(DTYPE_MAP.values())}"129 130        if block_indices is not None:131            try:132                first_block_index, last_block_index = block_indices.split(":")133                first_block_index, last_block_index = map(int, map(str.strip, (first_block_index, last_block_index)))134            except Exception as e:135                logger.error(f"Failed to parse --block_indices ({e}), must be start:end (e.g. 0:18)")136                raise137            block_indices = range(first_block_index, last_block_index)138        else:139            assert num_blocks is not None140            block_indices = range(num_blocks)  # TODO replace with proper load balancing141 142        block_config = BloomConfig.from_pretrained(143            converted_model_name_or_path, use_auth_token=use_auth_token144        )145 146        # initialize modules147        blocks = {}148        for block_index in block_indices:149            module_uid = f"{prefix}.{block_index}"150            block = load_pretrained_block(151                converted_model_name_or_path,152                block_index,153                block_config,154                torch_dtype=torch_dtype,155                use_auth_token=use_auth_token,156            )157            for param in block.parameters():158                param.requires_grad = False159 160            blocks[module_uid] = TransformerBackend(161                module_uid,162                block,163                memory_cache=memory_cache,164                args_schema=(BatchTensorDescriptor(1, 2048, block_config.hidden_size, compression=compression),),165                kwargs_schema={},166                outputs_schema=(BatchTensorDescriptor(1, 2048, block_config.hidden_size, compression=compression),),167                min_batch_size=min_batch_size,168                max_batch_size=max_batch_size,169            )170 171        num_handlers = num_handlers if num_handlers is not None else len(blocks) * 4172 173        return cls(174            dht,175            blocks,176            num_connection_handlers=num_handlers,177            device=device,178            stats_report_interval=stats_report_interval,179            update_period=update_period,180            expiration=expiration,181            start=start,182        )183 184    def run_in_background(self, await_ready=True, timeout=None):185        """186        Starts Server in a background thread. if await_ready, this method will wait until background server187        is ready to process incoming requests or for :timeout: seconds max.188        """189        self.start()190        if await_ready and not self.ready.wait(timeout=timeout):191            raise TimeoutError("Server didn't notify .ready in {timeout} seconds")192 193    @property194    def ready(self) -> mp.synchronize.Event:195        """196        An event (multiprocessing.Event) that is set when the server is ready to process requests.197 198        Example199        =======200        >>> server.start()201        >>> server.ready.wait(timeout=10)202        >>> print("Server ready" if server.ready.is_set() else "Server didn't start in 10 seconds")203        """204        return self.runtime.ready  # mp.Event that is true if self is ready to process batches205 206    def shutdown(self):207        """208        Gracefully terminate the server, process-safe.209        Please note that terminating server otherwise (e.g. by killing processes) may result in zombie processes.210        If you did already cause a zombie outbreak, your only option is to kill them with -9 (SIGKILL).211        """212        self.ready.clear()213 214        for process in self.conn_handlers:215            process.terminate()216            process.join()217        logger.debug("Connection handlers terminated")218 219        if self.module_backends:220            self.dht_handler_thread.stop.set()221            self.dht_handler_thread.join()222 223        if self.checkpoint_saver is not None:224            self.checkpoint_saver.stop.set()225            self.checkpoint_saver.join()226 227        self.dht.shutdown()228        self.dht.join()229 230        logger.debug(f"Shutting down runtime")231 232        self.runtime.shutdown()233        logger.info("Server shutdown succesfully")234 235 236class ModuleAnnouncerThread(threading.Thread):237    """Periodically announces that this server hosts the specified modules, visible to all DHT peers"""238 239    def __init__(240        self, module_backends, dht: DHT, update_period: float = 30, expiration: Optional[int] = None, **kwargs241    ):242        super().__init__(**kwargs)243        if expiration is None:244            expiration = max(2 * update_period, MAX_DHT_TIME_DISCREPANCY_SECONDS)245        self.module_backends = module_backends246        self.dht = dht247        self.update_period = update_period248        self.expiration = expiration249        self.stop = threading.Event()250 251    def run(self) -> None:252        declare_active_modules(self.dht, self.module_backends.keys(), get_dht_time() + self.expiration)253        while not self.stop.wait(self.update_period):254            declare_active_modules(self.dht, self.module_backends.keys(), get_dht_time() + self.expiration)255