()
| 84 | return parser.parse_args() |
| 85 | |
| 86 | def init_distributed(): |
| 87 | world_size = int(os.environ.get("WORLD_SIZE", 1)) |
| 88 | local_rank = int(os.environ.get("LOCAL_RANK", 0)) |
| 89 | rank = int(os.environ.get("RANK", 0)) |
| 90 | print(f"Inference on multiple gpus, this gpu {local_rank}, rank {rank}, world_size {world_size}") |
| 91 | torch.cuda.set_device(local_rank) |
| 92 | dist.init_process_group("nccl") |
| 93 | return world_size, local_rank, rank |
| 94 | |
| 95 | def data_collator(batch): |
| 96 | ids = [] |