()
| 174 | |
| 175 | |
| 176 | def initialize_distributed() -> None: |
| 177 | rank = int(os.getenv("RANK", "0")) |
| 178 | local_rank = int(os.getenv("LOCAL_RANK", "0")) |
| 179 | world_size = int(os.getenv("WORLD_SIZE", "1")) |
| 180 | |
| 181 | torch.cuda.set_device(local_rank) |
| 182 | dist.init_process_group(backend="nccl", rank=rank, world_size=world_size) |
| 183 | logger.info(f"🔢 Initialized process group; rank: {rank}, size: {world_size}") |
| 184 | return local_rank |
| 185 | |
| 186 | |
| 187 | def get_device(device_spec: Union[str, int, List[int]]) -> torch.device: |