"""Runs only inside the caller's existing pinned PyTorch runtime image."""
import datetime
import json
import os
from pathlib import Path
import statistics
import sys
import time

import torch
import torch.distributed as dist
import torch.multiprocessing as mp


def worker(rank, max_mib, iterations):
    torch.set_num_threads(1)
    torch.cuda.set_device(rank)
    properties = torch.cuda.get_device_properties(rank)
    raw_uuid = str(properties.uuid)
    gpu_uuid = raw_uuid if raw_uuid.startswith('GPU-') else 'GPU-' + raw_uuid
    metadata = dict(rank=rank, device_index=rank, uuid=gpu_uuid, name=properties.name,
                    torch_version=torch.__version__, torch_cuda_build_version=torch.version.cuda,
                    torch_nccl_build_version=list(torch.cuda.nccl.version()))
    print('ANVIL_NCCL_RANK ' + json.dumps(metadata), flush=True)
    dist.init_process_group('nccl', init_method='file:///tmp/anvil-nccl-rendezvous',
                            rank=rank, world_size=2,
                            timeout=datetime.timedelta(seconds=20),
                            device_id=torch.device('cuda', rank))
    sizes = sorted(set([2, 256, 16384, max_mib * 1024 * 1024 // 4]))
    measurements = []
    for count in sizes:
        indices = torch.arange(count, dtype=torch.float32, device=rank)
        samples = []
        for iteration in range(iterations + 3):
            base = (indices + iteration) % 97
            tensor = base + rank + 1
            expected = base * 2 + 3
            torch.cuda.synchronize()
            start = time.perf_counter()
            dist.all_reduce(tensor, op=dist.ReduceOp.SUM)
            torch.cuda.synchronize()
            elapsed = time.perf_counter() - start
            if not torch.equal(tensor, expected):
                raise RuntimeError(f'all-reduce corruption on rank {rank}, size {count}, iteration {iteration}')
            if iteration >= 3:
                samples.append(elapsed)
        median = statistics.median(samples)
        measurements.append(dict(bytes=count * 4, iterations=iterations, correct=True,
                                 median_ms=median * 1000, min_ms=min(samples) * 1000,
                                 algorithmic_gbps=count * 4 / median / 1e9))
        del indices, base, tensor, expected
    dist.destroy_process_group()
    Path(f'/tmp/anvil-nccl-rank-{rank}.json').write_text(json.dumps({
        **metadata, 'correct': True, 'measurements': measurements,
        'allocated_peak_bytes': torch.cuda.max_memory_allocated(rank),
    }))


if __name__ == '__main__':
    if torch.cuda.device_count() != 2:
        raise RuntimeError('probe requires exactly two exposed GPUs')
    mp.spawn(worker, args=(int(sys.argv[1]), int(sys.argv[2])), nprocs=2, join=True)
    ranks = [json.loads(Path(f'/tmp/anvil-nccl-rank-{rank}.json').read_text()) for rank in range(2)]
    print('ANVIL_NCCL_RESULT ' + json.dumps({'ranks': ranks}), flush=True)
