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