#!/usr/bin/env python3
"""Test przepustowości treningu MicroPLLM: 1 karta vs 3 karty DDP (asymetryczne accum).

Mierzy tok/s na realnym forward+backward (model 108M z pre-trainu, dane z train.bin),
bez zapisu checkpointów — czysty benchmark.

Uruchomienie:
  1 karta:  nix-shell shell.nix --run "python train_ddp_test.py --gpu 0 --steps 30"
  3 karty:  nix-shell shell.nix --run "torchrun --nproc_per_node=3 train_ddp_test.py --steps 30"

Dokładność pomiaru: pomija pierwsze 5 kroków (CUDA warmup), raportuje tok/s
z reszty. Zgodność DDP: na końcu weryfikuje, że wagi na rank!=0 zgadzają się
z rank 0 (all-reduce max abs diff) — pewność, że gradienty faktycznie spływają.
"""
import argparse
import os
import time

os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")

import numpy as np
import torch
import torch.distributed as dist

BASE = "/media/disk5/micro_pl_lm"
SEQ_LEN = 2048
VOCAB = 32000

ARCH = {
    "vocab_size": VOCAB,
    "d_model": 768,
    "n_layers": 16,
    "n_heads": 8,
    "mlp_layers": 4,
    "d_mlp": 512,
    "max_seq_len": 2048,
}

# asymetryczne accum per rank (równe czasy przy 30k/21k/15k tok/s)
ACCUM_PER_RANK = {0: 10, 1: 7, 2: 5}
MICRO_BATCH = 8


def setup_distributed():
    """Init DDP jeśli uruchomiono przez torchrun; zwraca (rank, world_size, device)."""
    if "RANK" in os.environ and "WORLD_SIZE" in os.environ:
        rank = int(os.environ["RANK"])
        world = int(os.environ["WORLD_SIZE"])
        local = int(os.environ.get("LOCAL_RANK", rank))
        torch.cuda.set_device(local)
        dist.init_process_group("nccl")
        return rank, world, f"cuda:{local}"
    return 0, 1, "cuda:0"


def cleanup_distributed():
    if dist.is_initialized():
        dist.destroy_process_group()


def get_batch(tokens, rng, batch_size, device):
    n = len(tokens)
    ix = rng.integers(0, n - SEQ_LEN - 1, size=batch_size)
    x = np.stack([tokens[i:i + SEQ_LEN].astype(np.int64) for i in ix])
    y = np.stack([tokens[i + 1:i + SEQ_LEN + 1].astype(np.int64) for i in ix])
    return (
        torch.from_numpy(x).pin_memory().to(device, non_blocking=True),
        torch.from_numpy(y).pin_memory().to(device, non_blocking=True),
    )


