"""SpectrumNet (arXiv 2408.15252, CC0, github.com/ShuhangZhang/FDRadiomap): km-scale radio maps for learning range.

Data: scripts/get_spectrumnet.sh --full -> data/spectrumnet/full/PPData5D-success/ (git-ignored):
  png/<TT.Env>/T<TT>C<c>D<dddd>_n<nn>_f<ff>_ss_z<zz>.png   128 x 128 uint8, 10 m pixels (1.28 km), one Tx per map
  npz/T<TT>C<c>D<dddd>_n<nn>_bdtr.npz                        inBldg_zyx (3, 128, 128) 0/1 at the 3 heights, terrain_yx (m)
  f00..f04 = 0.15 / 1.5 / 1.7 / 3.5 / 22 GHz, z00..z02 = Rx 1.5 / 30 / 200 m above the terrain (manuscript Appendix A).
  n = Tx sampling within the area. A map holds one to several transmitters (superposed, see the 200 m layer's ring
  patterns); their positions are not stored and are detected from the 1.5 m and 30 m maps (find_txs).

Encoding (fitted by `inspect`, runs/spectrumnet/inspect.json 'encoding'): 0 = no ray / building, otherwise
  path gain dB = (gray - GRAY_0DB) / GRAY_PER_DB, i.e. 1.44 gray per dB, no per-image normalisation.

Subcommands
  python -m rcm_ml.spectrumnet inspect                     -> runs/spectrumnet/inspect.json, figures/08_spectrumnet_inspect.png
  python -m rcm_ml.spectrumnet train --name geom [--dist] [--gate]  -> runs/spectrumnet/<name>/{config,log,test}.json (+ best.pt)
  python -m rcm_ml.spectrumnet mosaic --runs a,b           -> runs/spectrumnet/mosaic_2km.json, figures/08_spectrumnet_2km.png

The U-Net is models.UNet (base 16, depth 5) as in runs/poc and runs/capped, at the native 128 px / 10 m, on the urban
3.5 GHz maps with Rx at 1.5 m. Targets use the RadioMapSeer training scale: floor -127.2 dB (THRESH), top -47.84 dB.
The targets are nearly binary (no ray, or about free-space level), so MSE alone predicts a small positive mean, i.e.
'covered', where coverage is uncertain. --gate adds the 'covered' logit of train.py --gate (BCE, prediction set to the
floor where the logit is negative).
"""
from __future__ import annotations

import argparse
import json
import math
import os
import time
from pathlib import Path

import numpy as np
import torch
import torch.nn.functional as F
from PIL import Image

from .data import DB_RANGE, M1, P_TRNC, THRESH
from .models import build
from .train import DIST_SCALE, git_commit, pick_device

REPO = Path(__file__).resolve().parent.parent
ROOT = Path(os.environ.get("SPECTRUMNET_ROOT", REPO / "data/spectrumnet/full/PPData5D-success"))
OUT = REPO / "runs/spectrumnet"
FIGS = REPO / "docs/research/figures"
FREQ_MHZ = [150.0, 1500.0, 1700.0, 3500.0, 22000.0]
HEIGHTS = [1.5, 30.0, 200.0]
N, PX = 128, 10.0  # pixels, metres per pixel
URBAN = ["06.DenseUrban", "08.OrdinaryUrban"]
F35, Z15 = 3, 0  # 3.5 GHz, Rx 1.5 m
GRAY_PER_DB, GRAY_0DB = 1.44, 242.6  # from `inspect` (joint free-space fit of the three layers, single-Tx maps)
TX_REACH_M = 15.0  # a Tx pixel is at least as strong as free space at 15 m
FLOOR_DB = P_TRNC + THRESH * DB_RANGE  # -127.2 dB, the RadioMapSeer / RadioUNet floor
TOP_DB = M1  # -47.84 dB, top of the RadioMapSeer training scale
F59_SHIFT_DB = 20 * math.log10(5900 / 3500)  # 4.53 dB: free-space loss difference 3.5 -> 5.9 GHz
RINGS = [(0, 100), (100, 200), (200, 400), (400, 800), (800, 1200)]
TERRAIN_SCALE = 50.0  # terrain channel: (terrain - terrain at the Tx) / 50 m, clipped to +/- 2


def fspl_db(d_m, f_mhz):
    return 20 * np.log10(np.maximum(d_m, 1e-3) / 1000) + 20 * np.log10(f_mhz) + 32.44


def gray_to_db(g):
    g = np.asarray(g, np.float32)
    return np.where(g > 0, (g - GRAY_0DB) / GRAY_PER_DB, np.nan).astype(np.float32)


def inventory(env: str) -> dict[str, set]:
    """key 'T06C0D0001_n02' -> set of (f, z) present."""
    inv: dict[str, set] = {}
    for f in os.listdir(ROOT / "png" / env):
        if f.endswith(".png"):
            inv.setdefault(f[:14], set()).add((int(f[16:18]), int(f[23:25])))
    return inv


def load_png(env, key, f, z):
    return np.asarray(Image.open(ROOT / "png" / env / f"{key}_f{f:02d}_ss_z{z:02d}.png"))


def load_npz(key):
    z = np.load(ROOT / "npz" / f"{key}_bdtr.npz")
    return z["inBldg_zyx"], z["terrain_yx"].astype(np.float32)


