"""Train / evaluate learned radio-map models on RadioMapSeer.

Variants (identical U-Net backbone, only inputs differ):
  geom     [buildings, Tx one-hot]                        RadioUNet_C-style
  feats    [buildings, Tx, physics map]                   physics as an input feature
  hybrid   [buildings, Tx, physics map] + residual        network learns correction to the physics map
The physics map is B1 (straight-line multi-wall) or B2 (street-canyon geodesic), see baseline.py.

Options against painting dead areas as covered (off by default, so the PoC runs are unchanged):
  --dist     extra input: log distance to the Tx, so pixels beyond the receptive field (~170 m at 2 m/px) know
             their range even when the Tx pixel is out of sight (larger maps, --offtx windows)
  --gate     second output: logit of 'covered' (target above the floor), trained with BCE next to the MSE; the
             prediction is set to the floor where the logit is negative (unet_predict)
  --offtx P  with probability P a batch is cut to windows of --offtx-size px (at 1 m) that do not contain the Tx,
             so the network sees far-from-Tx areas with real labels (the Tx is only in the distance channel)

Physics-anchored hybrid (option A against far-field coverage):
  --variant hybrid --cap C [--cap-down D]   the output is the physics map plus a correction bounded to +C / -D dB
             (D defaults to C), and the physics
             channel is NOT clipped at the floor (it carries how far below -127 dB a pixel is, down to UNDER_DB), so
             where the physics says 'far below the floor' the network cannot paint coverage. Loss: MSE where the
             target is covered, hinge (only positive predictions penalised) where it is at the floor.

Protocol (RadioUNet, github.com/RonLevie/RadioUNet): official split, all 80 Tx, targets clipped at the
noise floor (gray 0.2, -127 dB) and rescaled, MSE, Adam 1e-4, batch 15, 50 epochs, lr x0.1 after 30,
best-on-val checkpoint. Inputs and targets are built on the device from uint8 memmaps.

Examples
  python -m rcm_ml.train --variant geom --epochs 50 --device mps
  python -m rcm_ml.train --variant hybrid --maps 50 --epochs 50 --equal-steps --seed 1 --device mps
"""
from __future__ import annotations

import argparse
import json
import math
import os
import queue
import subprocess
import threading
import time
import warnings
from pathlib import Path

import numpy as np
import torch
import torch.nn.functional as F

from .b2 import UNREACHED, turns_from_cache
from .baseline import FEATS, fit
from .data import DB_RANGE, N, P_TRNC, THRESH
from .models import build

ABOVE = 20.0 / DB_RANGE  # raw gray of -127 dB (the paper's SNR = 0 threshold)
UNDER_DB = 80.0  # --cap: the unclipped physics channel stops 80 dB below the floor (-207 dB)
DIST_SCALE = 3.5  # distance channel: 1 - log10(d / 1 m) / 3.5, i.e. 1 at the Tx, 0 at ~3.2 km


def dist_channel(rr, cc, r, c):
    """rr, cc: pixel coordinate grids (1 m pixels); r, c: Tx position (broadcastable). Works for any map size."""
    return 1.0 - torch.log10(torch.clamp(torch.hypot(rr - r, cc - c), min=1.0)) / DIST_SCALE


def git_commit() -> str:
    root = Path(__file__).resolve().parent.parent
    try:
        h = subprocess.check_output(["git", "rev-parse", "HEAD"], cwd=root, text=True).strip()
        dirty = subprocess.check_output(["git", "status", "--porcelain", "--", "rcm_ml"], cwd=root, text=True).strip() != ""
        return h + ("-dirty" if dirty else "")
    except (OSError, subprocess.CalledProcessError):
        return "unknown"


def pick_device(name: str) -> torch.device:
    if name == "mps":
        if not torch.backends.mps.is_available():
            raise SystemExit("MPS requested but torch.backends.mps.is_available() is False. Run natively on macOS, not in Docker/VM.")
        if os.environ.get("PYTORCH_ENABLE_MPS_FALLBACK") == "1":
            warnings.warn("PYTORCH_ENABLE_MPS_FALLBACK=1: unsupported ops will silently run on the CPU")
    return torch.device(name)


