"""Learned ray tracing, stage 1: one ray's full trajectory (docs/research/07-learned-ray-tracing.md).

Teacher: rcm_ml.raypaths (our 2D tracer, specular reflections, up to B bounces). Two representations:

  step    autoregressive next-hit model. For the current leg (origin o, direction d) the building raster is
          sampled on a strip in the ray's own frame (forward bins of 0.5 m from -2 m to the tile edge, 13 lanes
          +-3 m across). A small CNN gives, per forward bin, a hit hazard h_i, a sub-bin offset and the wall
          normal (in the ray frame). P(first hit in bin i) = s(h_i) prod_{j<i} (1 - s(h_j)) (transmittance, as
          in volume rendering); the ray leaves the tile if it survives every bin before the edge. The next
          direction is the mirror of d about the predicted normal. Translation/rotation equivariant by
          construction, cost independent of the map size.
  direct  one CNN pass over the whole tile (occupancy, Tx, launch angle as channels) regresses all B+1 hit
          points and leg types at once (the "image in, polyline out" alternative).

  python -m rcm_ml.raynet data                         # trajectories + 0.5 m occupancy rasters -> data/raynet
  python -m rcm_ml.raynet train --model step   --minutes 20 --out runs/raynet/s1_step
  python -m rcm_ml.raynet train --model direct --minutes 20 --out runs/raynet/s1_direct
"""
from __future__ import annotations

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

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

from .raypaths import EXIT, REFLECT, STOP, pad_paths, scene, trace_rays, tx_xy

DATA = Path(os.environ.get("RAYNET_DATA", Path(__file__).resolve().parent.parent / "data/raynet"))
RES = 0.5          # m per occupancy pixel
NOCC = 512         # 256 m tile at 0.5 m
T0, DT, NB = -2.0, 0.5, 736  # forward bins t_i = T0 + i DT (covers the 362 m diagonal)
LANES = 13         # lateral samples, DT apart, centred on the ray
TOL = (0.5, 1.0, 2.0)


# ---------------------------------------------------------------- data

def occupancy(polys, n=NOCC, ss=4):
    """Anti-aliased building coverage at RES m (ss x ss supersampling), uint8 0..255. Pixel j covers [j, j+1) * RES."""
    k = n * ss
    im = Image.new("L", (k, k))
    d = ImageDraw.Draw(im)
    s = k / 256.0
    for p in polys:
        d.polygon([(x * s - 0.5, y * s - 0.5) for x, y in p], fill=255)
    a = np.asarray(im, np.float32).reshape(n, ss, n, ss).mean((1, 3))
    return np.round(a).astype(np.uint8)


def ray_dirs_normals(ang, pts, kind, wall, segs):
    """Direction of every leg and the hit wall's unit normal, oriented towards the incoming ray."""
    R, K = kind.shape
    d = np.zeros((R, K, 2), np.float32)
    nrm = np.zeros((R, K, 2), np.float32)
    dx, dy = np.cos(ang), np.sin(ang)
    for k in range(K):
        d[:, k, 0], d[:, k, 1] = dx, dy
        w = wall[:, k]
        ok = w >= 0
        e = segs[np.maximum(w, 0)]
        ex, ey = e[:, 2] - e[:, 0], e[:, 3] - e[:, 1]
        L = np.hypot(ex, ey)
        nx, ny = -ey / L, ex / L
        flip = nx * dx + ny * dy > 0
        nx, ny = np.where(flip, -nx, nx), np.where(flip, -ny, ny)
        nrm[:, k, 0], nrm[:, k, 1] = np.where(ok, nx, 0), np.where(ok, ny, 0)
        dot = dx * nx + dy * ny
        dx, dy = np.where(ok, dx - 2 * dot * nx, dx), np.where(ok, dy - 2 * dot * ny, dy)
    return d, nrm