def find_txs(g, g30=None, f=F35):
    """Transmitters of a map. Candidates: local maxima (7 x 7 px) of the 1.5 m map g at least as strong as free space at
    TX_REACH_M, at least 4 px apart. With the 30 m map g30 (Rx above the street canyons) a candidate is kept only if the
    smoothed 30 m map has a local maximum (9 x 9 px) within 2 px: this drops bright street spots near a Tx. Falls back
    to the brightest candidate. Returns a (k, 2) int array, brightest first."""
    from scipy.ndimage import gaussian_filter, maximum_filter

    db = np.nan_to_num(gray_to_db(g), nan=-999.0)
    cand = np.argwhere((g == maximum_filter(g, 7)) & (db >= -fspl_db(TX_REACH_M, FREQ_MHZ[f])))
    cand = cand[np.argsort(-g[cand[:, 0], cand[:, 1]], kind="stable")]
    keep: list = []
    for p in cand:
        if all(max(abs(p[0] - q[0]), abs(p[1] - q[1])) > 3 for q in keep):
            keep.append(p)
    if not keep:
        keep = [np.array(np.unravel_index(int(np.argmax(g)), g.shape))]
    if g30 is not None:
        s = gaussian_filter(np.nan_to_num(gray_to_db(g30), nan=-200.0), 1.0)
        lm = (s == maximum_filter(s, 9)) & (s > -150)
        ok = [p for p in keep if lm[max(p[0] - 2, 0):p[0] + 3, max(p[1] - 2, 0):p[1] + 3].any()]
        keep = ok or keep[:1]
    return np.array(keep, int)


def load_txs(env, key, inv):
    g30 = load_png(env, key, F35, 1) if (F35, 1) in inv[key] else None
    return find_txs(load_png(env, key, F35, Z15), g30)


def area_of(key):
    return key[:10]


def los_mask(bld, r, c, steps=256):
    """True where the straight line from the Tx pixel to the pixel centre crosses no building pixel."""
    rr, cc = np.mgrid[:N, :N]
    t = np.linspace(0, 1, steps)[:, None, None]
    i = np.rint(r + t * (rr - r)).astype(int)
    j = np.rint(c + t * (cc - c)).astype(int)
    return ~bld[i, j].any(0)


def dist_m(txs):
    """Distance (m) of every pixel to the nearest Tx."""
    rr, cc = np.mgrid[:N, :N]
    return PX * np.min([np.hypot(rr - r, cc - c) for r, c in txs], 0)


def ring_counts(db, outdoor, d, floor):
    out = {}
    for lo, hi in RINGS:
        m = outdoor & (d >= lo) & (d < hi)
        out[f"{lo}-{hi} m"] = [int((np.nan_to_num(db, nan=-999)[m] > floor).sum()), int(m.sum())]
    return out


def add_counts(acc, c):
    for k, (a, b) in c.items():
        acc.setdefault(k, [0, 0])
        acc[k][0] += a
        acc[k][1] += b


def pct(acc):
    return {k: round(100 * a / b, 1) if b else None for k, (a, b) in acc.items()}


# ----------------------------------------------------------------------------------------------------------- inspect

def fit_encoding(samples, h_grid):
    """samples: list per map of (layer, horizontal distance m, gray). Joint fit gray = a - s * FSPL(d3), d3 with the Tx
    at height h. Returns per-h fits and a bootstrap over maps at h = 1.5 m."""
    allv = np.concatenate(samples)
    hr = np.array(HEIGHTS)[allv[:, 0].astype(int)]
    fits = []
    for h in h_grid:
        x = -fspl_db(np.hypot(np.maximum(allv[:, 1], 2.5), hr - h), FREQ_MHZ[F35])
        A = np.c_[np.ones(len(x)), x]
        (a, s), *_ = np.linalg.lstsq(A, allv[:, 2], rcond=None)
        res = A @ [a, s] - allv[:, 2]
        fits.append({"tx_height_m": h, "gray_0db": round(float(a), 2), "gray_per_db": round(float(s), 4),
                     "rmse_gray": round(float(np.sqrt((res ** 2).mean())), 2),
                     "mean_residual_gray_by_layer": {f"{HEIGHTS[z]} m": round(float(res[allv[:, 0] == z].mean()), 2) for z in range(3)}})
    rng = np.random.default_rng(0)
    boot = []
    for _ in range(100):
        v = np.concatenate([samples[i] for i in rng.integers(0, len(samples), len(samples))])
        x = -fspl_db(np.hypot(np.maximum(v[:, 1], 2.5), np.array(HEIGHTS)[v[:, 0].astype(int)] - 1.5), FREQ_MHZ[F35])
        boot.append(np.linalg.lstsq(np.c_[np.ones(len(x)), x], v[:, 2], rcond=None)[0])
    boot = np.array(boot)
    return fits, {"gray_0db_ci95": np.percentile(boot[:, 0], [2.5, 97.5]).round(2).tolist(),
                  "gray_per_db_ci95": np.percentile(boot[:, 1], [2.5, 97.5]).round(4).tolist()}