class Cache:
    """One split of the uint8 cache. Buildings are small (one per map) and live in RAM; the rest stays memmapped."""

    def __init__(self, cache: Path, name: str, target: str, maps_limit: int | None = None, phys: str = "B1"):
        meta = json.loads((cache / f"{name}_meta.json").read_text())
        self.maps = json.loads((cache / f"{name}_maps.json").read_text())
        map_idx = {m: i for i, m in enumerate(self.maps)}
        self.bld = np.load(cache / f"{name}_buildings.npy")
        self.arrs = {k: np.load(cache / f"{name}_{k}.npy", mmap_mode="r") for k in phys_arrays(phys)}
        self.tgt = np.load(cache / f"{name}_{target}.npy", mmap_mode="r")
        self.bidx = np.array([map_idx[m["map"]] for m in meta])
        self.rc = np.array([[m["r"], m["c"]] for m in meta], np.int64)
        self.map_of = np.array([m["map"] for m in meta])
        idx = np.arange(len(meta))
        if maps_limit:
            keep = set(self.maps[:maps_limit])  # nested subsets: the first k maps of the split order
            idx = idx[np.isin(self.map_of, list(keep))]
        self.idx = idx
        self.n_maps = len(set(self.map_of[idx].tolist()))

    def __len__(self):
        return len(self.idx)

    def batch(self, ids: np.ndarray) -> dict[str, torch.Tensor]:
        ids = np.sort(ids)  # sorted reads are friendlier to the memmaps
        out = {"bld": torch.from_numpy(self.bld[self.bidx[ids]]), "y": torch.from_numpy(self.tgt[ids]),
               "rc": torch.from_numpy(self.rc[ids])}
        for k, a in self.arrs.items():
            v = a[ids]
            out[k] = torch.from_numpy(v.astype(np.int32) if v.dtype == np.uint16 else v)
        return out


def phys_arrays(phys: str) -> list[str]:
    return {"B1": ["inside", "walls"], "B2": ["inside", "walls", "geod", "turns"]}[phys]


def batches(ds: Cache, batch: int, shuffle: bool, rng: np.random.Generator | None, drop_last: bool, prefetch: int = 6):
    """Background-thread prefetch of uint8 batches (memmap reads overlap with GPU work)."""
    order = rng.permutation(ds.idx) if shuffle else ds.idx
    stop = len(order) - (len(order) % batch if drop_last else 0)
    q: queue.Queue = queue.Queue(prefetch)

    def work():
        for k in range(0, stop, batch):
            q.put(ds.batch(order[k:k + batch]))
        q.put(None)

    threading.Thread(target=work, daemon=True).start()
    while (b := q.get()) is not None:
        yield b


