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()