def inspect(a):
    t0 = time.time()
    res = {"commit": git_commit(), "root": str(ROOT.relative_to(REPO)), "inventory": {}}
    for env in sorted(os.listdir(ROOT / "png")):
        inv = inventory(env)
        res["inventory"][env] = {"keys": len(inv), "areas": len({area_of(k) for k in inv}),
                                 "complete_15": sum(len(v) == 15 for v in inv.values()),
                                 "f03_z00": sum((F35, Z15) in v for v in inv.values())}
    npz_vals = set()
    rng = np.random.default_rng(0)
    enc_samples, freq_diff = [], {z: {i: [] for i in range(4)} for z in range(3)}
    tx_stats = {"maps": 0, "confirmed_at_30m": 0, "n_tx": {}, "tx_on_building": 0, "tx_edge_dist_px": [],
                "cross_freq": {f"f{f:02d}": {"f03_tx_found": 0, "f03_tx": 0, "extra_tx": 0} for f in (0, 4)}}
    nz_min, hist = 255, np.zeros(256, np.int64)
    bld_frac = {z: [] for z in range(3)}
    bld_zero = [0, 0]
    terr = []
    cov = {f"{env} f{f:02d}": {} for env in URBAN for f in range(5)}
    cov59 = {env: {} for env in URBAN}
    cov1 = {env: {} for env in URBAN}
    cov_z = {f"{env} z{z:02d}": {} for env in URBAN for z in (1, 2)}
    bld_ring = {env: {} for env in URBAN}
    ring_hist = {env: {f"{lo}-{hi} m": np.zeros(256, np.int64) for lo, hi in RINGS} for env in URBAN}  # outdoor gray histograms
    examples = []
    for env in URBAN:
        inv = inventory(env)
        keys = [str(k) for k in rng.permutation(sorted(k for k, v in inv.items() if (F35, Z15) in v))]
        n_enc = 0
        for key in keys:
            bz, terrain = load_npz(key)
            npz_vals |= set(np.unique(bz).tolist())
            bld = bz[0] > 0
            g = load_png(env, key, F35, Z15)
            txs = load_txs(env, key, inv)
            r, c = txs[0]
            tx_stats["maps"] += 1
            tx_stats["confirmed_at_30m"] += int((F35, 1) in inv[key])
            tx_stats["n_tx"][len(txs)] = tx_stats["n_tx"].get(len(txs), 0) + 1
            tx_stats["tx_on_building"] += int(bld[txs[:, 0], txs[:, 1]].sum())
            tx_stats["tx_edge_dist_px"] += [min(p[0], p[1], N - 1 - p[0], N - 1 - p[1]) for p in txs]
            for f in (0, 4):  # the same transmitters must show up at other frequencies
                if (f, Z15) in inv[key]:
                    tf = find_txs(load_png(env, key, f, Z15), f=f)  # unconfirmed candidates at that frequency
                    near = lambda p, q: (np.abs(q - p).max(1) <= 1).any()
                    cf = tx_stats["cross_freq"][f"f{f:02d}"]
                    cf["f03_tx"] += len(txs)
                    cf["f03_tx_found"] += int(sum(near(p, tf) for p in txs))
                    cf["extra_tx"] += int(sum(not near(q, txs) for q in tf))
            nz = g[g > 0]
            if nz.size:
                nz_min = min(nz_min, int(nz.min()))
            hist += np.bincount(g[~bld].ravel(), minlength=256)
            for z in range(3):
                bld_frac[z].append(float(bz[z].mean()))
            bld_zero[0] += int((g[bld] == 0).sum()); bld_zero[1] += int(bld.sum())
            terr.append(float(terrain.max() - terrain.min()))
            d = dist_m(txs)
            outdoor = ~bld
            for f in range(5):
                if (f, Z15) in inv[key]:
                    add_counts(cov[f"{env} f{f:02d}"], ring_counts(gray_to_db(load_png(env, key, f, Z15)), outdoor, d, FLOOR_DB))
            for lo, hi in RINGS:
                ring_hist[env][f"{lo}-{hi} m"] += np.bincount(g[outdoor & (d >= lo) & (d < hi)], minlength=256)
            db = gray_to_db(g)
            add_counts(cov59[env], ring_counts(db - F59_SHIFT_DB, outdoor, d, FLOOR_DB))
            if len(txs) == 1:
                add_counts(cov1[env], ring_counts(db, outdoor, d, FLOOR_DB))
            for z in (1, 2):
                if (F35, z) in inv[key]:
                    add_counts(cov_z[f"{env} z{z:02d}"], ring_counts(gray_to_db(load_png(env, key, F35, z)), bz[z] == 0, d, FLOOR_DB))
            add_counts(bld_ring[env], {k: [int(bld[(d >= lo) & (d < hi)].sum()), int(((d >= lo) & (d < hi)).sum())]
                                       for k, (lo, hi) in zip([f"{lo}-{hi} m" for lo, hi in RINGS], RINGS)})
            if len(txs) == 1 and len(inv[key]) == 15 and n_enc < a.fit_maps:  # encoding fit on single-Tx maps only
                n_enc += 1
                los = los_mask(bld, r, c)
                rows = []
                m0 = los & (d >= 10) & (d <= 100) & (g > 0)  # 1.5 m layer: LOS street pixels before the two-ray breakpoint
                rows.append(np.c_[np.zeros(m0.sum()), d[m0], g[m0]])
                for z, dmax in ((1, 30), (2, 300)):  # 30 m layer right around the Tx, 200 m layer (mostly LOS) within 300 m
                    gz = load_png(env, key, F35, z)
                    mz = (gz > 0) & (d <= dmax)
                    rows.append(np.c_[np.full(mz.sum(), z), d[mz], gz[mz]])
                enc_samples.append(np.concatenate(rows).astype(np.float64))
                for z in range(3):
                    A = [load_png(env, key, f, z).astype(np.float64) for f in range(5)]
                    for i in range(4):
                        m = (A[i] > 0) & (A[i + 1] > 0)
                        freq_diff[z][i].append(A[i][m] - A[i + 1][m])
            if len(examples) < 2 * (URBAN.index(env) + 1) and len(inv[key]) == 15 and 0.1 < bld.mean() < 0.3 and len(txs) <= 2:
                examples.append((env, key))
    fits, boot = fit_encoding(enc_samples, a.h_grid)
    res["encoding"] = {
        "rule": "path_gain_db = (gray - gray_0db) / gray_per_db for gray > 0; gray 0 = no ray found or building",
        "adopted": {"gray_0db": GRAY_0DB, "gray_per_db": GRAY_PER_DB},
        "fit_maps": len(enc_samples), "fit_points": int(sum(len(s) for s in enc_samples)),
        "joint_fspl_fit_by_tx_height": fits, "bootstrap_tx_1.5m": boot,
        "frequency_differences": {
            f"{HEIGHTS[z]} m": {f"f{i:02d}-f{i + 1:02d}": {"mean_gray": round(float(np.concatenate(freq_diff[z][i]).mean()), 2),
                                                          "fspl_db": round(float(20 * np.log10(FREQ_MHZ[i + 1] / FREQ_MHZ[i])), 2),
                                                          "gray_per_db": round(float(np.concatenate(freq_diff[z][i]).mean() / (20 * np.log10(FREQ_MHZ[i + 1] / FREQ_MHZ[i]))), 3)}
                                for i in range(4)} for z in range(3)},
        "nonzero_gray_min": nz_min, "nonzero_min_db": round(float((nz_min - GRAY_0DB) / GRAY_PER_DB), 1),
        "gray_of_floor": round(GRAY_0DB + GRAY_PER_DB * FLOOR_DB, 1),
        "outdoor_db_percentiles_nonzero": dict(zip(["p1", "p10", "p50", "p90", "p99"], np.round((np.percentile(
            np.repeat(np.arange(1, 256), hist[1:]), [1, 10, 50, 90, 99]) - GRAY_0DB) / GRAY_PER_DB, 1).tolist())),
        "outdoor_zero_pct": round(100 * float(hist[0] / hist.sum()), 1)}
    e = tx_stats.pop("tx_edge_dist_px")
    for cf in tx_stats["cross_freq"].values():
        cf["recall_pct"] = round(100 * cf["f03_tx_found"] / max(cf["f03_tx"], 1), 1)
    tx_stats["n_tx"] = dict(sorted(tx_stats["n_tx"].items()))
    res["tx"] = {**tx_stats, "tx_edge_dist_px_percentiles": dict(zip(["p10", "p50", "p90"], np.percentile(e, [10, 50, 90]).tolist())),
                 "rule": f"Tx = local maxima (7x7) of the 1.5 m map at least as strong as free space at {TX_REACH_M:g} m, >= 4 px "
                         "apart, kept where the smoothed 30 m map has a local maximum within 2 px (find_txs)",
                 "cross_freq_note": "recall = share of the 3.5 GHz Tx that are also candidates (within 1 px) at 150 MHz / 22 GHz"}
    res["buildings"] = {"npz_values": sorted(int(v) for v in npz_vals),
                        "building_pct_by_layer": {f"{HEIGHTS[z]} m": round(100 * float(np.mean(bld_frac[z])), 2) for z in range(3)},
                        "building_pixels_zero_in_1.5m_map_pct": round(100 * bld_zero[0] / max(bld_zero[1], 1), 2),
                        "terrain_range_m_percentiles": dict(zip(["p10", "p50", "p90"], np.round(np.percentile(terr, [10, 50, 90]), 1).tolist())),
                        "building_pct_by_ring": {env: pct(v) for env, v in bld_ring.items()}}
    res["coverage"] = {"floor_db": round(FLOOR_DB, 1), "rx": "1.5 m unless stated, outdoor pixels (no building at that layer)",
                       "by_freq_1.5m": {k: pct(v) for k, v in cov.items()},
                       "f03_1.5m_5.9GHz_equivalent": {env: pct(v) for env, v in cov59.items()},
                       "f03_1.5m_single_tx_maps": {env: pct(v) for env, v in cov1.items()},
                       "distance": "to the nearest Tx",
                       "f03_higher_rx": {k: pct(v) for k, v in cov_z.items()},
                       "pixels": {k: {r: v[1] for r, v in cov[f"{k} f03"].items()} for k in URBAN}}
    fl = int(math.floor(GRAY_0DB + GRAY_PER_DB * FLOOR_DB))  # highest gray at or below the floor
    res["coverage"]["f03_1.5m_detail"] = {env: {r: {
        "no_ray_pct": round(100 * float(h[0] / h.sum()), 1), "weak_pct": round(100 * float(h[1:fl + 1].sum() / h.sum()), 1),
        "median_db_covered": round(float((np.searchsorted(np.cumsum(h[fl + 1:]), h[fl + 1:].sum() / 2) + fl + 1 - GRAY_0DB) / GRAY_PER_DB), 1),
        "fspl_db_mid": round(float(-fspl_db((lo + hi) / 2, FREQ_MHZ[F35])), 1),
        "two_ray_db_mid": round(float(-(40 * np.log10((lo + hi) / 2) - 20 * np.log10(HEIGHTS[0] * HEIGHTS[0]))), 1)}
        for r, h, (lo, hi) in zip(hs, hs.values(), RINGS)} for env, hs in ring_hist.items()}
    res["references"] = references()
    res["seconds"] = round(time.time() - t0, 1)
    OUT.mkdir(parents=True, exist_ok=True)
    (OUT / "inspect.json").write_text(json.dumps(res, indent=1))
    print(json.dumps({k: res[k] for k in ("encoding", "tx", "buildings", "coverage")}, indent=1))
    inspect_figure(res, enc_samples, examples)


