"""
MRI 2-slice observation geometry vs low-rank DVF reconstruction error.
Question: does orthogonal (sag+cor) beat parallel-2-sagittal under equal
2-slice budget, for a PCA/low-rank 3D motion model fit from partial slice
observations? See NOTES.md / task plan for full design.
"""
import sys
import json
import time
import numpy as np

from mriio import load_motionfield

np.random.seed(0)  # not used for anything stochastic outside seeded rngs below

AXIS_NAMES = ["ap", "si", "lr"]  # axis0=AP(x), axis1=SI(y), axis2=LR(z)
AP, SI, LR = 0, 1, 2

DATA_DIR = "data"
SUBJECTS = {
    "A": {"cycles": [1, 2, 3, 4], "n_points_expected": 8219, "ref_cycle": 3},
    "B": {"cycles": [1, 2, 3, 4, 5], "n_points_expected": 10125, "ref_cycle": 3},
}
K_MAX = 8
K_DEFAULT = 3
SPACINGS = [10, 20, 30, 50]
SIGMA_SAG_DEFAULT = 0.5
ALPHA_DEFAULT = 1.0
SEEDS = [0, 1, 2]
ROI_RADIUS_MM = 20.0
LAYER_TOL_MM = 2.5  # half of 5mm grid spacing
RIDGE_REL = 1e-6


def log(msg):
    print(f"[{time.time()-T0:7.1f}s] {msg}", flush=True)


def axis_calibration_gate():
    """Independent sanity check on ds15 subject-A data: axis assignment and
    reference-vs-own-t0 offset magnitudes must match plan anchors (approx)."""
    log("=== axis calibration gate (ds15, subject A) ===")
    cycs = {c: load_motionfield(f"{DATA_DIR}/motionfield_subjA_cyc{c}_ds15.txt")
            for c in [1, 2, 3, 4]}
    ref = cycs[3]["pos"][:, 0, :]
    span_ref = ref.max(axis=0) - ref.min(axis=0)
    allpos = np.concatenate([cycs[c]["pos"].reshape(-1, 3) for c in cycs], axis=0)
    span_all = allpos.max(axis=0) - allpos.min(axis=0)
    maxdisp = np.zeros(3)
    for c, d in cycs.items():
        disp = d["pos"] - ref[:, None, :]
        maxdisp = np.maximum(maxdisp, np.abs(disp).max(axis=(0, 1)))
    own_ref = cycs[1]["pos"][:, 0, :]
    off = own_ref - ref
    mean_off = np.abs(off).mean(axis=0).sum() / 3.0  # scalar mean like plan's "3.81mm"
    mean_off_3d = np.linalg.norm(off, axis=1).mean()
    max_off_3d = np.linalg.norm(off, axis=1).max()
    for i, name in enumerate(AXIS_NAMES):
        log(f"axis{i}({name}): ref-span={span_ref[i]:.2f}mm all-span={span_all[i]:.2f}mm "
            f"max|disp|={maxdisp[i]:.2f}mm")
    log(f"cyc1-own-t0 vs cyc3-ref-t0: mean_3d_offset={mean_off_3d:.2f}mm "
        f"max_3d_offset={max_off_3d:.2f}mm (plan anchor: mean 3.81mm, max 9.46mm)")

    # hard invariant: SI (axis1) must be the dominant axis both in spatial
    # span and in displacement magnitude (respiratory motion physiology)
    assert span_ref[SI] > span_ref[AP] and span_ref[SI] > span_ref[LR], \
        "axis1 not the largest-span axis -- axis mapping likely wrong"
    assert maxdisp[SI] > maxdisp[AP] and maxdisp[SI] > maxdisp[LR], \
        "axis1 not the largest-displacement axis -- axis mapping likely wrong"
    # approximate magnitude checks vs plan anchors (loose tolerance, "约")
    assert 120 * 0.85 <= span_ref[AP] <= 120 * 1.15, span_ref[AP]
    assert 156 * 0.80 <= span_ref[SI] <= 156 * 1.20 or 156 * 0.80 <= span_all[SI] <= 156 * 1.20, (span_ref[SI], span_all[SI])
    assert 77 * 0.80 <= span_ref[LR] <= 77 * 1.20 or 77 * 0.80 <= span_all[LR] <= 77 * 1.20, (span_ref[LR], span_all[LR])
    assert 17.8 * 0.75 <= maxdisp[SI] <= 17.8 * 1.30, maxdisp[SI]
    assert abs(mean_off_3d - 3.81) < 1.0 and abs(max_off_3d - 9.46) < 2.0, (mean_off_3d, max_off_3d)
    log("PASS: axis calibration gate ok\n")


