(self, config)
| 363 | raise RuntimeError(f"Failed to produce phase1 RDMA request for room={room} after {retries} attempts") from last_exc |
| 364 | |
| 365 | def init(self, config): |
| 366 | self._sync_runtime_config(config) |
| 367 | self.encoder_engine_rank = int(self.config.get("encoder_engine_rank", self.encoder_engine_rank)) |
| 368 | self.transformer_engine_rank = int(self.config.get("transformer_engine_rank", self.transformer_engine_rank)) |
| 369 | self.decoder_engine_rank = int(self.config.get("decoder_engine_rank", self.decoder_engine_rank)) |
| 370 | shared_slots = int(self.config.get("rdma_buffer_slots", self._request_slots)) |
| 371 | shared_slot_size = int(self.config.get("rdma_buffer_slot_size", 4096)) |
| 372 | self._request_server_ip = str(self.config.get("rdma_request_host", self._request_server_ip)) |
| 373 | self._request_handshake_port = int(self.config.get("rdma_request_handshake_port", self._request_handshake_port)) |
| 374 | self._request_slots = shared_slots |
| 375 | self._request_slot_size = shared_slot_size |
| 376 | self._phase1_server_ip = str(self.config.get("rdma_phase1_host", self._phase1_server_ip)) |
| 377 | self._phase1_handshake_port = int(self.config.get("rdma_phase1_handshake_port", self._phase1_handshake_port)) |
| 378 | self._phase1_slots = shared_slots |
| 379 | self._phase1_slot_size = shared_slot_size |
| 380 | |
| 381 | # Seed everything if seed is in config |
| 382 | if "seed" in self.config: |
| 383 | seed_all(self.config["seed"]) |
| 384 | |
| 385 | data_bootstrap_addr = self.config.get("data_bootstrap_addr", "127.0.0.1") |
| 386 | data_bootstrap_room = self.config.get("data_bootstrap_room", 0) |
| 387 | |
| 388 | phase1_deadline = time.time() + 30.0 |
| 389 | while self._phase1_rdma_buffer is None and time.time() < phase1_deadline: |
| 390 | try: |
| 391 | self._ensure_phase1_meta_buffer() |
| 392 | except Exception: |
| 393 | self.logger.exception("Failed to connect phase1 RDMA buffer, will retry") |
| 394 | if self._phase1_rdma_buffer is None: |
| 395 | time.sleep(0.1) |
| 396 | |
| 397 | if self._phase1_rdma_buffer is None: |
| 398 | raise RuntimeError("phase1 RDMA buffer is not ready") |
| 399 | |
| 400 | if data_bootstrap_addr is None or data_bootstrap_room is None: |
| 401 | return |
| 402 | |
| 403 | buffer_sizes = estimate_encoder_buffer_sizes(self.config) |
| 404 | request = AllocationRequest( |
| 405 | bootstrap_room=data_bootstrap_room, |
| 406 | buffer_sizes=buffer_sizes, |
| 407 | ) |
| 408 | handle = self.alloc_memory(request) |
| 409 | data_ptrs = [buf.addr for buf in handle.buffers] |
| 410 | data_lens = [buf.nbytes for buf in handle.buffers] |
| 411 | data_args = DataArgs( |
| 412 | sender_engine_rank=self.encoder_engine_rank, |
| 413 | receiver_engine_rank=self.transformer_engine_rank, |
| 414 | data_ptrs=data_ptrs, |
| 415 | data_lens=data_lens, |
| 416 | data_item_lens=data_lens, |
| 417 | ib_device=None, |
| 418 | ) |
| 419 | self.data_mgr.init(data_args, data_bootstrap_room) |
| 420 | self.data_sender[data_bootstrap_room] = DataSender(self.data_mgr, data_bootstrap_addr, data_bootstrap_room) |
| 421 | |
| 422 | def load_models(self): |
no test coverage detected