def references():
    """Coverage by ring of the other references, on rings 0-100/100-200/200-400 m (WinProp rings merged by pixel count)."""
    w = json.loads((REPO / "runs/convergence/winprop_rings.json").read_text())
    merge = {"0-100 m": ["0-50 m", "50-100 m"], "100-200 m": ["100-150 m", "150-200 m"], "200-400 m": ["200-260 m", "260-400 m"]}
    out = {"winprop_tiles (runs/convergence/winprop_rings.json, 5.9 GHz)": {
        sim: {r: round(sum(w["covered_pct"][sim][p] * w["outdoor_pixels"][p] for p in ps) / sum(w["outdoor_pixels"][p] for p in ps), 1)
              for r, ps in merge.items()} for sim in ("DPM", "IRT2", "IRT4")}}
    c = json.loads((REPO / "runs/convergence/report.json").read_text())
    out["2 km mosaic receivers (runs/convergence/report.json, 5.9 GHz)"] = {
        "tracer int4": {r: v["covered_pct"] for r, v in c["tracer"]["int4"]["rings"].items()},
        "Sionna + two-corner d5_1e6": {r: v["covered_pct"] for r, v in c["sionna"]["d5_1e6"]["rings_with_dd"].items()}}
    return out


def inspect_figure(res, enc_samples, examples):
    import matplotlib
    matplotlib.use("Agg")
    import matplotlib.pyplot as plt
    fig = plt.figure(figsize=(18, 10))
    ax = fig.add_subplot(2, 3, 1)
    allv = np.concatenate(enc_samples)
    cols = ["#1f77b4", "#ff7f0e", "#2ca02c"]
    for z in range(3):
        v = allv[allv[:, 0] == z]
        x = fspl_db(np.hypot(np.maximum(v[:, 1], 2.5), HEIGHTS[z] - 1.5), FREQ_MHZ[F35])
        bins = np.arange(np.floor(x.min()), np.ceil(x.max()) + 1, 1.0)
        idx = np.digitize(x, bins)
        med = [(bins[i - 1] + 0.5, np.median(v[idx == i, 2])) for i in np.unique(idx) if (idx == i).sum() > 30]
        med = np.array(med)
        ax.scatter(x[::20], v[::20, 2], s=1, alpha=0.15, color=cols[z])
        ax.plot(med[:, 0], med[:, 1], "o-", ms=3, color=cols[z], label=f"Rx {HEIGHTS[z]:g} m (median)")
    xs = np.linspace(45, 110, 10)
    ax.plot(xs, GRAY_0DB - GRAY_PER_DB * xs, "k--", label=f"gray = {GRAY_0DB} - {GRAY_PER_DB} x FSPL")
    ax.set_xlabel("free-space loss, 3.5 GHz, Tx 1.5 m (dB)"); ax.set_ylabel("PNG gray"); ax.legend(fontsize=8)
    ax.set_title("Encoding: LOS pixels vs free-space loss")
    ax = fig.add_subplot(2, 3, 2)
    rings = [f"{lo}-{hi} m" for lo, hi in RINGS[:4]]
    x = np.arange(len(rings))
    series = {f"SpectrumNet {env[3:]} 3.5 GHz": res["coverage"]["by_freq_1.5m"][f"{env} f03"] for env in URBAN}
    series.update({f"SpectrumNet {env[3:]} 5.9 GHz-equiv.": res["coverage"]["f03_1.5m_5.9GHz_equivalent"][env] for env in URBAN})
    for sim, v in res["references"]["winprop_tiles (runs/convergence/winprop_rings.json, 5.9 GHz)"].items():
        series[f"WinProp {sim} (256 m tiles)"] = v
    series["our tracer int4 (2 km mosaic)"] = res["references"]["2 km mosaic receivers (runs/convergence/report.json, 5.9 GHz)"]["tracer int4"]
    w = 0.8 / len(series)
    for i, (name, v) in enumerate(series.items()):
        ax.bar(x + i * w - 0.4 + w / 2, [v.get(r) or 0 for r in rings], w, label=name)
    ax.set_xticks(x, rings); ax.set_ylabel("outdoor % above -127 dB"); ax.legend(fontsize=7)
    ax.set_title("Coverage by distance, Rx 1.5 m")
    for i, (env, key) in enumerate(examples[:4]):
        ax = fig.add_subplot(2, 6, 7 + i) if i < 4 else None
        g = load_png(env, key, F35, Z15)
        bz, _ = load_npz(key)
        txs = load_txs(env, key, inventory(env))
        db = gray_to_db(g)
        ax.set_facecolor("black")
        im = ax.imshow(db, vmin=-160, vmax=-50, cmap="viridis")
        ax.contour(np.nan_to_num(db, nan=-999), levels=[FLOOR_DB], colors="red", linewidths=0.5)
        ax.imshow(np.where(bz[0] > 0, 1.0, np.nan), cmap="Greys", vmin=0, vmax=1.6)
        ax.plot(txs[:, 1], txs[:, 0], "*", color="#d62728", ms=10, mec="white", ls="none")
        ax.set_title(f"{key} 3.5 GHz 1.5 m", fontsize=8); ax.set_xticks([]); ax.set_yticks([])
    ax = fig.add_subplot(2, 6, 11)
    if examples:
        env, key = examples[0]
        g2 = gray_to_db(load_png(env, key, F35, 2))
        ax.imshow(g2, cmap="viridis"); ax.set_title(f"{key} Rx 200 m", fontsize=8); ax.set_xticks([]); ax.set_yticks([])
    fig.colorbar(im, ax=fig.axes[2:7], fraction=0.015, pad=0.01).set_label("path gain dB (black: no ray, grey: building, red: -127 dB)")
    FIGS.mkdir(parents=True, exist_ok=True)
    plt.savefig(FIGS / "08_spectrumnet_inspect.png", dpi=70, bbox_inches="tight")


