"""Cell-site selection on RadioMapSeer test maps: how well do fast predictors pick transmitter sites?

Each test map has 80 candidate transmitter positions with WinProp maps (DPM and IRT2). Coverage of a site = share of
outdoor pixels whose path gain is above a threshold. For every method we predict all 80 maps, then
  - one site: pick the site with the highest predicted coverage; regret = true coverage of the best site minus true
    coverage of the pick (percentage points); also whether the pick is the true best / in the true top 5, and the rank
    correlation of predicted vs true coverage over the 80 sites;
  - three sites: greedy choice maximising predicted joint coverage (a pixel is covered if any chosen site covers it);
    regret against greedy on the true maps.
Truth: DPM and IRT2 in turn. Methods: the other WinProp simulator, physics maps B1/B2, our calibrated ray tracer,
and the PoC U-Nets (128 px, upsampled). With --sionna-dir: tuned Sionna (rough walls, angle-dependent diffraction,
optionally + our tracer's two-corner paths) on an 8 m grid (each grid value fills its 8 m block), scored only on the
maps that have a Sionna file, together with all other methods on the same maps. Timings: seconds to predict all 80
sites of one map.

  python -m rcm_ml.cell_placement --out runs/mps/cell_placement.json --fig docs/research/figures/05_cell_placement.png
"""
from __future__ import annotations

import argparse
import json
import time
from multiprocessing import Pool
from pathlib import Path

import numpy as np
import torch

from .data import DB_RANGE, N, P_TRNC, THRESH
from .raytrace import load_polygons, radio_map
from .tracer_dd import dd_map
from .train import Cache, Featurizer, git_commit, load_unet, pick_device, unet_predict

CACHE = Path("data/rms_cache")
THRESHOLDS_DB = (-100, -110, -120)
UNETS = {"U-Net geometry only": "geom_s128_m501_seed0", "U-Net + straight-line physics (B1)": "feats_s128_m501_seed0",
         "U-Net + street-canyon physics (B2)": "feats-B2_s128_m501_seed0", "U-Net + physics residual": "hybrid_s128_m501_seed0",
         "U-Net geometry only, trained on IRT2": "geom_irt2_s128_m501_seed0"}


def gray_thr(db):
    return (db - P_TRNC) / DB_RANGE


def tracer_one(args):
    m, r, c, ple, ds, rl = args
    g, _ = radio_map(load_polygons(m), (r, c), ple=ple, dslope=ds, rloss=rl)
    return (np.clip(g, 0, 1) * 255).astype(np.uint8)