def load_subject(name):
    info = SUBJECTS[name]
    cyc_data = {}
    for c in info["cycles"]:
        d = load_motionfield(f"{DATA_DIR}/motionfield_subj{name}_cyc{c}_ds5.txt")
        assert d["n_points"] == info["n_points_expected"], (name, c, d["n_points"])
        cyc_data[c] = d
    ref = cyc_data[info["ref_cycle"]]["pos"][:, 0, :].copy()
    disp = {c: cyc_data[c]["pos"] - ref[:, None, :] for c in cyc_data}  # (N,T,3)
    return {"ref": ref, "disp": disp, "n_points": info["n_points_expected"],
            "cycles": info["cycles"]}


def nearest_layer_value(target, axis_vals):
    return axis_vals[np.argmin(np.abs(axis_vals - target))]


def layer_mask(ref, axis_idx, layer_val, tol=LAYER_TOL_MM):
    return np.abs(ref[:, axis_idx] - layer_val) <= tol


def cluster_axis_values(coords, tol=LAYER_TOL_MM):
    """Return sorted representative grid-layer coordinates along one axis."""
    vals = np.sort(np.unique(np.round(coords, 1)))
    layers = []
    cur = [vals[0]]
    for v in vals[1:]:
        if v - cur[-1] <= tol:
            cur.append(v)
        else:
            layers.append(np.mean(cur))
            cur = [v]
    layers.append(np.mean(cur))
    return np.array(layers)


def build_geometry(ref):
    """Compute per-subject fixed geometry: axis grid layers, ROI sites,
    and slice point index sets for each arm/spacing config."""
    ap_layers = cluster_axis_values(ref[:, AP])
    si_layers = cluster_axis_values(ref[:, SI])
    lr_layers = cluster_axis_values(ref[:, LR])

    center = np.median(ref, axis=0)
    center_ap = nearest_layer_value(center[AP], ap_layers)
    center_si = nearest_layer_value(center[SI], si_layers)
    center_lr = nearest_layer_value(center[LR], lr_layers)

    # ROI sites: dome (max SI), inferior (min SI), lateral (extreme LR),
    # each near central position on the other two axes.
    def nearest_point(target):
        d = np.linalg.norm(ref - target[None, :], axis=1)
        return np.argmin(d)

    p90_si = np.percentile(ref[:, SI], 90)
    p10_si = np.percentile(ref[:, SI], 10)
    lr_span = ref[:, LR].max() - ref[:, LR].min()
    lateral_lr = center_lr + 0.4 * lr_span * np.sign(ref[:, LR].mean() - center_lr + 1e-9)
    if lateral_lr == center_lr:
        lateral_lr = center_lr + 0.4 * lr_span

    roi_centers = {
        "liver_dome": nearest_point(np.array([center_ap, p90_si, center_lr])),
        "liver_inferior": nearest_point(np.array([center_ap, p10_si, center_lr])),
        "liver_lateral": nearest_point(np.array([center_ap, center_si, lateral_lr])),
    }
    roi_masks = {}
    for site, pidx in roi_centers.items():
        d = np.linalg.norm(ref - ref[pidx][None, :], axis=1)
        roi_masks[site] = d <= ROI_RADIUS_MM

    def sag_mask(lr_val):
        return layer_mask(ref, LR, nearest_layer_value(lr_val, lr_layers))

    def cor_mask(ap_val):
        return layer_mask(ref, AP, nearest_layer_value(ap_val, ap_layers))

    geo = {
        "center_ap": center_ap, "center_si": center_si, "center_lr": center_lr,
        "roi_masks": roi_masks,
        "sag_mask": sag_mask, "cor_mask": cor_mask,
    }
    return geo