# ------------------------------------------------------------------------------------------------------------- train

def split_areas(areas, seed=0, frac=(0.7, 0.1, 0.2)):
    """Split by area (all Tx placements of an area land in the same split)."""
    u = sorted(set(areas))
    p = np.random.default_rng(seed).permutation(len(u))
    n1, n2 = int(frac[0] * len(u)), int((frac[0] + frac[1]) * len(u))
    part = {}
    for i, j in enumerate(p):
        part[u[j]] = "train" if i < n1 else "val" if i < n2 else "test"
    return part


def load_urban(envs=URBAN):
    """All maps with a 3.5 GHz / 1.5 m png: target dB (NaN = no signal), buildings, terrain, Tx."""
    keys, envl, dbs, blds, terrs, txs = [], [], [], [], [], []
    for env in envs:
        inv = inventory(env)
        for key in sorted(k for k, v in inv.items() if (F35, Z15) in v):
            g = load_png(env, key, F35, Z15)
            bz, terrain = load_npz(key)
            keys.append(key); envl.append(env)
            dbs.append(gray_to_db(g)); blds.append(bz[0] > 0); terrs.append(terrain); txs.append(load_txs(env, key, inv))
    return {"keys": keys, "env": np.array(envl), "db": np.stack(dbs), "bld": np.stack(blds), "terrain": np.stack(terrs),
            "tx": txs}