def build_split(maps, n_tx, n_ang, bounces, rng, tx_ids=None):
    from .data import load_tx
    from .raytrace import load_polygons
    out = {k: [] for k in ("map", "tx", "ang", "pts", "kind", "wall", "cosi", "dir", "nrm")}
    for m in maps:
        polys = load_polygons(int(m))
        _, segs, _ = scene(polys)
        tids = tx_ids if tx_ids is not None else rng.choice(80, n_tx, replace=False)
        for t in tids:
            tx = tx_xy(load_tx(int(m), int(t)))
            ang = rng.uniform(0, 2 * np.pi, n_ang)
            pts, kind, wall, cosi, _ = pad_paths(trace_rays(segs, tx, ang, bounces), bounces)
            d, nrm = ray_dirs_normals(ang, pts, kind, wall, segs)
            for k, v in (("map", np.full(n_ang, m)), ("tx", np.tile(np.float32(tx), (n_ang, 1))), ("ang", ang.astype(np.float32)),
                         ("pts", pts), ("kind", kind), ("wall", wall), ("cosi", cosi), ("dir", d), ("nrm", nrm)):
                out[k].append(v)
    return {k: np.concatenate(v) for k, v in out.items()}


def make_data(out=DATA, bounces=4, seed=0):
    from .data import split
    from .raytrace import load_polygons
    out.mkdir(parents=True, exist_ok=True)
    occ = np.zeros((701, NOCC, NOCC), np.uint8)
    for m in range(701):
        try:
            occ[m] = occupancy(load_polygons(m))
        except FileNotFoundError:
            pass
    np.save(out / "occ.npy", occ)
    tr, va, te = split(official=True)
    rng = np.random.default_rng(seed)
    cfg = {"train": (tr, 16, 32, None), "val": (va, 2, 32, None), "test": (te, 2, 64, (0, 1))}
    meta = {"bounces": bounces, "seed": seed}
    for name, (maps, ntx, nang, tids) in cfg.items():
        t0 = time.time()
        D = build_split(maps, ntx, nang, bounces, rng, tids)
        np.savez(out / f"{name}.npz", **D)
        meta[name] = {"maps": len(maps), "tx_per_map": ntx, "angles_per_tx": nang, "rays": int(len(D["map"])), "seconds": round(time.time() - t0, 1)}
        print(name, meta[name], flush=True)
    # raster alignment: building coverage on the ray at the true hit points of test legs, and +-0.25 / 0.5 m around them
    D = truncate(dict(np.load(out / "test.npz")), bounces)
    Lg = leg_table(D)
    h = np.nonzero(Lg["cls"] >= 0)[0]
    oc = torch.from_numpy(occ).float() / 255
    g = lambda k: torch.from_numpy(Lg[k][h])
    t = T0 + DT * (g("cls").float() + g("off"))
    meta["coverage_at_hit"] = {}
    for dt in (-0.5, -0.25, 0.0, 0.25, 0.5):
        p = g("o") + (t + dt)[:, None] * g("d")
        z = torch.zeros(1)
        c = sample_frame(oc, g("map"), p, g("d"), z, z)[:, 0, 0, 0]
        meta["coverage_at_hit"][str(dt)] = float(c.mean())
    from .train import git_commit
    meta["commit"] = git_commit()
    (out / "meta.json").write_text(json.dumps(meta, indent=1))
    Path("runs/raynet").mkdir(parents=True, exist_ok=True)
    Path("runs/raynet/data.json").write_text(json.dumps(meta, indent=1))


def truncate(D, bounces):
    """Trajectories of a B'-bounce tracer from B-bounce ones (B' <= B): same legs, the last allowed hit becomes STOP."""
    K = bounces + 1
    E = {k: D[k] for k in ("map", "tx", "ang")}
    for k in ("kind", "wall", "cosi", "dir", "nrm"):
        E[k] = D[k][:, :K].copy()
    E["pts"] = D["pts"][:, : K + 1].copy()
    E["kind"][E["kind"][:, -1] == REFLECT, -1] = STOP
    return E


# ---------------------------------------------------------------- strip sampling (shared by train and rollout)

def exit_t(o, d, n=256.0):
    big = torch.full_like(o[:, 0], 1e9)
    tx = torch.where(d[:, 0] > 1e-9, (n - o[:, 0]) / d[:, 0], torch.where(d[:, 0] < -1e-9, -o[:, 0] / d[:, 0], big))
    ty = torch.where(d[:, 1] > 1e-9, (n - o[:, 1]) / d[:, 1], torch.where(d[:, 1] < -1e-9, -o[:, 1] / d[:, 1], big))
    return torch.minimum(tx, ty).clamp(min=0)