def arm_observation_indices(geo, arm, spacing=None):
    """Return list of (point_mask, plane_type) for the arm's slices.
    plane_type 'sag' -> observe (AP,SI) with sigma_sag; 'cor' -> observe
    (LR,SI) with alpha*sigma_sag."""
    if arm == "baseline-single-sagittal":
        return [(geo["sag_mask"](geo["center_lr"]), "sag")]
    if arm == "parallel-2sag":
        d = spacing
        return [(geo["sag_mask"](geo["center_lr"] - d / 2), "sag"),
                (geo["sag_mask"](geo["center_lr"] + d / 2), "sag")]
    if arm == "parallel-2cor":
        d = spacing
        return [(geo["cor_mask"](geo["center_ap"] - d / 2), "cor"),
                (geo["cor_mask"](geo["center_ap"] + d / 2), "cor")]
    if arm == "orthogonal-sag-cor":
        return [(geo["sag_mask"](geo["center_lr"]), "sag"),
                (geo["cor_mask"](geo["center_ap"]), "cor")]
    raise ValueError(arm)


def build_A_and_sigma(geo, arm, spacing, n_points, U_K, sigma_sag, alpha):
    """Gather flattened-index rows (point*3+axis) for observed in-plane
    components of each slice, return (idx_array, sigma_array, plane_pts)."""
    slices = arm_observation_indices(geo, arm, spacing)
    idx_list, sigma_list = [], []
    for mask, ptype in slices:
        pts = np.where(mask)[0]
        if ptype == "sag":
            comps = [AP, SI]
            sig = sigma_sag
        else:
            comps = [LR, SI]
            sig = alpha * sigma_sag
        for c in comps:
            idx_list.append(pts * 3 + c)
            sigma_list.append(np.full(pts.shape[0], sig))
    idx = np.concatenate(idx_list)
    sigma = np.concatenate(sigma_list)
    A = U_K[idx, :]
    return idx, sigma, A


def solve_and_reconstruct(A, mu_idx, y, reg_rel=RIDGE_REL):
    """Ridge-regularized least squares c_hat for batched frames.
    A: (m,K), mu_idx: (m,), y: (m,n_frames) noisy observations."""
    K = A.shape[1]
    AtA = A.T @ A
    reg = reg_rel * np.trace(AtA) / K if K > 0 else 0.0
    G = AtA + reg * np.eye(K)
    r = y - mu_idx[:, None]
    AtR = A.T @ r
    c_hat = np.linalg.solve(G, AtR)  # (K, n_frames)
    cond = np.linalg.cond(G)
    return c_hat, cond


def disp_to_flat(d):
    """(N,T,3) -> (3N,T) with axis order [p0_ap,p0_si,p0_lr,p1_ap,...]."""
    return d.transpose(1, 0, 2).reshape(d.shape[1], -1).T


def pca_fit(train_disp_list, K_max=K_MAX):
    D = np.concatenate([disp_to_flat(d) for d in train_disp_list], axis=1)  # (3N, n_train_frames)
    mu = D.mean(axis=1)
    Dc = D - mu[:, None]
    U, S, Vt = np.linalg.svd(Dc, full_matrices=False)
    K_eff = min(K_max, U.shape[1])
    return mu, U[:, :K_eff]


def rng_for(subj_idx, test_cyc, seed, arm_idx, spacing):
    sp = spacing if spacing else 0
    seed_val = (subj_idx * 1_000_003 + test_cyc * 10_007 + seed * 101 + arm_idx * 7 + sp) % (2**32)
    return np.random.default_rng(seed_val)


ARM_LIST = ["baseline-single-sagittal", "parallel-2sag", "parallel-2cor",
            "orthogonal-sag-cor", "oracle-full-dvf"]
ARM_IDX = {a: i for i, a in enumerate(ARM_LIST)}


