import argparse from datetime import timedelta import os import time import torch import torch.distributed as dist def parse_args(): parser = argparse.ArgumentParser(description="Distributed PyTorch test script") parser.add_argument("--master_addr", type=str, required=True, help="Master node address") parser.add_argument("--master_port", type=str, required=True, help="Master node port") return parser.parse_args() if __name__ == "__main__": args = parse_args() # Set environment variables based on parsed arguments os.environ["MASTER_ADDR"] = args.master_addr os.environ["MASTER_PORT"] = args.master_port os.environ["RANK"] = os.environ["SLURM_PROCID"] os.environ["LOCAL_RANK"] = os.environ["SLURM_LOCALID"] rank = int(os.environ["RANK"]) local_rank = int(os.environ["LOCAL_RANK"]) world_size = int(os.environ["SLURM_JOB_NUM_NODES"]) * torch.cuda.device_count() print(f"Initializing process group on rank {rank}") torch.cuda.set_device(local_rank) dist.init_process_group( backend="nccl", timeout=timedelta(minutes=2), rank=rank, world_size=world_size, device_id=torch.device(f"cuda:{local_rank}"), ) print(f"Done initializing on rank: {dist.get_rank()}") dist.barrier() if dist.get_rank() == 0: time.sleep(2) print(f"DONE! world: {dist.get_world_size()}") dist.barrier() dist.destroy_process_group()