def to_train_scale(db):
    return np.clip((np.nan_to_num(db, nan=-999.0) - FLOOR_DB) / (TOP_DB - FLOOR_DB), 0.0, 1.0).astype(np.float32)


def inputs(bld, terrain, txs, dist: bool):
    """(B, C, H, W) float32: buildings, Tx one-hot (all Tx), terrain relative to the strongest Tx, [log distance to the
    nearest Tx, as train.dist_channel]. txs: per map a (k, 2) array. Works for any H, W."""
    B, H, W = bld.shape
    x = np.zeros((B, 3 + int(dist), H, W), np.float32)
    x[:, 0] = bld
    rr, cc = np.mgrid[:H, :W]
    for i, t in enumerate(txs):
        x[i, 1, t[:, 0], t[:, 1]] = 1
        x[i, 2] = np.clip((terrain[i] - terrain[i][t[0, 0], t[0, 1]]) / TERRAIN_SCALE, -2, 2)
        if dist:
            d = PX * np.min([np.hypot(rr - r, cc - c) for r, c in t], 0)
            x[i, 3] = 1 - np.log10(np.maximum(d, 1.0)) / DIST_SCALE
    return x


def augment(x, y, rng):
    """Random dihedral transform per sample (rot90 x flip)."""
    xs, ys = [], []
    for i in range(len(x)):
        k, fl = int(rng.integers(4)), bool(rng.integers(2))
        a, b = torch.rot90(x[i], k, (-2, -1)), torch.rot90(y[i], k, (-2, -1))
        if fl:
            a, b = torch.flip(a, (-1,)), torch.flip(b, (-1,))
        xs.append(a); ys.append(b)
    return torch.stack(xs), torch.stack(ys)


def predict_db(model, x, dev, gate=False, batch=32):
    out = []
    with torch.no_grad():
        for i in range(0, len(x), batch):
            p = model(torch.from_numpy(x[i:i + batch]).to(dev))
            v = p[:, 0].clamp(0, 1)
            if gate:
                v = torch.where(p[:, 1] > 0, v, torch.zeros_like(v))
            out.append(v.float().cpu().numpy())
    p = np.concatenate(out)
    return np.where(p > 0, FLOOR_DB + p * (TOP_DB - FLOOR_DB), np.nan).astype(np.float32)


def evaluate(pred_db, true_db, bld, tx):
    """Outdoor RMSE with both maps clipped at the floor, RMSE where the truth is above the floor, coverage by ring."""
    fl = lambda v: np.maximum(np.nan_to_num(v, nan=-999.0), FLOOR_DB)
    out = ~bld
    e = (fl(pred_db) - fl(true_db))[out]
    tc = np.nan_to_num(true_db, nan=-999) > FLOOR_DB
    pc = np.nan_to_num(pred_db, nan=-999) > FLOOR_DB
    ea = (fl(pred_db) - fl(true_db))[out & tc]
    rp, rt = {}, {}
    for i in range(len(tx)):
        d = dist_m(tx[i])
        add_counts(rp, ring_counts(pred_db[i], out[i], d, FLOOR_DB))
        add_counts(rt, ring_counts(true_db[i], out[i], d, FLOOR_DB))
    return {"maps": int(len(tx)), "rmse_db_outdoor_floor_clipped": round(float(np.sqrt((e ** 2).mean())), 3),
            "rmse_db_where_truth_covered": round(float(np.sqrt((ea ** 2).mean())), 3),
            "bias_db_where_truth_covered": round(float(ea.mean()), 3),
            "dead_painted_covered_pct": round(100 * float((pc & ~tc & out).sum() / max((~tc & out).sum(), 1)), 2),
            "covered_painted_dead_pct": round(100 * float((~pc & tc & out).sum() / max((tc & out).sum(), 1)), 2),
            "covered_pct_by_ring": {"truth": pct(rt), "pred": pct(rp)}}


def radial_baseline(train_db, train_bld, train_tx, bins_m=10.0):
    """Mean training target (floor-clipped dB) of outdoor pixels per 10 m distance bin: what distance alone predicts."""
    s, n = np.zeros(200), np.zeros(200)
    for db, b, t in zip(train_db, train_bld, train_tx):
        k = np.minimum((dist_m(t) / bins_m).astype(int), 199)[~b]
        np.add.at(s, k, np.maximum(np.nan_to_num(db, nan=-999.0), FLOOR_DB)[~b])
        np.add.at(n, k, 1)
    prof = s / np.maximum(n, 1)
    return prof