def eval_arm(subj_idx, test_cyc, seed, arm, spacing, geo, mu, U_K, n_points,
             d_true_flat, sigma_sag, alpha, collect_full=False, roi_masks=None):
    """Returns dict of accumulated metrics for this (fold,seed,arm,spacing)."""
    n_frames = d_true_flat.shape[1]
    if arm == "oracle-full-dvf":
        idx = np.arange(3 * n_points)
        sigma = np.zeros_like(idx, dtype=np.float64)
        A = U_K
        y = d_true_flat.copy()
    else:
        idx, sigma, A = build_A_and_sigma(geo, arm, spacing, n_points, U_K, sigma_sag, alpha)
        rng = rng_for(subj_idx, test_cyc, seed, ARM_IDX[arm], spacing)
        noise = rng.normal(0.0, 1.0, size=(idx.shape[0], n_frames)) * sigma[:, None]
        y = d_true_flat[idx, :] + noise
    mu_idx = mu[idx]
    c_hat, cond = solve_and_reconstruct(A, mu_idx, y)
    d_hat_flat = mu[:, None] + U_K @ c_hat  # (3N, n_frames)

    err = (d_hat_flat - d_true_flat).reshape(n_points, 3, n_frames)
    err_mag = np.sqrt((err ** 2).sum(axis=1))  # (N, n_frames)

    out = {
        "sum_sq": float((err_mag ** 2).sum()),
        "count": int(err_mag.size),
        "sum_sq_ap": float((err[:, AP, :] ** 2).sum()),
        "sum_sq_si": float((err[:, SI, :] ** 2).sum()),
        "sum_sq_lr": float((err[:, LR, :] ** 2).sum()),
        "n_obs": int(idx.shape[0]),
        "log10_cond": float(np.log10(cond)) if cond > 0 and np.isfinite(cond) else float("nan"),
    }
    if collect_full:
        out["err_mag_flat"] = err_mag.ravel()
    if roi_masks is not None:
        roi_sq = {}
        for site, mask in roi_masks.items():
            hat_c = d_hat_flat.reshape(n_points, 3, n_frames)[mask].mean(axis=0)  # (3,n_frames)
            true_c = d_true_flat.reshape(n_points, 3, n_frames)[mask].mean(axis=0)
            e = np.linalg.norm(hat_c - true_c, axis=0)  # (n_frames,)
            roi_sq[site] = float((e ** 2).sum())
        out["roi_sum_sq"] = roi_sq
        out["roi_count"] = n_frames
    return out