def sample_frame(occ, mid, o, d, t, u):
    """Bilinear building coverage at o + t d + u d_perp (d_perp = d rotated +90 deg) for the grid t (T,) x u (U,).
    occ (M, NOCC, NOCC) float in [0, 1] on device; mid (B,) map index; o, d (B, 2) in metres.
    Returns (B, 2, U, T): coverage and an outside-the-tile indicator."""
    px = o[:, 0, None, None] + t[None, None, :] * d[:, 0, None, None] - u[None, :, None] * d[:, 1, None, None]
    py = o[:, 1, None, None] + t[None, None, :] * d[:, 1, None, None] + u[None, :, None] * d[:, 0, None, None]
    out = ((px < 0) | (py < 0) | (px >= 256) | (py >= 256)).float()
    gx, gy = px / RES - 0.5, py / RES - 0.5
    x0, y0 = torch.floor(gx), torch.floor(gy)
    wx, wy = gx - x0, gy - y0
    flat = occ.reshape(-1)
    base = (mid.long() * NOCC * NOCC)[:, None, None]
    acc = torch.zeros_like(px)
    for ddx, ddy, w in ((0, 0, (1 - wx) * (1 - wy)), (1, 0, wx * (1 - wy)), (0, 1, (1 - wx) * wy), (1, 1, wx * wy)):
        xi, yi = (x0 + ddx).long(), (y0 + ddy).long()
        ok = (xi >= 0) & (yi >= 0) & (xi < NOCC) & (yi < NOCC)
        idx = base + yi.clamp(0, NOCC - 1) * NOCC + xi.clamp(0, NOCC - 1)
        acc = acc + w * ok * flat[idx.reshape(-1)].reshape(px.shape)
    return torch.stack([acc, out], 1)