def greedy(cov_maps, outdoor, k):
    """cov_maps: [sites, H, W] bool. Greedy k sites maximising the union over outdoor pixels."""
    chosen, cur = [], np.zeros(outdoor.shape, bool)
    for _ in range(k):
        gain = ((cov_maps | cur) & outdoor).reshape(len(cov_maps), -1).sum(1)
        gain[chosen] = -1
        j = int(np.argmax(gain))
        chosen.append(j)
        cur |= cov_maps[j]
    return chosen


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--maps", type=int, default=99)
    ap.add_argument("--k", type=int, default=3)
    ap.add_argument("--workers", type=int, default=10)
    ap.add_argument("--out", default="runs/mps/cell_placement.json")
    ap.add_argument("--fig", default=None)
    ap.add_argument("--unets", default=None, help="extra U-Nets as 'name=run_dir;name=run_dir'")
    ap.add_argument("--sionna-dir", default=None, help="runs/mps/cell_sionna: per-map Sionna files (scripts/cell_sionna.sh)")
    a = ap.parse_args()
    dev = pick_device("mps")
    te = {t: Cache(CACHE, "test", t, None, "B2") for t in ("dpm", "irt2")}
    ds = te["dpm"]
    maps = list(dict.fromkeys(ds.map_of.tolist()))[: a.maps]
    SION = ["Sionna tuned (8 m grid)", "Sionna tuned + two-corner paths (8 m grid)", "WinProp DPM read on the 8 m grid"]
    if a.sionna_dir:
        maps = [m for m in maps if (Path(a.sionna_dir) / f"{m}.npz").exists()]
        ddloss = json.loads(Path("runs/mps/sionna_cfg/scat09add_r100k/result.json").read_text())["best"]["ddloss"]
    cal = json.loads(Path("runs/mps/raytrace_3_calibrated.json").read_text())
    tr_args = (cal["path_loss_exponent"], float(cal["calibration"]["best"]["dslope"]), float(cal["calibration"]["best"]["rloss"]))
    models = {}
    if a.unets:
        UNETS.update(dict(kv.split("=", 1) for kv in a.unets.split(";")))
    for name, run in UNETS.items():
        models[name] = load_unet(run if "/" in run else f"runs/poc/{run}", dev)
    phys = {"Physics map B1 (straight-line walls)": json.loads(Path("runs/poc/feats_s128_m501_seed0/config.json").read_text())["phys_coef"],
            "Physics map B2 (street canyon)": json.loads(Path("runs/poc/feats-B2_s128_m501_seed0/config.json").read_text())["phys_coef"]}
    phys_feat = {n: Featurizer(dev, c, "B2" if "B2" in n else "B1", "feats", N, 0.0) for n, c in phys.items()}
    truths = ("WinProp DPM", "WinProp IRT2")
    methods = list(truths) + list(phys) + ["Our ray tracer, calibrated to IRT2"] + list(UNETS) + (SION if a.sionna_dir else [])
    rec = {t: {m: {th: {"regret1": [], "top1": [], "top5": [], "spearman": [], "regretk": []} for th in THRESHOLDS_DB} for m in methods} for t in truths}
    secs = {m: [] for m in methods}
    ctx = {t: {th: {"best": [], "median_site": [], "random_pick_regret": [], "best3": []} for th in THRESHOLDS_DB} for t in truths}
    example = None
    pool = Pool(a.workers)
    radio_map([], (128, 128), 360, 90)
    for i, m in enumerate(maps):
        ids = np.flatnonzero(ds.map_of == m)
        b = ds.batch(ids)
        outdoor = b["bld"][0].numpy() == 0
        grays = {"WinProp DPM": b["y"].numpy().astype(np.float32) / 255, "WinProp IRT2": te["irt2"].batch(ids)["y"].numpy().astype(np.float32) / 255}
        with torch.no_grad():
            for n, f in phys_feat.items():
                t0 = time.time()
                g = f.physics_gray({k: v.to(dev) for k, v in b.items()})
                if dev.type == "mps":
                    torch.mps.synchronize()
                grays[n] = g.cpu().numpy()
                secs[n].append(time.time() - t0)
            for n, (f, net, cfg) in models.items():
                t0 = time.time()
                x, *_ = f(b)
                p = unet_predict(net, x, cfg.get("gate", False)).clamp(0, 1)
                g = THRESH + (1 - THRESH) * p
                g = torch.nn.functional.interpolate(g, scale_factor=N // cfg["size"], mode="nearest")[:, 0]
                if dev.type == "mps":
                    torch.mps.synchronize()
                grays[n] = g.cpu().numpy()
                secs[n].append(time.time() - t0)
        t0 = time.time()
        rc = b["rc"].numpy()
        grays["Our ray tracer, calibrated to IRT2"] = np.stack(pool.map(tracer_one, [(m, int(r), int(c), *tr_args) for r, c in rc])).astype(np.float32) / 255
        secs["Our ray tracer, calibrated to IRT2"].append(time.time() - t0)
        if a.sionna_dir:
            z = np.load(Path(a.sionna_dir) / f"{m}.npz")
            zm = json.loads(str(z["meta"]))
            secs[SION[0]].append(zm["s_per_tx_median"] * 80)
            meta_tx = [x["tx"] for x in json.loads((CACHE / "test_meta.json").read_text())]
            site = np.array([meta_tx[k] for k in ids])  # transmitter id of each cache sample
            ddm = dict(pool.map(dd_map, [(m, int(t)) for t in site]))
            for name, with_dd in zip(SION, (False, True)):
                G = np.zeros((len(ids), N, N), np.float32)
                for j, t in enumerate(site):
                    k = z["rx_tx"] == t
                    pw = z["pw_cal"][k] + (ddm[(m, int(t))][z["rx_r"][k], z["rx_c"][k]] * 10.0 ** (-ddloss / 10.0) if with_dd else 0.0)
                    with np.errstate(divide="ignore"):
                        v = np.clip((10 * np.log10(pw) - P_TRNC) / DB_RANGE, 0, 1)
                    r0, c0 = z["rx_r"][k] - 4, z["rx_c"][k] - 4  # 8 m block around each grid point
                    for dr in range(8):
                        for dc in range(8):
                            G[j, r0 + dr, c0 + dc] = v
                grays[name] = G
            # the grid alone: the true DPM maps sampled at the same 8 m points and filled the same way
            gd = grays["WinProp DPM"][:, 4::8, 4::8]
            grays[SION[2]] = np.repeat(np.repeat(gd, 8, axis=1), 8, axis=2)
        for n in grays:
            grays[n][:, ~outdoor] = 0
        n_out = outdoor.sum()
        for th in THRESHOLDS_DB:
            cov = {n: g > gray_thr(th) for n, g in grays.items()}
            frac = {n: (c & outdoor).reshape(len(ids), -1).sum(1) / n_out for n, c in cov.items()}
            for t in truths:
                tf = frac[t]
                order = np.argsort(-tf)
                best_k = greedy(cov[t], outdoor, a.k)
                true_k = ((cov[t][best_k].any(0)) & outdoor).sum() / n_out
                c_ = ctx[t][th]
                c_["best"].append(float(tf.max()) * 100); c_["median_site"].append(float(np.median(tf)) * 100)
                c_["random_pick_regret"].append(float(tf.max() - tf.mean()) * 100); c_["best3"].append(float(true_k) * 100)
                for n in methods:
                    pick = int(np.argmax(frac[n]))
                    r_ = rec[t][n][th]
                    r_["regret1"].append(float(tf.max() - tf[pick]) * 100)
                    r_["top1"].append(float(pick == order[0]))
                    r_["top5"].append(float(pick in order[:5]))
                    pr, trk = np.argsort(np.argsort(-frac[n])), np.argsort(np.argsort(-tf))
                    r_["spearman"].append(float(np.corrcoef(pr, trk)[0, 1]))
                    ch = greedy(cov[n], outdoor, a.k)
                    r_["regretk"].append(float(true_k - ((cov[t][ch].any(0)) & outdoor).sum() / n_out) * 100)
            if th == -110 and example is None and i == 0:
                example = {"map": m, "outdoor": outdoor, "rc": rc, "frac": {n: frac[n] for n in ("WinProp DPM", "U-Net + street-canyon physics (B2)", "Physics map B1 (straight-line walls)")},
                           "cov_best": {n: cov[n][int(np.argmax(frac[n]))] for n in ("WinProp DPM", "U-Net + street-canyon physics (B2)")},
                           "picks": {n: int(np.argmax(frac[n])) for n in ("WinProp DPM", "U-Net + street-canyon physics (B2)", "Physics map B1 (straight-line walls)")},
                           "bld": ~outdoor}
        if (i + 1) % 10 == 0:
            print(f"{i + 1}/{len(maps)} maps", flush=True)
    pool.close()
    summ = {t: {n: {str(th): {"one_site_regret_pp_mean": round(float(np.mean(v["regret1"])), 2),
                                "one_site_regret_pp_median": round(float(np.median(v["regret1"])), 2),
                                "one_site_regret_pp_p90": round(float(np.percentile(v["regret1"], 90)), 2),
                                "picks_true_best_pct": round(float(np.mean(v["top1"])) * 100, 1),
                                "picks_true_top5_pct": round(float(np.mean(v["top5"])) * 100, 1),
                                "rank_correlation_mean": round(float(np.nanmean(v["spearman"])), 3),
                                f"{a.k}_site_regret_pp_mean": round(float(np.mean(v["regretk"])), 2)}
                          for th, v in d.items()} for n, d in rec[t].items()} for t in truths}
    if a.sionna_dir:
        secs[SION[1]] = secs[SION[0]]
        secs.pop(SION[2], None)
    res = {"commit": git_commit(), "maps": len(maps), "map_ids": [int(m) for m in maps], "sites_per_map": 80, "k": a.k, "thresholds_db": list(THRESHOLDS_DB),
           "seconds_per_map_all_80_sites_median": {n: round(float(np.median(v)), 4) for n, v in secs.items() if v},
           "unet_runs": UNETS, "results": summ,
           "context": {t: {str(th): {k: round(float(np.mean(v)), 2) for k, v in d.items()} for th, d in ctx[t].items()} for t in truths}}
    Path(a.out).write_text(json.dumps(res, indent=1))
    print(json.dumps(res["seconds_per_map_all_80_sites_median"], indent=1))
    if a.fig and example:
        import matplotlib
        matplotlib.use("Agg")
        import matplotlib.pyplot as plt
        e = example
        fig, axs = plt.subplots(1, 3, figsize=(17, 5.6), gridspec_kw={"width_ratios": [1, 1, 1.3]})
        for ax, n in zip(axs[:2], ("WinProp DPM", "U-Net + street-canyon physics (B2)")):
            ax.imshow(np.where(e["bld"], 0.35, np.where(e["cov_best"][n], 1.0, 0.0)), cmap="Greys_r", vmin=0, vmax=1, interpolation="nearest")
            ax.scatter(e["rc"][:, 1], e["rc"][:, 0], s=10, c="#1f77b4")
            p = e["picks"][n]
            ax.plot(e["rc"][p, 1], e["rc"][p, 0], marker="*", ms=16, color="#d62728", mec="white")
            tf = e["frac"]["WinProp DPM"][p]
            ax.set_title(f"Pick: {n}\ntrue coverage {tf:.0%} (white: that method's map above −110 dB)", fontsize=9.5)
            ax.set_xticks([]); ax.set_yticks([])
        ax = axs[2]
        t = e["frac"]["WinProp DPM"] * 100
        for n, col in (("U-Net + street-canyon physics (B2)", "#d62728"), ("Physics map B1 (straight-line walls)", "#888888")):
            ax.scatter(t, e["frac"][n] * 100, s=14, color=col, label=n)
        lim = [0, 100]
        ax.plot(lim, lim, color="#222222", lw=0.8)
        ax.set_xlabel("true coverage of each of the 80 sites (WinProp DPM), %")
        ax.set_ylabel("predicted coverage, %")
        ax.legend(frameon=False, fontsize=9)
        ax.set_title(f"Test map {e['map']}: predicted vs true coverage at −110 dB", fontsize=10)
        plt.savefig(a.fig, dpi=75, bbox_inches="tight")


if __name__ == "__main__":
    main()