def main():
    global T0
    T0 = time.time()
    axis_calibration_gate()

    log("loading ds5 subjects A and B ...")
    data = {"A": load_subject("A"), "B": load_subject("B")}
    for s in data:
        log(f"  subject {s}: n_points={data[s]['n_points']} cycles={data[s]['cycles']}")

    geo = {s: build_geometry(data[s]["ref"]) for s in data}
    for s in geo:
        for site, m in geo[s]["roi_masks"].items():
            log(f"  subject {s} ROI '{site}': {m.sum()} points within {ROI_RADIUS_MM}mm")

    folds = [("A", c) for c in data["A"]["cycles"]] + [("B", c) for c in data["B"]["cycles"]]
    log(f"total folds (leave-one-cycle-out): {len(folds)}")

    subj_idx_map = {"A": 0, "B": 1}

    # ---------- PASS 1: spacing selection (pooled sum_sq/count) ----------
    log("\n=== PASS 1: spacing selection sweep (default K/sigma/alpha) ===")
    pooled = {}  # (arm, spacing) -> [sum_sq, count]
    for subj, test_cyc in folds:
        d = data[subj]
        train_disp = [d["disp"][c] for c in d["cycles"] if c != test_cyc]
        mu, U_K = pca_fit(train_disp, K_max=K_DEFAULT)
        d_true_flat = disp_to_flat(d["disp"][test_cyc])
        for seed in SEEDS:
            for arm in ["parallel-2sag", "parallel-2cor"]:
                for sp in SPACINGS:
                    r = eval_arm(subj_idx_map[subj], test_cyc, seed, arm, sp, geo[subj],
                                 mu, U_K, d["n_points"], d_true_flat,
                                 SIGMA_SAG_DEFAULT, ALPHA_DEFAULT)
                    key = (arm, sp)
                    if key not in pooled:
                        pooled[key] = [0.0, 0]
                    pooled[key][0] += r["sum_sq"]
                    pooled[key][1] += r["count"]
        log(f"  fold {subj}/cyc{test_cyc} done")

    best_spacing = {}
    for arm in ["parallel-2sag", "parallel-2cor"]:
        cands = {sp: np.sqrt(pooled[(arm, sp)][0] / pooled[(arm, sp)][1]) for sp in SPACINGS}
        best = min(cands, key=cands.get)
        best_spacing[arm] = best
        log(f"  {arm} spacing sweep pooled RMSE(mm): " +
            ", ".join(f"d={sp}:{v:.3f}" for sp, v in cands.items()) + f"  -> best d={best}")

    # ---------- PASS 2: canonical arms, full detail, per-seed pooling ----------
    log("\n=== PASS 2: canonical 5 arms, full metrics, 3 seeds x 9 folds ===")
    canonical_spacing = {"baseline-single-sagittal": None,
                          "parallel-2sag": best_spacing["parallel-2sag"],
                          "parallel-2cor": best_spacing["parallel-2cor"],
                          "orthogonal-sag-cor": None,
                          "oracle-full-dvf": None}

    # accumulators: per (arm, seed) pooled sums; also raw per-fold-seed pairs for prediction test
    per_seed_acc = {arm: {seed: {"sum_sq": 0.0, "count": 0, "sum_sq_ap": 0.0, "sum_sq_si": 0.0,
                                  "sum_sq_lr": 0.0, "n_obs": [], "log10_cond": [],
                                  "err_mag": [], "roi_sum_sq": {s: 0.0 for s in ["liver_dome", "liver_inferior", "liver_lateral"]},
                                  "roi_count": 0}
                          for seed in SEEDS} for arm in ARM_LIST}
    paired_27 = []  # list of dicts per (fold,seed): rmse per arm, needed for prediction test

    for subj, test_cyc in folds:
        d = data[subj]
        train_disp = [d["disp"][c] for c in d["cycles"] if c != test_cyc]
        mu, U_K = pca_fit(train_disp, K_max=K_DEFAULT)
        d_true_flat = disp_to_flat(d["disp"][test_cyc])
        for seed in SEEDS:
            fold_result = {}
            for arm in ARM_LIST:
                sp = canonical_spacing[arm]
                r = eval_arm(subj_idx_map[subj], test_cyc, seed, arm, sp, geo[subj],
                             mu, U_K, d["n_points"], d_true_flat,
                             SIGMA_SAG_DEFAULT, ALPHA_DEFAULT,
                             collect_full=True, roi_masks=geo[subj]["roi_masks"])
                acc = per_seed_acc[arm][seed]
                acc["sum_sq"] += r["sum_sq"]; acc["count"] += r["count"]
                acc["sum_sq_ap"] += r["sum_sq_ap"]; acc["sum_sq_si"] += r["sum_sq_si"]; acc["sum_sq_lr"] += r["sum_sq_lr"]
                acc["n_obs"].append(r["n_obs"]); acc["log10_cond"].append(r["log10_cond"])
                acc["err_mag"].append(r["err_mag_flat"])
                for site in acc["roi_sum_sq"]:
                    acc["roi_sum_sq"][site] += r["roi_sum_sq"][site]
                acc["roi_count"] += r["roi_count"]
                fold_result[arm] = {
                    "dvf_rmse3d_mm": float(np.sqrt(r["sum_sq"] / r["count"])),
                    "rmse_lr_mm": float(np.sqrt(r["sum_sq_lr"] / r["count"])),
                    "rmse_si_mm": float(np.sqrt(r["sum_sq_si"] / r["count"])),
                    "rmse_ap_mm": float(np.sqrt(r["sum_sq_ap"] / r["count"])),
                }
            paired_27.append({"subj": subj, "test_cyc": test_cyc, "seed": seed, **fold_result})
        log(f"  fold {subj}/cyc{test_cyc} full-metrics done")

    geo_summary = {
        s: {
            "center_ap": geo[s]["center_ap"], "center_si": geo[s]["center_si"],
            "center_lr": geo[s]["center_lr"],
            "roi_n_points": {site: int(m.sum()) for site, m in geo[s]["roi_masks"].items()},
            "n_points": data[s]["n_points"],
        } for s in geo
    }
    spacing_sweep_rmse = {
        f"{arm}|d={sp}": float(np.sqrt(pooled[(arm, sp)][0] / pooled[(arm, sp)][1]))
        for (arm, sp) in pooled
    }
    return {
        "geo_summary": geo_summary, "folds": folds, "best_spacing": best_spacing,
        "per_seed_acc": per_seed_acc, "paired_27": paired_27,
        "canonical_spacing": canonical_spacing, "spacing_sweep_rmse": spacing_sweep_rmse,
    }


if __name__ == "__main__":
    out = main()
    import pickle
    with open("run_raw.pkl", "wb") as f:
        pickle.dump(out, f)
    log("\nsaved run_raw.pkl, done with core computation.")