class Featurizer:
    """uint8 batch on the device -> (inputs, target) in the training scale."""

    def __init__(self, device, coef, phys: str, variant: str, size: int, thresh: float, dist: bool = False, unclipped: bool = False):
        """unclipped: physics channel in the training scale without the floor clip (negative below -127 dB), for --cap."""
        self.device, self.variant, self.size, self.thresh, self.phys, self.dist = device, variant, size, thresh, phys, dist
        self.unclipped = unclipped
        self.coef = torch.tensor(np.asarray(coef), dtype=torch.float32, device=device)
        g = torch.arange(N, device=device, dtype=torch.float32)
        self.rr, self.cc = torch.meshgrid(g, g, indexing="ij")

    def physics_gray(self, b):
        return torch.clamp((self.physics_db(b) - P_TRNC) / DB_RANGE, 0.0, 1.0)

    def physics_db(self, b):
        r = b["rc"][:, 0, None, None].float()
        c = b["rc"][:, 1, None, None].float()
        d = torch.clamp(torch.hypot(self.rr - r, self.cc - c), min=1.0)
        ins = b["inside"].float()
        cols = {"one": torch.ones_like(d), "logd": torch.log10(d), "inside": ins, "walls": b["walls"].float(),
                "los": (ins == 0).float()}
        if self.phys == "B2":
            g = b["geod"].float()  # uint16 arrives as int32 (see Cache.batch)
            blocked = g == UNREACHED
            cols["loggeod"] = torch.where(blocked, torch.zeros_like(d), torch.log10(torch.clamp(g, min=1.0)))
            cols["turns"] = turns_from_cache(b["turns"].float())
            cols["blocked"] = blocked.float()
        return sum(self.coef[j] * cols[n] for j, n in enumerate(FEATS[self.phys]))

    def physics_unclipped(self, b):
        return unclipped_train_scale(self.physics_db(b), self.thresh)

    def to_train_scale(self, g):
        if self.thresh > 0:
            g = torch.clamp(g - self.thresh, min=0.0) / (1 - self.thresh)
        return g

    def __call__(self, b):
        b = {k: v.to(self.device, non_blocking=True) for k, v in b.items()}
        bld = b["bld"].float()[:, None]  # cached as 0/1
        B = bld.shape[0]
        tx = torch.zeros_like(bld)
        tx[torch.arange(B, device=self.device), 0, b["rc"][:, 0], b["rc"][:, 1]] = 1.0
        y_raw = b["y"].float()[:, None] / 255.0
        chans = [bld, tx]
        if self.variant != "geom":
            chans.append((self.physics_unclipped(b) if self.unclipped else self.to_train_scale(self.physics_gray(b)))[:, None])
        if self.dist:  # last, so the physics map stays channel 2 (hybrid residual)
            chans.append(dist_channel(self.rr, self.cc, b["rc"][:, 0, None, None].float(), b["rc"][:, 1, None, None].float())[:, None])
        x = torch.cat(chans, 1)
        y = self.to_train_scale(y_raw)
        if self.size != N:
            k = N // self.size
            parts = [F.avg_pool2d(x[:, 0:1], k), F.max_pool2d(x[:, 1:2], k)]
            if x.shape[1] > 2:
                parts.append(F.avg_pool2d(x[:, 2:], k))
            x, y, y_raw, bld = torch.cat(parts, 1), F.avg_pool2d(y, k), F.avg_pool2d(y_raw, k), F.avg_pool2d(bld, k)
        return x, y, y_raw, bld


class Metrics:
    """Streaming RMSEs. 'thr' = the paper's floor-clipped scale, 'raw' = the decoded PNG scale; both in dB."""

    def __init__(self, thresh: float):
        self.t = thresh
        self.s = {k: 0.0 for k in ("thr", "thr_out", "raw", "above")}
        self.n = dict.fromkeys(self.s, 0)
        self.fc = [0, 0, 0, 0]  # dead outdoor px painted covered, dead outdoor px, covered outdoor px painted dead, covered outdoor px

    def add(self, p_train, y_raw, bld):
        t = self.t
        p_raw = t + (1 - t) * p_train.clamp(0, 1) if t > 0 else p_train.clamp(0, 1)
        e_thr = (p_raw.clamp(min=THRESH) - y_raw.clamp(min=THRESH)) ** 2
        e_raw = (p_raw - y_raw) ** 2
        out = bld < 0.5
        ab = y_raw > ABOVE
        for k, e, m in (("thr", e_thr, None), ("thr_out", e_thr, out), ("raw", e_raw, None), ("above", e_raw, ab)):
            self.s[k] += float((e if m is None else e[m]).sum())
            self.n[k] += int(e.numel() if m is None else m.sum())
        dead, live, pc = out & (y_raw <= THRESH), out & (y_raw > THRESH), p_raw > THRESH
        self.fc[0] += int((dead & pc).sum()); self.fc[1] += int(dead.sum())
        self.fc[2] += int((live & ~pc).sum()); self.fc[3] += int(live.sum())

    def result(self):
        r = {k: (self.s[k] / max(self.n[k], 1)) ** 0.5 * DB_RANGE for k in self.s}
        return {"rmse_gray_paper": r["thr"] / DB_RANGE / (1 - THRESH), "rmse_db_thr": r["thr"], "rmse_db_thr_outdoor": r["thr_out"],
                "rmse_db_raw": r["raw"], "rmse_db_above_-127": r["above"],
                "dead_painted_covered_pct": round(100 * self.fc[0] / max(self.fc[1], 1), 2),
                "covered_painted_dead_pct": round(100 * self.fc[2] / max(self.fc[3], 1), 2)}