def train(a):
    t_start = time.time()
    torch.manual_seed(a.seed)
    rng = np.random.default_rng(a.seed)
    dev = pick_device(a.device)
    D = load_urban()
    part = split_areas([e + area_of(k) for e, k in zip(D["env"], D["keys"])], seed=0)
    sp = np.array([part[e + area_of(k)] for e, k in zip(D["env"], D["keys"])])
    X = inputs(D["bld"], D["terrain"], D["tx"], a.dist)
    Y = to_train_scale(D["db"])[:, None]
    idx = {s: np.where(sp == s)[0] for s in ("train", "val", "test")}
    sel = lambda ii: [D["tx"][i] for i in ii]
    model = build("unet", X.shape[1], a.base, out_ch=2 if a.gate else 1).to(dev)
    opt = torch.optim.Adam(model.parameters(), a.lr)
    total = a.samples // a.batch
    sched = torch.optim.lr_scheduler.StepLR(opt, step_size=max(1, int(a.lr_drop * total)), gamma=0.1)
    out = OUT / a.name
    out.mkdir(parents=True, exist_ok=True)
    cfg = {**vars(a), "maps": {s: int(len(v)) for s, v in idx.items()}, "areas": {s: len({area_of(D["keys"][i]) + D["env"][i] for i in v}) for s, v in idx.items()},
           "inputs": ["buildings 1.5 m", "Tx one-hot", f"terrain - terrain(Tx), /{TERRAIN_SCALE} m"] + (["log distance"] if a.dist else []),
           "total_steps": total, "encoding": {"gray_0db": GRAY_0DB, "gray_per_db": GRAY_PER_DB}, "floor_db": round(FLOOR_DB, 2),
           "top_db": TOP_DB, "commit": git_commit(), "torch": torch.__version__}
    (out / "config.json").write_text(json.dumps(cfg, indent=1))
    print(json.dumps(cfg), flush=True)
    Xt, Yt = torch.from_numpy(X[idx["train"]]), torch.from_numpy(Y[idx["train"]])
    log, best, step = [], math.inf, 0
    t0 = time.time()
    while step < total:
        perm = rng.permutation(len(Xt))
        for k in range(0, len(perm) - a.batch + 1, a.batch):
            b = perm[k:k + a.batch]
            x, y = augment(Xt[b], Yt[b], rng)
            p = model(x.to(dev))
            y = y.to(dev)
            loss = F.mse_loss(p[:, :1], y)
            if a.gate:
                loss = loss + a.gate_weight * F.binary_cross_entropy_with_logits(p[:, 1:2], (y > 0).float())
            opt.zero_grad(set_to_none=True)
            loss.backward()
            opt.step()
            sched.step()
            step += 1
            if step % a.val_every == 0 or step >= total:
                model.eval()
                v = evaluate(predict_db(model, X[idx["val"]], dev, a.gate), D["db"][idx["val"]], D["bld"][idx["val"]], sel(idx["val"]))
                model.train()
                rec = {"step": step, "epoch": round(step * a.batch / len(Xt), 1), "t_min": round((time.time() - t0) / 60, 2),
                       "train_loss": round(loss.item(), 6), "lr": opt.param_groups[0]["lr"],
                       "val_rmse_db_outdoor": v["rmse_db_outdoor_floor_clipped"], "val_rmse_db_covered": v["rmse_db_where_truth_covered"]}
                log.append(rec)
                print(json.dumps(rec), flush=True)
                if v["rmse_db_outdoor_floor_clipped"] < best:
                    best = v["rmse_db_outdoor_floor_clipped"]
                    torch.save(model.state_dict(), out / "best.pt")
                (out / "log.json").write_text(json.dumps(log))
            if step >= total:
                break
    train_min = (time.time() - t0) / 60
    model.load_state_dict(torch.load(out / "best.pt", map_location=dev))
    model.eval()
    te = idx["test"]
    pred = predict_db(model, X[te], dev, a.gate)
    res = {"name": a.name, "seed": a.seed, "commit": cfg["commit"], "config": cfg, "train_minutes": round(train_min, 1),
           "val_best_rmse_db_outdoor": best, "test": {"all": evaluate(pred, D["db"][te], D["bld"][te], sel(te))}}
    groups = {env: D["env"][te] == env for env in URBAN}
    groups["single_tx"] = np.array([len(D["tx"][i]) == 1 for i in te])
    for name, m in groups.items():
        res["test"][name] = evaluate(pred[m], D["db"][te][m], D["bld"][te][m], [t for t, k in zip(sel(te), m) if k])
    tr = idx["train"]
    prof = radial_baseline(D["db"][tr], D["bld"][tr], sel(tr))
    rad = np.stack([prof[np.minimum((dist_m(t) / 10).astype(int), 199)] for t in sel(te)]).astype(np.float32)
    res["test_radial_baseline"] = evaluate(np.where(rad > FLOOR_DB, rad, np.nan), D["db"][te], D["bld"][te], sel(te))
    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"]["all"]), flush=True)


def load_model(run, dev):
    cfg = json.loads((Path(run) / "config.json").read_text())
    model = build("unet", 3 + int(cfg["dist"]), cfg["base"], out_ch=2 if cfg.get("gate") else 1).to(dev)
    model.load_state_dict(torch.load(Path(run) / "best.pt", map_location=dev))
    return model.eval(), cfg


# ------------------------------------------------------------------------------------------------------------ mosaic

