(self, config)
| 123 | return True |
| 124 | |
| 125 | def init(self, config): |
| 126 | self._sync_runtime_config(config) |
| 127 | self.encoder_engine_rank = int(self.config.get("encoder_engine_rank", self.encoder_engine_rank)) |
| 128 | self.transformer_engine_rank = int(self.config.get("transformer_engine_rank", self.transformer_engine_rank)) |
| 129 | self.decoder_engine_rank = int(self.config.get("decoder_engine_rank", self.decoder_engine_rank)) |
| 130 | shared_slots = int(self.config.get("rdma_buffer_slots", self._phase2_slots)) |
| 131 | shared_slot_size = int(self.config.get("rdma_buffer_slot_size", 4096)) |
| 132 | self._phase2_server_ip = str(self.config.get("rdma_phase2_host", self._phase2_server_ip)) |
| 133 | self._phase2_handshake_port = int(self.config.get("rdma_phase2_handshake_port", self._phase2_handshake_port)) |
| 134 | self._phase2_slots = shared_slots |
| 135 | self._phase2_slot_size = shared_slot_size |
| 136 | |
| 137 | if "seed" in self.config: |
| 138 | seed_all(self.config["seed"]) |
| 139 | |
| 140 | data_bootstrap_addr = self.config.get("data_bootstrap_addr", "127.0.0.1") |
| 141 | data_bootstrap_room = self.config.get("data_bootstrap_room", 0) |
| 142 | |
| 143 | if data_bootstrap_addr is None or data_bootstrap_room is None: |
| 144 | return |
| 145 | |
| 146 | if str(os.getenv("IS_CENTRALIZED", "0")).strip().lower() not in {"1", "true", "yes", "on"}: |
| 147 | try: |
| 148 | self._ensure_phase2_request_buffer() |
| 149 | except Exception: |
| 150 | self.logger.exception("Failed to connect phase2 RDMA buffer, will retry") |
| 151 | |
| 152 | buffer_sizes = estimate_transformer_buffer_sizes(self.config) |
| 153 | request = AllocationRequest( |
| 154 | bootstrap_room=data_bootstrap_room, |
| 155 | buffer_sizes=buffer_sizes, |
| 156 | ) |
| 157 | handle = self.alloc_memory(request) |
| 158 | data_ptrs = [buf.addr for buf in handle.buffers] |
| 159 | data_lens = [buf.nbytes for buf in handle.buffers] |
| 160 | data_args = DataArgs( |
| 161 | sender_engine_rank=self.transformer_engine_rank, |
| 162 | receiver_engine_rank=self.decoder_engine_rank, |
| 163 | data_ptrs=data_ptrs, |
| 164 | data_lens=data_lens, |
| 165 | data_item_lens=data_lens, |
| 166 | ib_device=None, |
| 167 | ) |
| 168 | self.data_mgr.init(data_args, data_bootstrap_room) |
| 169 | phase2_bootstrap_addr = str(self.config.get("transformer_node_address", data_bootstrap_addr)) |
| 170 | self.data_receiver[data_bootstrap_room] = DataReceiver(self.data_mgr, phase2_bootstrap_addr, data_bootstrap_room) |
| 171 | self.data_receiver[data_bootstrap_room].init() |
| 172 | |
| 173 | def load_models(self): |
| 174 | self.logger.info("Loading Decoder Models...") |
no test coverage detected