def main():
    ap = argparse.ArgumentParser(description="DDP throughput test MicroPLLM")
    ap.add_argument("--gpu", type=int, default=0, help="GPU (tylko tryb 1-kartowy)")
    ap.add_argument("--steps", type=int, default=30, help="kroki pomiarowe (+5 warmup)")
    ap.add_argument("--micro_batch", type=int, default=MICRO_BATCH)
    ap.add_argument("--asymmetric", action="store_true", default=True,
                    help="accum proporcjonalny do prędkości karty (10/7/5)")
    ap.add_argument("--equal", action="store_true", help="równe accum (porównanie: czekanie na najwolniejszą)")
    ap.add_argument("--equal_mb", action="store_true", help="identyczny micro_batch na wszystkich kartach")
    ap.add_argument("--data", default=f"{BASE}/data/train.bin")
    ap.add_argument("--micro", action="store_true", help="mikro-test: 4 warstwy zamiast 16 (szybka weryfikacja DDP)")
    args = ap.parse_args()

    rank, world, device = setup_distributed()
    is_ddp = world > 1

    if args.micro:
        ARCH["n_layers"] = 4

    import sys
    sys.path.insert(0, BASE)
    from model import MicroPLLM

    torch.manual_seed(42 + rank)  # różne porcje danych per rank
    rng = np.random.default_rng(42 + rank)

    model = MicroPLLM(**ARCH).to(device)
    opt = torch.optim.AdamW(model.parameters(), lr=1e-4)

    if is_ddp:
        # gradient as bucket view + broadcast wag z rank0 (start z identycznych wag)
        for p in model.parameters():
            dist.broadcast(p.data, src=0)
        ddp_model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[int(device.split(':')[-1])])
    else:
        ddp_model = model

    # DDP wymaga TEJ SAMEJ liczby backwardów na krok na każdym ranku —
    # asymetrię realizujemy RÓŻNYM micro_batchem, nie accum.
    # (przy asymetrycznym accum DDP zawiesza się na all-reduce)
    accum = 2  # identyczny dla wszystkich ranków
    # micro_batch według WOLNEGO VRAM-u i prędkości (LM Studio zjada 10GB GPU0):
    # VRAM wolny: 13.8/14.9/11.2GB; logits+grad = micro×2048×32k×4B
    #   micro 8 → 2.1GB, micro 6 → 1.6GB; reszta (wagi+opt+aktywacje) ~5-6GB
    MB_PER_RANK = {0: 8, 1: 8, 2: 6}
    micro_batch = MB_PER_RANK.get(rank, 16) if not args.equal_mb else args.micro_batch
    args.micro_batch = micro_batch

    tokens = np.memmap(args.data, dtype=np.uint16, mode="r")
    n_params = sum(p.numel() for p in model.parameters())

    if rank == 0:
        mb_list = [MB_PER_RANK.get(r, 16) if not args.equal_mb else args.micro_batch for r in range(world)]
        mode = f"DDP×{world} (micro/rank: {'/'.join(map(str, mb_list))})" if is_ddp else "SOLO (1 karta)"
        arch = " (MICRO 4 warstwy)" if args.micro else ""
        print(f"=== BENCHMARK: {mode}{arch} ===", flush=True)
        print(f"parametry: {n_params:,} | accum {accum} × seq {SEQ_LEN}", flush=True)
        if is_ddp:
            world_batch = sum(mb_list) * SEQ_LEN * accum
            print(f"światowy batch/krok: {world_batch:,} tok", flush=True)

    WARMUP = 5
    step_times = []
    tokens_per_step = args.micro_batch * SEQ_LEN * accum  # tok liczone przez TĘ kartę

    model.train()
    global_step = 0
    while global_step < args.steps + WARMUP:
        t0 = time.time()
        opt.zero_grad(set_to_none=True)
        for micro in range(accum):
            x, y = get_batch(tokens, rng, args.micro_batch, device)
            with torch.autocast("cuda", dtype=torch.bfloat16):
                h = model.tok_emb(x)
                for block in model.blocks:
                    h = torch.utils.checkpoint.checkpoint(
                        lambda b, hh: b(hh), block, h, use_reentrant=False)
                logits = model.lm_head(model.ln_f(h))
                loss = torch.nn.functional.cross_entropy(
                    logits.reshape(-1, logits.size(-1)).float(), y.reshape(-1))
            loss.backward()
        # DDP uśrednia gradienty przy każdym backward; accum skaluje loss
        # (uproszczenie do testu — w produkcji loss/accum przed backward)
        gnorm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
        opt.step()
        if is_ddp:
            dist.barrier()  # równy start następnego kroku
        dt = time.time() - t0
        if global_step >= WARMUP:
            step_times.append(dt)
        global_step += 1
        if rank == 0 and (global_step % 10 == 0 or global_step == args.steps + WARMUP):
            print(f"  krok {global_step}/{args.steps + WARMUP}: {dt:.2f}s", flush=True)

    # === POMIAR ===
    if is_ddp:
        # toks liczone przez wszystkich ranków w kroku (różne micro per rank):
        mb_list = [MB_PER_RANK.get(r, 16) if not args.equal_mb else args.micro_batch for r in range(world)]
        toks_all = sum(mb_list) * SEQ_LEN * accum
    else:
        toks_all = tokens_per_step

    med = sorted(step_times)[len(step_times)//2]
    tps = toks_all / med

    # === Weryfikacja zgodności wag (czy DDP naprawdę synchronizuje) ===
    sync_ok = True
    if is_ddp:
        with torch.no_grad():
            w0 = model.blocks[0].attn.wq.weight
            max_diff = w0.clone()
            dist.all_reduce(max_diff, op=dist.ReduceOp.MAX)
            min_diff = w0.clone()
            dist.all_reduce(min_diff, op=dist.ReduceOp.MIN)
            spread = (max_diff - min_diff).abs().max().item()
        sync_ok = spread == 0.0
        if rank == 0:
            print(f"synchronizacja wag: spread={spread:.2e} → {'OK (identyczne)' if sync_ok else 'ROZJAZD!'}", flush=True)

    if rank == 0:
        print(f"\n=== WYNIK ===", flush=True)
        print(f"krok (mediana): {med*1000:.0f} ms | tok/krok (świat): {toks_all:,}", flush=True)
        print(f"PRZEPŁYWOWOŚĆ: {tps:,.0f} tok/s", flush=True)
        print(f"przyśpieszenie vs 3090 solo (30k): {tps/30000:.2f}x", flush=True)
        print(f"projekcja chinchilla 2.8B tok: {2.8e9/tps/3600:.1f} h", flush=True)

    cleanup_distributed()


if __name__ == "__main__":
    main()