def mosaic(a):
    """Apply SpectrumNet U-Nets to the 2 km RadioMapSeer mosaic (runs/cloudrf/mosaic_2km_anchor) at 10 m and compare
    outdoor coverage by distance with the methods in compare.json (same 1 m outdoor mask, rings and floor)."""
    from .bigmap import mosaic as build_mosaic
    from .cloudrf_compare import RINGS as CRINGS, compare

    scene = json.loads((REPO / a.scene / "scene.json").read_text())
    bld, _, _, maps = build_mosaic(scene["k"])
    assert maps == scene["maps"]
    n, tx = scene["size_m"], tuple(scene["tx_rc"])
    f = int(PX)
    npad = int(math.ceil(n / (16 * f)) * 16 * f)  # 2080 m: multiple of 10 m and of the U-Net's 16 px
    bp = np.zeros((npad, npad), np.float32)
    bp[:n, :n] = bld
    b10 = bp.reshape(npad // f, f, npad // f, f).mean((1, 3)) >= 0.5
    tx10 = [np.array([[tx[0] // f, tx[1] // f]])]
    b10[tx[0] // f, tx[1] // f] = False
    dev = pick_device(a.device)
    rr, cc = np.mgrid[:n, :n]
    dist = np.hypot(rr - tx[0], cc - tx[1])
    outdoor = ~bld & (dist < a.radius)
    cmp = json.loads((REPO / a.scene / "compare.json").read_text())
    ref = None
    if a.compare_npz and Path(a.compare_npz).exists():
        ref = np.load(a.compare_npz)["Our tracer (calibrated)"].astype(np.float32)
    res = {"commit": git_commit(), "scene": a.scene, "floor_db": round(FLOOR_DB, 1), "radius_m": a.radius,
           "grid": f"buildings averaged to {f} m cells (>= 50 % = building), padded to {npad} m with open ground; prediction "
                   f"repeated to 1 m and scored on the 1 m outdoor mask",
           "building_pct_10m": round(100 * float(b10[:n // f, :n // f].mean()), 1), "building_pct_1m": round(100 * float(bld.mean()), 1),
           "f59_shift_db": round(F59_SHIFT_DB, 2), "methods": {},
           "reference_methods_from_compare_json": {name: {r: (v or {}).get("covered_pct") for r, v in rings.items()}
                                                   for name, rings in cmp["vs_tracer_by_distance"].items()}}
    maps_db = {}
    for run in a.runs.split(","):
        model, cfg = load_model(REPO / run, dev)
        x = inputs(b10[None], np.zeros((1, npad // f, npad // f), np.float32), tx10, cfg["dist"])
        t0 = time.time()
        p10 = predict_db(model, x, dev, cfg.get("gate", False))[0]
        ms = (time.time() - t0) * 1000
        p1 = np.repeat(np.repeat(p10, f, 0), f, 1)[:n, :n]
        for tag, shift in (("3.5 GHz", 0.0), ("5.9 GHz-equiv.", F59_SHIFT_DB)):
            name = f"SpectrumNet U-Net {cfg['name']} ({tag})"
            db = p1 - shift
            maps_db[name] = db
            res["methods"][name] = {"run": run, "ms": round(ms, 1), "covered_pct_by_distance": {
                f"{lo}-{hi} m": compare(db, db, outdoor & (dist >= lo) & (dist < hi))["covered_pct"] for lo, hi in CRINGS}}
            if ref is not None:
                res["methods"][name]["vs_tracer_by_distance"] = {
                    f"{lo}-{hi} m": compare(db, ref, outdoor & (dist >= lo) & (dist < hi)) for lo, hi in CRINGS}
    (OUT / "mosaic_2km.json").write_text(json.dumps(res, indent=1))
    print(json.dumps(res, indent=1))
    if a.fig:
        import matplotlib
        matplotlib.use("Agg")
        import matplotlib.pyplot as plt
        # gated runs: the 5.9 GHz-equivalent map has the same coverage (the gate decides), so it is not shown twice
        show = ({"Our tracer (calibrated), 5.9 GHz": ref} if ref is not None else {}) | {
            k: v for k, v in maps_db.items() if not ("gate" in k and "5.9" in k)}
        cols = 4
        rows = math.ceil(len(show) / cols)
        fig, axs = plt.subplots(rows, cols, figsize=(5.4 * cols, 5.6 * rows))
        for ax in axs.flat[len(show):]:
            ax.axis("off")
        for ax, (name, m) in zip(axs.flat, show.items()):
            ax.set_facecolor("black")
            im = ax.imshow(m, cmap="viridis", vmin=-160, vmax=-50, interpolation="nearest")
            ax.contour(np.nan_to_num(m, nan=-999), levels=[FLOOR_DB], colors="red", linewidths=0.4)
            ax.imshow(np.where(bld, 1.0, np.nan), cmap="Greys", vmin=0, vmax=1.6, interpolation="nearest")
            ax.plot(tx[1], tx[0], marker="*", color="#d62728", ms=12, mec="white")
            for rad in (200, 400, 800):
                ax.add_patch(plt.Circle((tx[1], tx[0]), rad, fill=False, ec="white", lw=0.6, ls="--"))
            ax.set_title(name, fontsize=9); ax.set_xticks([]); ax.set_yticks([])
        fig.colorbar(im, ax=axs.ravel().tolist(), fraction=0.012, pad=0.01).set_label("path gain, dB")
        plt.savefig(FIGS / "08_spectrumnet_2km.png", dpi=55, bbox_inches="tight")


def main():
    ap = argparse.ArgumentParser()
    sub = ap.add_subparsers(dest="cmd", required=True)
    p = sub.add_parser("inspect")
    p.add_argument("--fit-maps", type=int, default=150, help="maps per environment used for the encoding fit")
    p.add_argument("--h-grid", type=float, nargs="+", default=[0, 1.5, 3, 5, 10, 20, 30])
    p = sub.add_parser("train")
    p.add_argument("--name", required=True)
    p.add_argument("--dist", action="store_true", help="extra input: log distance to the Tx")
    p.add_argument("--gate", action="store_true", help="second output: 'covered' logit (BCE), prediction floored where negative")
    p.add_argument("--gate-weight", type=float, default=0.02, help="BCE weight next to the MSE (train.py default)")
    p.add_argument("--base", type=int, default=16)
    p.add_argument("--batch", type=int, default=15)
    p.add_argument("--lr", type=float, default=1e-4)
    p.add_argument("--lr-drop", type=float, default=0.6)
    p.add_argument("--samples", type=int, default=480960, help="training samples seen (runs/capped: 12 epochs x 40080)")
    p.add_argument("--val-every", type=int, default=1000)
    p.add_argument("--seed", type=int, default=0)
    p.add_argument("--device", default="mps")
    p = sub.add_parser("mosaic")
    p.add_argument("--runs", required=True, help="comma-separated run dirs (relative to the repo)")
    p.add_argument("--scene", default="runs/cloudrf/mosaic_2km_anchor")
    p.add_argument("--compare-npz", default=None, help="cloudrf_compare maps (compare.npz) for agreement with our tracer")
    p.add_argument("--radius", type=float, default=1100)
    p.add_argument("--device", default="mps")
    p.add_argument("--fig", action="store_true")
    a = ap.parse_args()
    {"inspect": inspect, "train": train, "mosaic": mosaic}[a.cmd](a)


if __name__ == "__main__":
    main()