def sample_strip(occ, mid, o, d):
    """(B, 2, LANES, NB): the ray-frame strip of the step model."""
    dev = o.device
    t = T0 + DT * torch.arange(NB, device=dev, dtype=torch.float32)
    u = DT * (torch.arange(LANES, device=dev, dtype=torch.float32) - LANES // 2)
    return sample_frame(occ, mid, o, d, t, u)


PATCH, PRES = 32, 0.25  # normal-refinement patch: 32 x 32 samples at 0.25 m (8 m), centred on the hit, ray frame


def sample_patch(occ, mid, p, d):
    g = PRES * (torch.arange(PATCH, device=p.device, dtype=torch.float32) - (PATCH - 1) / 2)
    return sample_frame(occ, mid, p, d, g, g)


def bin_mask(o, d):
    """Valid hit bins: t_i >= 0 and before the tile edge. Also returns the exit distance."""
    te = exit_t(o, d)
    t = T0 + DT * torch.arange(NB, device=o.device, dtype=torch.float32)
    return (t[None] >= -DT / 2) & (t[None] <= te[:, None] + DT / 2), te


# ---------------------------------------------------------------- models

class PatchNet(nn.Module):
    """Wall normal (ray frame) from an 8 m patch around the hit point."""

    def __init__(self, ch=32):
        super().__init__()
        c = lambda i, o, s=1: [nn.Conv2d(i, o, 3, stride=s, padding=1), nn.GELU()]
        self.f = nn.Sequential(*c(2, ch), *c(ch, ch), *c(ch, 2 * ch, 2), *c(2 * ch, 2 * ch), *c(2 * ch, 2 * ch, 2), *c(2 * ch, 2 * ch),
                               nn.Flatten(), nn.Linear(2 * ch * (PATCH // 4) ** 2, 256), nn.GELU(), nn.Linear(256, 2))

    def forward(self, patch):
        return self.f(patch)


class StepNet(nn.Module):
    """Lanes are channels: (2 x LANES) -> ch, then residual dilated 1D convolutions along the ray.
    patch=True adds PatchNet, which replaces the strip's normal at inference."""

    def __init__(self, ch=64, dil=(1, 2, 4, 8, 1, 2), patch=False):
        super().__init__()
        self.inp = nn.Sequential(nn.Conv1d(2 * LANES, ch, 5, padding=2), nn.GELU(), nn.Conv1d(ch, ch, 5, padding=2), nn.GELU())
        self.c1d = nn.ModuleList([nn.Conv1d(ch, ch, 5, padding=2 * k, dilation=k) for k in dil])
        self.head = nn.Conv1d(ch, 4, 1)  # hazard logit, offset, normal (2, ray frame)
        self.patch = PatchNet() if patch else None

    def forward(self, strip):
        h = self.inp(strip.flatten(1, 2))
        for c in self.c1d:
            h = h + F.gelu(c(h))
        return self.head(h)  # (B, 4, NB)


def step_logp(haz, valid):
    """log P(first hit in bin i) for every bin, and log P(no hit before the edge)."""
    haz = haz.masked_fill(~valid, -30.0)
    ls_hit, ls_miss = F.logsigmoid(haz), F.logsigmoid(-haz)
    cum = torch.cumsum(ls_miss, 1) - ls_miss  # sum over j < i
    lp = (ls_hit + cum).masked_fill(~valid, -1e4)
    return lp, ls_miss.sum(1)


def ray_frame(v, d):
    """World vector -> (along d, across d) with across = d rotated +90 deg."""
    return torch.stack([v[:, 0] * d[:, 0] + v[:, 1] * d[:, 1], -v[:, 0] * d[:, 1] + v[:, 1] * d[:, 0]], 1)


def world(v, d):
    return torch.stack([v[:, 0] * d[:, 0] - v[:, 1] * d[:, 1], v[:, 0] * d[:, 1] + v[:, 1] * d[:, 0]], 1)


def step_predict(net, occ, mid, o, d):
    """One leg: (kind EXIT/hit, hit distance t, unit normal in world frame). kind is EXIT or REFLECT (caller turns
    the last REFLECT into STOP)."""
    out = net(sample_strip(occ, mid, o, d))
    valid, te = bin_mask(o, d)
    lp, lsurv = step_logp(out[:, 0], valid)
    best, i = lp.max(1)
    ex = lsurv > best
    off = 0.5 * torch.tanh(out[:, 1].gather(1, i[:, None]))[:, 0]
    t = torch.where(ex, te, T0 + DT * (i.float() + off)).clamp(min=1e-3)
    nl = out[:, 2:4].gather(2, i[:, None, None].expand(-1, 2, 1))[..., 0]
    if net.patch is not None:
        nl = net.patch(sample_patch(occ, mid, o + t[:, None] * d, d))
    n = world(F.normalize(nl, dim=1), d)
    return ex, t, n


class DirectNet(nn.Module):
    """CNN over the 256 m tile at 1 m: [occupancy, Tx gaussian, cos a, sin a, x, y] -> K legs x (x, y, 4 kind logits)."""

    def __init__(self, K, ch=(32, 64, 128, 128, 256, 256), hid=1024):
        super().__init__()
        L, c0 = [], 6
        for c in ch:
            L += [nn.Conv2d(c0, c, 3, stride=2, padding=1), nn.BatchNorm2d(c), nn.GELU(), nn.Conv2d(c, c, 3, padding=1), nn.BatchNorm2d(c), nn.GELU()]
            c0 = c
        self.enc = nn.Sequential(*L)
        self.K = K
        self.fc = nn.Sequential(nn.Flatten(), nn.Linear(c0 * 16, hid), nn.GELU(), nn.Linear(hid + 4, hid), nn.GELU(), nn.Linear(hid, K * 6))

    def forward(self, occ256, tx, ang):
        B = tx.shape[0]
        dev = tx.device
        g = torch.arange(256, device=dev, dtype=torch.float32) + 0.5
        X, Y = g[None, None, :].expand(B, 256, 256), g[None, :, None].expand(B, 256, 256)
        txg = torch.exp(-((X - tx[:, 0, None, None]) ** 2 + (Y - tx[:, 1, None, None]) ** 2) / 8.0)
        ca, sa = torch.cos(ang)[:, None, None].expand(B, 256, 256), torch.sin(ang)[:, None, None].expand(B, 256, 256)
        x = torch.stack([occ256, txg, ca, sa, X / 128 - 1, Y / 128 - 1], 1)
        h = self.enc(x)
        h = self.fc[2](self.fc[1](self.fc[0](h)))
        h = torch.cat([h, torch.stack([torch.cos(ang), torch.sin(ang), tx[:, 0] / 128 - 1, tx[:, 1] / 128 - 1], 1)], 1)
        y = self.fc[5](self.fc[4](self.fc[3](h))).view(B, self.K, 6)
        return 128 * (y[..., :2] + 1), y[..., 2:]  # points in metres, kind logits (REFLECT, EXIT, STOP, none)


# ---------------------------------------------------------------- evaluation

def endpoint_metrics(pred_pts, pred_kind, D, tag):
    """pred_pts (R, K+1, 2), pred_kind (R, K) with -1 = no leg; D the reference trajectories (same K)."""
    tk, tp = D["kind"], D["pts"]
    R, K = tk.shape
    res = {}
    seq_ok = {tau: np.ones(R, bool) for tau in TOL}
    for k in range(K):
        has = tk[:, k] >= 0
        kind_ok = pred_kind[:, k] == tk[:, k]
        err = np.linalg.norm(pred_pts[:, k + 1] - tp[:, k + 1], axis=1)
        e = err[has & kind_ok]
        res[f"leg{k}"] = {"rays": int(has.sum()), "kind_acc": float(kind_ok[has].mean()) if has.any() else None,
                          "err_m_median": float(np.median(e)) if len(e) else None, "err_m_p90": float(np.percentile(e, 90)) if len(e) else None,
                          "err_m_mean": float(e.mean()) if len(e) else None,
                          **{f"within_{tau}m": float(((err < tau) & kind_ok)[has].mean()) if has.any() else None for tau in TOL}}
        for tau in TOL:
            seq_ok[tau] &= np.where(has, kind_ok & (err < tau), pred_kind[:, k] < 0)
    res["exact_match"] = {f"{tau}m": float(seq_ok[tau].mean()) for tau in TOL}
    return {tag: res}


@torch.no_grad()
def rollout_step(net, occ, D, bounces, dev, teacher=False, bs=4096):
    """Autoregressive trajectories. teacher=True: every leg starts from the reference origin and direction
    (one-step accuracy); else from the model's own previous prediction."""
    R = len(D["map"])
    K = bounces + 1
    P = np.full((R, K + 1, 2), np.nan, np.float32)
    Kd = np.full((R, K), -1, np.int8)
    for s in range(0, R, bs):
        sl = slice(s, s + bs)
        mid = torch.from_numpy(D["map"][sl]).to(dev)
        o = torch.from_numpy(D["tx"][sl]).to(dev)
        d = torch.stack([torch.cos(torch.from_numpy(D["ang"][sl])), torch.sin(torch.from_numpy(D["ang"][sl]))], 1).to(dev)
        alive = torch.ones(len(o), dtype=torch.bool, device=dev)
        P[sl, 0] = o.cpu().numpy()
        for k in range(K):
            if teacher:
                alive = torch.from_numpy(D["kind"][sl, k] >= 0).to(dev)
                o = torch.from_numpy(np.nan_to_num(D["pts"][sl, k])).to(dev)
                d = torch.from_numpy(D["dir"][sl, k]).to(dev)
            ex, t, n = step_predict(net, occ, mid, o, d)
            p = o + t[:, None] * d
            kind = torch.where(ex, torch.full_like(t, EXIT), torch.full_like(t, REFLECT if k < K - 1 else STOP)).long()
            kind = torch.where(alive, kind, torch.full_like(kind, -1))
            P[sl, k + 1] = torch.where(alive[:, None], p, torch.full_like(p, float("nan"))).cpu().numpy()
            Kd[sl, k] = kind.cpu().numpy()
            dot = (d * n).sum(1, keepdim=True)
            o, d = p, F.normalize(d - 2 * dot * n, dim=1)
            alive = alive & ~ex
    return P, Kd


@torch.no_grad()
def predict_direct(net, occ256, D, dev, bs=256):
    R = len(D["map"])
    Ps, Ks = [], []
    for s in range(0, R, bs):
        sl = slice(s, s + bs)
        mid = torch.from_numpy(D["map"][sl]).to(dev).long()
        p, kl = net(occ256[mid], torch.from_numpy(D["tx"][sl]).to(dev), torch.from_numpy(D["ang"][sl]).to(dev))
        Ps.append(p.cpu().numpy())
        Ks.append(kl.argmax(-1).cpu().numpy())
    p, kc = np.concatenate(Ps), np.concatenate(Ks)
    kind = np.where(kc == 3, -1, kc).astype(np.int8)
    # a ray ends at its first EXIT/STOP: later legs are dropped
    for k in range(1, kind.shape[1]):
        dead = (kind[:, k - 1] != REFLECT)
        kind[dead, k] = -1
    P = np.concatenate([D["tx"][:, None], p], 1)
    return P, kind


# ---------------------------------------------------------------- training

def leg_table(D):
    """Flatten trajectories into legs: map, origin, dir, target class (bin or -1 = exit), offset, normal (ray frame)."""
    k = D["kind"] >= 0
    r, j = np.nonzero(k)
    o = D["pts"][r, j]
    p = D["pts"][r, j + 1]
    d = D["dir"][r, j]
    t = np.linalg.norm(p - o, axis=1)
    kind = D["kind"][r, j]
    b = np.round((t - T0) / DT).astype(np.int64)
    off = (t - T0) / DT - b
    n = D["nrm"][r, j]
    nl = np.stack([n[:, 0] * d[:, 0] + n[:, 1] * d[:, 1], -n[:, 0] * d[:, 1] + n[:, 1] * d[:, 0]], 1)
    cls = np.where(kind == EXIT, -1, b)
    return {"map": D["map"][r], "o": o.astype(np.float32), "d": d, "cls": cls, "off": off.astype(np.float32), "nl": nl.astype(np.float32)}


@torch.no_grad()
def normal_error(net, occ, D, dev, bs=2048):
    """Angle (deg) between predicted and true wall normal at the true hit points: from the strip head and, if
    present, from the patch net."""
    L = leg_table(D)
    hit = np.nonzero(L["cls"] >= 0)[0]
    errs = {"strip": [], "patch": []}
    for s in range(0, len(hit), bs):
        idx = hit[s:s + bs]
        g = lambda k: torch.from_numpy(L[k][idx]).to(dev)
        o, d, ci, nl, mid = g("o"), g("d"), g("cls"), g("nl"), g("map")
        out = net(sample_strip(occ, mid, o, d))
        cands = {"strip": out[:, 2:4].gather(2, ci[:, None, None].expand(-1, 2, 1))[..., 0]}
        if net.patch is not None:
            p = o + (T0 + DT * (ci.float() + g("off")))[:, None] * d
            cands["patch"] = net.patch(sample_patch(occ, mid, p, d))
        for k, v in cands.items():
            c = (F.normalize(v, dim=1) * nl).sum(1).clamp(-1, 1)
            errs[k].append(torch.rad2deg(torch.acos(c)).cpu().numpy())
    return {k: {f"p{q}": float(np.percentile(np.concatenate(v), q)) for q in (50, 75, 90, 99)} for k, v in errs.items() if v}


def step_loss(net, occ, L, idx, dev, jitter=0.1):
    g = lambda k: torch.from_numpy(L[k][idx]).to(dev)
    o, d, cls, off, nl = g("o"), g("d"), g("cls"), g("off"), g("nl")
    out = net(sample_strip(occ, g("map"), o, d))
    valid, _ = bin_mask(o, d)
    lp, lsurv = step_logp(out[:, 0], valid)
    hit = cls >= 0
    ci = cls.clamp(min=0)
    nll = -torch.where(hit, lp.gather(1, ci[:, None])[:, 0], lsurv)
    po = 0.5 * torch.tanh(out[:, 1].gather(1, ci[:, None])[:, 0])
    pn = F.normalize(out[:, 2:4].gather(2, ci[:, None, None].expand(-1, 2, 1))[..., 0], dim=1)
    l_off = ((po - off).abs() * hit).sum() / hit.sum().clamp(min=1)
    l_n = ((pn - nl).abs().sum(1) * hit).sum() / hit.sum().clamp(min=1)
    loss = nll.mean() + 4 * l_off + 4 * l_n
    parts = {"nll": nll.mean().item(), "off": l_off.item(), "n": l_n.item()}
    if net.patch is not None and hit.any():
        h = torch.nonzero(hit)[:, 0]
        t = T0 + DT * (ci[h].float() + off[h]) + jitter * torch.randn(len(h), device=dev)
        p = o[h] + t[:, None] * d[h]
        qn = F.normalize(net.patch(sample_patch(occ, g("map")[h], p, d[h])), dim=1)
        l_p = (qn - nl[h]).abs().sum(1).mean()
        loss = loss + 4 * l_p
        parts["patch_n"] = l_p.item()
    return loss, parts


def direct_loss(net, occ256, D, idx, dev):
    mid = torch.from_numpy(D["map"][idx]).to(dev).long()
    p, kl = net(occ256[mid], torch.from_numpy(D["tx"][idx]).to(dev), torch.from_numpy(D["ang"][idx]).to(dev))
    tk = torch.from_numpy(D["kind"][idx].astype(np.int64)).to(dev)
    tgt = torch.where(tk < 0, torch.full_like(tk, 3), tk)
    tp = torch.from_numpy(np.nan_to_num(D["pts"][idx, 1:])).to(dev)
    has = (tk >= 0).float()
    l_pts = ((F.smooth_l1_loss(p, tp, reduction="none", beta=1.0).sum(-1)) * has).sum() / has.sum()
    l_k = F.cross_entropy(kl.reshape(-1, 4), tgt.reshape(-1))
    return l_pts / 10 + l_k, {"pts": l_pts.item(), "kind": l_k.item()}


def load_split(name, bounces):
    D = dict(np.load(DATA / f"{name}.npz"))
    return truncate(D, bounces)


def main_train(a):
    from .train import git_commit
    torch.manual_seed(a.seed)
    rng = np.random.default_rng(a.seed)
    dev = torch.device(a.device)
    occ = torch.from_numpy(np.load(DATA / "occ.npy")).to(dev).float() / 255
    tr = load_split("train", a.train_bounces if a.model == "step" else a.bounces)
    va, te = load_split("val", a.bounces), load_split("test", a.bounces)
    if a.model == "step":
        net = StepNet(patch=a.patch).to(dev)
        L = leg_table(tr)
        n_items = len(L["map"])
        loss_fn = lambda idx: step_loss(net, occ, L, idx, dev)
        evaluate = lambda D, teacher=False: endpoint_metrics(*rollout_step(net, occ, D, a.bounces, dev, teacher), D, "x")["x"]
    else:
        net = DirectNet(a.bounces + 1).to(dev)
        occ256 = F.avg_pool2d(occ[:, None], 2)[:, 0]
        n_items = len(tr["map"])
        loss_fn = lambda idx: direct_loss(net, occ256, tr, idx, dev)
        evaluate = lambda D, teacher=False: endpoint_metrics(*predict_direct(net, occ256, D, dev), D, "x")["x"]
    n_par = sum(p.numel() for p in net.parameters())
    print(a.model, "params", n_par, "train items", n_items, flush=True)
    opt = torch.optim.AdamW(net.parameters(), lr=a.lr, weight_decay=1e-4)
    out = Path(a.out)
    out.mkdir(parents=True, exist_ok=True)
    log, best, step, t0 = [], -1.0, 0, time.time()
    total = a.minutes * 60
    next_eval = a.eval_every
    while time.time() - t0 < total:
        perm = rng.permutation(n_items)
        for s in range(0, n_items - a.batch + 1, a.batch):
            frac = (time.time() - t0) / total
            if frac >= 1:
                break
            for gp in opt.param_groups:
                gp["lr"] = a.lr * 0.5 * (1 + np.cos(np.pi * min(frac, 1)))
            net.train()
            loss, parts = loss_fn(np.sort(perm[s:s + a.batch]))
            opt.zero_grad()
            loss.backward()
            opt.step()
            step += 1
            if time.time() - t0 >= next_eval or time.time() - t0 >= total:
                next_eval += a.eval_every
                net.eval()
                v = evaluate(va)
                em = v["exact_match"]["1.0m"]
                log.append({"step": step, "seconds": round(time.time() - t0), "samples": step * a.batch, "loss": float(loss), **parts, "val_exact_1m": em})
                print(log[-1], flush=True)
                if em > best:
                    best = em
                    torch.save(net.state_dict(), out / "best.pt")
    net.load_state_dict(torch.load(out / "best.pt", map_location=dev))
    net.eval()
    t1 = time.time()
    rollout = evaluate(te)
    if dev.type == "mps":
        torch.mps.synchronize()
    t_eval = time.time() - t1
    res = {"commit": git_commit(), "model": a.model, "params": n_par, "config": vars(a), "train_items": n_items, "steps": step,
           "train_seconds": round(time.time() - t0 - t_eval), "test_rays": int(len(te["map"])), "bounces": a.bounces,
           "test_rollout": rollout, "test_eval_seconds": round(t_eval, 2), "val_best_exact_1m": best, "log": log}
    if a.model == "step":
        res["test_teacher_forced"] = evaluate(te, True)
        res["test_normal_err_deg"] = normal_error(net, occ, te, dev)
        for b in (4,):
            te4 = load_split("test", b)
            res[f"test_rollout_{b}_bounces"] = endpoint_metrics(*rollout_step(net, occ, te4, b, dev), te4, "x")["x"]
    (out / "test.json").write_text(json.dumps(res, indent=1))
    print(json.dumps({k: res[k] for k in ("test_rollout",)}, indent=1))


def main_eval(a):
    """Extra metrics for a trained step run (normal error at the true hits) -> <run>/eval.json."""
    from .train import git_commit
    dev = torch.device(a.device)
    cfg = json.loads((Path(a.run) / "test.json").read_text())["config"]
    occ = torch.from_numpy(np.load(DATA / "occ.npy")).to(dev).float() / 255
    net = StepNet(patch=cfg.get("patch", False)).to(dev)
    net.load_state_dict(torch.load(Path(a.run) / "best.pt", map_location=dev))
    net.eval()
    res = {"commit": git_commit(), "run": a.run, "test_normal_err_deg": normal_error(net, occ, load_split("test", 2), dev)}
    (Path(a.run) / "eval.json").write_text(json.dumps(res, indent=1))
    print(res)


def main():
    ap = argparse.ArgumentParser()
    sub = ap.add_subparsers(dest="cmd", required=True)
    d = sub.add_parser("data")
    d.add_argument("--bounces", type=int, default=4)
    t = sub.add_parser("train")
    t.add_argument("--model", choices=("step", "direct"), required=True)
    t.add_argument("--bounces", type=int, default=2, help="reflections per ray at evaluation (the tracer uses 2)")
    t.add_argument("--train-bounces", type=int, default=4, help="step model: legs from rays with up to this many reflections")
    t.add_argument("--minutes", type=float, default=20)
    t.add_argument("--eval-every", type=float, default=120, help="seconds")
    t.add_argument("--batch", type=int, default=256)
    t.add_argument("--lr", type=float, default=1e-3)
    t.add_argument("--seed", type=int, default=0)
    t.add_argument("--device", default="mps")
    t.add_argument("--patch", action="store_true", help="step model: refine the wall normal from an 8 m patch at the hit")
    t.add_argument("--out", required=True)
    e = sub.add_parser("eval")
    e.add_argument("--run", required=True)
    e.add_argument("--device", default="mps")
    a = ap.parse_args()
    if a.cmd == "data":
        make_data(bounces=a.bounces)
    elif a.cmd == "eval":
        main_eval(a)
    else:
        main_train(a)


if __name__ == "__main__":
    main()