def unclipped_train_scale(db, thresh: float):
    """Path gain in dB -> training scale (0 = floor, 1 = -47.8 dB), negative below the floor down to -UNDER_DB. Works on
    tensors and numpy arrays."""
    t = ((db - P_TRNC) / DB_RANGE - thresh) / (1 - thresh)
    lo, hi = -UNDER_DB / ((1 - thresh) * DB_RANGE), 1.0
    return t.clamp(lo, hi) if isinstance(t, torch.Tensor) else np.clip(t, lo, hi)


def cap_train_scale(cap_db: float, thresh: float) -> float:
    return cap_db / ((1 - thresh) * DB_RANGE)


def unet_predict(model, x, gate: bool):
    """Prediction in the training scale (0 = floor). With the gate, pixels whose 'covered' logit is negative are
    set to the floor."""
    p = model(x)
    v = p[:, :1]
    if gate:
        v = torch.where(p[:, 1:2] > 0, v, torch.zeros_like(v))
    return v


def load_unet(run, dev):
    """(featurizer, model in eval mode, config) of a trained run directory."""
    run = Path(run)
    cfg = json.loads((run / "config.json").read_text())
    cap = cfg.get("cap", 0.0)
    feat = Featurizer(dev, cfg["phys_coef"], cfg["phys"], cfg["variant"], cfg["size"], cfg["thresh"], dist=cfg.get("dist", False), unclipped=cap > 0)
    net = build(cfg.get("arch", "unet"), in_channels(cfg["variant"], cfg.get("dist", False)), cfg["base"],
                residual=(cfg["variant"] == "hybrid"), out_ch=2 if cfg.get("gate") else 1, cap=cap_train_scale(cap, cfg["thresh"]),
                cap_down=cap_train_scale(cfg.get("cap_down") or cap, cfg["thresh"])).to(dev)
    net.load_state_dict(torch.load(run / "best.pt", map_location=dev))
    return feat, net.eval(), cfg


def in_channels(variant: str, dist: bool) -> int:
    return (2 if variant == "geom" else 3) + int(dist)


def offtx_crop(x, y, rc, k: int, w: int, rng):
    """Cut every sample to a w x w window (model pixels) that does not contain its Tx (rc in 1 m pixels)."""
    n = x.shape[-1]
    xs, ys = [], []
    for i in range(len(x)):
        tr, tc = int(rc[i, 0]) // k, int(rc[i, 1]) // k
        for _ in range(50):
            r0, c0 = rng.integers(0, n - w + 1, 2)
            if not (r0 <= tr < r0 + w and c0 <= tc < c0 + w):
                break
        xs.append(x[i, :, r0:r0 + w, c0:c0 + w]); ys.append(y[i, :, r0:r0 + w, c0:c0 + w])
    return torch.stack(xs), torch.stack(ys)


