(self)
| 137 | self.num_steps = 0 |
| 138 | |
| 139 | def _init_workers(self): |
| 140 | # Lazy import the Worker to avoid importing torch.cuda/xformers |
| 141 | # before CUDA_VISIBLE_DEVICES is set in the Worker |
| 142 | from TD_Pipe.worker.worker import Worker |
| 143 | |
| 144 | assert self.parallel_config.world_size == 1, ( |
| 145 | "Ray is required if parallel_config.world_size > 1.") |
| 146 | |
| 147 | self.workers: List[Worker] = [] |
| 148 | distributed_init_method = f"tcp://{get_ip()}:{get_open_port()}" |
| 149 | self.driver_worker = Worker( |
| 150 | self.model_config, |
| 151 | self.parallel_config, |
| 152 | self.scheduler_config, |
| 153 | local_rank=0, |
| 154 | rank=0, |
| 155 | distributed_init_method=distributed_init_method, |
| 156 | is_driver_worker=True, |
| 157 | ) |
| 158 | self._run_workers("init_model") |
| 159 | self._run_workers("load_model") |
| 160 | |
| 161 | def _init_workers_ray(self, placement_group: "PlacementGroup", |
| 162 | **ray_remote_kwargs): |
no test coverage detected