@torch.no_grad()
def evaluate(model, ds: Cache, feat: Featurizer, amp: bool, batch: int = 32, channels_last: bool = False, gate: bool = False):
    model.eval()
    M = Metrics(feat.thresh)
    for b in batches(ds, batch, False, None, False):
        x, _, y_raw, bld = feat(b)
        if channels_last:
            x = x.contiguous(memory_format=torch.channels_last)
        with torch.autocast(feat.device.type, dtype=torch.float16, enabled=amp):
            p = unet_predict(model, x, gate)
        M.add(p.float(), y_raw, bld)
    model.train()
    return M.result()


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--cache", default="data/rms_cache")
    ap.add_argument("--variant", choices=["geom", "feats", "hybrid"], required=True)
    ap.add_argument("--phys", choices=["B1", "B2"], default="B1")
    ap.add_argument("--target", default="dpm", choices=["dpm", "irt2"])
    ap.add_argument("--eval-targets", default="dpm,irt2", help="test targets evaluated at the end")
    ap.add_argument("--size", type=int, default=256)
    ap.add_argument("--arch", choices=["unet", "radiounet"], default="unet")
    ap.add_argument("--base", type=int, default=32)
    ap.add_argument("--batch", type=int, default=15)
    ap.add_argument("--lr", type=float, default=1e-4)
    ap.add_argument("--epochs", type=float, default=50)
    ap.add_argument("--equal-steps", action="store_true",
                    help="epochs are counted over the full training set, so --maps 50 gets as many steps as 501 maps")
    ap.add_argument("--lr-drop", type=float, default=0.6, help="fraction of steps after which lr x0.1 (RadioUNet: 30 of 50 epochs)")
    ap.add_argument("--thresh", type=float, default=THRESH, help="noise-floor clip of the targets (0 = raw targets)")
    ap.add_argument("--minutes", type=float, default=None, help="wall-clock budget (benchmarks)")
    ap.add_argument("--val-every", type=int, default=None, help="steps between validations (default: one full-set epoch)")
    ap.add_argument("--maps", type=int, default=None, help="limit number of training maps")
    ap.add_argument("--seed", type=int, default=0)
    ap.add_argument("--amp", action="store_true", help="float16 autocast")
    ap.add_argument("--channels-last", action="store_true")
    ap.add_argument("--device", default="mps")
    ap.add_argument("--threads", type=int, default=8)
    ap.add_argument("--out", default="runs/mps")
    ap.add_argument("--name", default=None)
    ap.add_argument("--dist", action="store_true", help="extra input: log distance to the Tx")
    ap.add_argument("--gate", action="store_true", help="second output: 'covered' logit (BCE), prediction floored where negative")
    ap.add_argument("--gate-weight", type=float, default=0.02, help="BCE weight next to the MSE")
    ap.add_argument("--offtx", type=float, default=0.0, help="probability that a batch is cut to windows without the Tx")
    ap.add_argument("--offtx-size", type=int, default=128, help="window side in 1 m pixels")
    ap.add_argument("--cap", type=float, default=0.0, help="hybrid: bound the correction to +/- this many dB, unclipped physics input, hinge loss")
    ap.add_argument("--cap-down", type=float, default=None, help="bound for downward corrections, dB (default: --cap)")
    a = ap.parse_args()
    assert a.cap == 0 or (a.variant == "hybrid" and a.thresh > 0), "--cap needs --variant hybrid and the floor-clipped targets"
    t_start = time.time()
    torch.set_num_threads(a.threads)
    torch.manual_seed(a.seed)
    rng = np.random.default_rng(a.seed)
    dev = pick_device(a.device)
    cache = Path(a.cache)
    tr = Cache(cache, "train", a.target, a.maps, a.phys)
    va = Cache(cache, "val", a.target, None, a.phys)
    coef = fit(cache, a.phys, a.target, maps_limit=a.maps) if a.variant != "geom" else np.zeros(len(FEATS[a.phys]))
    feat = Featurizer(dev, coef, a.phys, a.variant, a.size, a.thresh, dist=a.dist, unclipped=a.cap > 0)
    in_ch = in_channels(a.variant, a.dist)
    model = build(a.arch, in_ch, a.base, residual=(a.variant == "hybrid"), out_ch=2 if a.gate else 1,
                  cap=cap_train_scale(a.cap, a.thresh), cap_down=cap_train_scale(a.cap_down or a.cap, a.thresh)).to(dev)
    if a.channels_last:
        model = model.to(memory_format=torch.channels_last)
    opt = torch.optim.Adam(model.parameters(), a.lr)
    full = len(Cache(cache, "train", a.target, None, a.phys)) if a.equal_steps else len(tr)
    steps_per_epoch = full // a.batch
    total = int(round(a.epochs * steps_per_epoch))
    val_every = a.val_every or steps_per_epoch
    sched = torch.optim.lr_scheduler.StepLR(opt, step_size=max(1, int(a.lr_drop * total)), gamma=0.1)
    scaler = torch.amp.GradScaler(dev.type, enabled=a.amp)
    name = a.name or f"{a.variant}{'' if a.phys == 'B1' or a.variant == 'geom' else '-' + a.phys}_{a.target}_s{a.size}_m{tr.n_maps}_seed{a.seed}"
    out = Path(a.out) / name
    out.mkdir(parents=True, exist_ok=True)
    cfg = {**vars(a), "train_maps": tr.n_maps, "train_samples": len(tr), "total_steps": total, "val_every": val_every,
           "phys_coef": [round(float(c), 4) for c in coef], "commit": git_commit(), "torch": torch.__version__}
    (out / "config.json").write_text(json.dumps(cfg, indent=1))
    print(json.dumps(cfg), flush=True)

    log, best, best_step, step, seen = [], math.inf, 0, 0, 0
    t0 = time.time()
    done = False
    while not done:
        for b in batches(tr, a.batch, True, rng, True):
            x, y, _, _ = feat(b)
            if a.offtx and rng.random() < a.offtx:
                x, y = offtx_crop(x, y, b["rc"], N // a.size, a.offtx_size * a.size // N, rng)
            if a.channels_last:
                x = x.contiguous(memory_format=torch.channels_last)
            with torch.autocast(dev.type, dtype=torch.float16, enabled=a.amp):
                p = model(x)
            if a.cap:  # the floor is 0: below it any prediction <= 0 is right
                v = p[:, :1].float()
                loss = torch.where(y > 0, (v - y) ** 2, F.relu(v) ** 2).mean()
            else:
                loss = F.mse_loss(p[:, :1].float(), y)
            if a.gate:
                loss = loss + a.gate_weight * F.binary_cross_entropy_with_logits(p[:, 1:2].float(), (y > 0).float())
            opt.zero_grad(set_to_none=True)
            scaler.scale(loss).backward()
            scaler.step(opt)
            scaler.update()
            sched.step()
            step += 1
            seen += len(x)
            timeout = a.minutes and time.time() - t0 > a.minutes * 60
            if step % val_every == 0 or step >= total or timeout:
                v = evaluate(model, va, feat, a.amp, channels_last=a.channels_last, gate=a.gate)
                rec = {"step": step, "epoch": round(step / steps_per_epoch, 2), "t_min": round((time.time() - t0) / 60, 2),
                       "samples": seen, "train_loss": loss.item(), "lr": opt.param_groups[0]["lr"],
                       **{f"val_{k}": round(val, 5) for k, val in v.items()}}
                log.append(rec)
                print(json.dumps(rec), flush=True)
                if v["rmse_db_thr"] < best:
                    best, best_step = v["rmse_db_thr"], step
                    torch.save(model.state_dict(), out / "best.pt")
                (out / "log.json").write_text(json.dumps(log))
            if step >= total or timeout:
                done = True
                break
    train_min = (time.time() - t0) / 60
    model.load_state_dict(torch.load(out / "best.pt", map_location=dev))
    res = {"variant": a.variant, "phys": a.phys if a.variant != "geom" else None, "target": a.target, "seed": a.seed,
           "size": a.size, "arch": a.arch, "base": a.base, "train_maps": tr.n_maps, "steps": step, "best_step": best_step,
           "epochs_full_set_equiv": round(seen / full, 2), "samples_seen": seen, "train_minutes": round(train_min, 1),
           "params_M": round(sum(p.numel() for p in model.parameters()) / 1e6, 2), "commit": cfg["commit"], "config": cfg,
           "val_best_rmse_db_thr": best, "test": {}}
    for tgt in a.eval_targets.split(","):
        te = Cache(cache, "test", tgt, None, a.phys)
        t1 = time.time()
        res["test"][tgt] = evaluate(model, te, feat, a.amp, channels_last=a.channels_last, gate=a.gate)
        res["test"][tgt]["ms_per_map"] = round((time.time() - t1) / len(te) * 1000, 2)
    res["wall_minutes"] = round((time.time() - t_start) / 60, 1)
    (out / "test.json").write_text(json.dumps(res, indent=1))
    print("TEST", json.dumps(res["test"]), flush=True)


if __name__ == "__main__":
    main()
