#!/usr/bin/env python3
"""Finite-k sample-size scaling panel for the hierarchical three-layer model.

This panel tests finite-size smoothing and sequential recovery on nested
datasets.  Since k is fixed, it does not estimate the growing-rank scaling-law
exponent of arXiv:2605.14567.
"""

from __future__ import annotations

import argparse
import csv
import json
import math
import os
import sys
from dataclasses import asdict, dataclass, replace
from pathlib import Path
from typing import Any

import numpy as np
import torch

os.environ.setdefault("MPLCONFIGDIR", "/tmp/matplotlib-codex-hierarchical3-scaling")
import matplotlib

matplotlib.use("Agg")
import matplotlib.pyplot as plt


WORKSPACE_ROOT = Path(__file__).resolve().parents[3]
if str(WORKSPACE_ROOT) not in sys.path:
    sys.path.insert(0, str(WORKSPACE_ROOT))

from experiments.brainstorm.resnetl3.hierarchical3_muon_bbp import (
    Config as TrainingConfig,
    OPTIMIZERS,
    clone_state,
    make_setup,
    new_model,
    run_trajectory,
    teacher_on,
)


ROOT = Path(__file__).resolve().parent
DEFAULT_OUT = ROOT / "results/hierarchical3_sample_scaling_latest"


@dataclass
class Config:
    outdir: str = str(DEFAULT_OUT)
    seed: int = 271828
    sample_sizes: str = "64,128,256,512,1024,2048"
    repeats: int = 5
    steps: int = 800
    gd_lr: float = 0.02
    muon_lr: float = 0.004
    vector_lr: float = 0.02


def parse_sizes(text: str) -> list[int]:
    values = sorted({int(value.strip()) for value in text.split(",") if value.strip()})
    if not values or values[0] <= 0:
        raise ValueError("sample_sizes must contain positive integers")
    return values


def write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
    fields: list[str] = []
    for row in rows:
        for key in row:
            if key not in fields:
                fields.append(key)
    with path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(handle, fieldnames=fields, lineterminator="\n")
        writer.writeheader()
        writer.writerows(rows)


def run_panel(cfg: Config) -> tuple[list[dict[str, Any]], dict[str, Any]]:
    sizes = parse_sizes(cfg.sample_sizes)
    training_cfg = TrainingConfig(
        seed=cfg.seed,
        steps=cfg.steps,
        diag_every=cfg.steps,
        repeats=1,
        train_n=max(sizes),
        gd_lr=cfg.gd_lr,
        muon_lr=cfg.muon_lr,
        vector_lr=cfg.vector_lr,
        bbp_checkpoints=2,
    )
    setup = make_setup(training_cfg)
    rows: list[dict[str, Any]] = []
    for repeat in range(cfg.repeats):
        model_seed = cfg.seed + 1000 * repeat
        base = new_model(training_cfg, setup, model_seed)
        init_state = clone_state(base)
        generator = torch.Generator().manual_seed(cfg.seed + 20_000 + repeat)
        z_max = torch.randn(max(sizes), training_cfg.k, generator=generator, dtype=setup.z.dtype)
        y_max = teacher_on(z_max, setup)
        for n in sizes:
            nested_z = z_max[:n]
            nested_y = y_max[:n]
            cfg_n = replace(training_cfg, train_n=n)
            for spec in OPTIMIZERS:
                trajectory, _ = run_trajectory(
                    cfg_n,
                    setup,
                    spec,
                    "empirical",
                    repeat,
                    init_state,
                    nested_z,
                    nested_y,
                )
                initial, final = trajectory[0], trajectory[-1]
                row: dict[str, Any] = {
                    "sample_size": n,
                    "repeat": repeat,
                    "optimizer": spec.name,
                    "initial_test_risk": initial["test_population_risk"],
                    "final_test_risk": final["test_population_risk"],
                    "risk_ratio": final["test_population_risk"] / initial["test_population_risk"],
                    "backtracking_rejects": sum(r["backtracking_rejects"] for r in trajectory),
                }
                for degree in (2, 3, 4):
                    key = f"E_order_{degree}"
                    row[f"order_{degree}_error_ratio"] = final[key] / max(initial[key], 1e-300)
                for coord in range(1, training_cfg.k + 1):
                    row[f"V_coord_{coord}_recovered"] = int(final[f"align_V_{coord}"] >= 0.8)
                    row[f"U_coord_{coord}_recovered"] = int(final[f"align_U_{coord}"] >= 0.8)
                rows.append(row)
            print(f"repeat {repeat + 1}/{cfg.repeats}, n={n}", flush=True)

    summary: dict[str, Any] = {"config": asdict(cfg), "finite_k_warning": True, "optimizers": {}}
    for spec in OPTIMIZERS:
        by_size: list[dict[str, Any]] = []
        for n in sizes:
            selected = [row for row in rows if row["optimizer"] == spec.name and row["sample_size"] == n]
            risks = np.asarray([row["final_test_risk"] for row in selected], dtype=float)
            entry: dict[str, Any] = {
                "sample_size": n,
                "risk_mean": float(np.mean(risks)),
                "risk_median": float(np.median(risks)),
                "risk_std": float(np.std(risks, ddof=1)) if len(risks) > 1 else 0.0,
                "risk_geometric_mean": float(math.exp(np.mean(np.log(np.maximum(risks, 1e-300))))),
            }
            for coord in range(1, training_cfg.k + 1):
                entry[f"V_coord_{coord}_recovery_probability"] = float(
                    np.mean([row[f"V_coord_{coord}_recovered"] for row in selected])
                )
            by_size.append(entry)
        tail = by_size[-min(4, len(by_size)) :]
        log_n = np.log([entry["sample_size"] for entry in tail])
        log_risk = np.log([max(entry["risk_median"], 1e-300) for entry in tail])
        slope = float(np.polyfit(log_n, log_risk, 1)[0]) if len(tail) >= 2 else float("nan")
        summary["optimizers"][spec.name] = {
            "by_size": by_size,
            "descriptive_tail_loglog_slope": slope,
            "slope_status": "finite-k descriptive fit; not the arXiv:2605.14567 exponent",
        }
    return rows, summary


def plot_panel(summary: dict[str, Any], outdir: Path) -> None:
    fig, axes = plt.subplots(1, 2, figsize=(12.2, 4.8), dpi=190)
    colors = dict(zip((spec.name for spec in OPTIMIZERS), plt.cm.viridis(np.linspace(0.05, 0.92, len(OPTIMIZERS)))))
    for spec in OPTIMIZERS:
        entries = summary["optimizers"][spec.name]["by_size"]
        sizes = np.asarray([entry["sample_size"] for entry in entries])
        median = np.asarray([entry["risk_median"] for entry in entries])
        axes[0].plot(sizes, median, "o-", color=colors[spec.name], label=spec.name)
        for coord, style in zip((1, 2, 3), ("-", "--", ":")):
            axes[1].plot(
                sizes,
                [entry[f"V_coord_{coord}_recovery_probability"] for entry in entries],
                color=colors[spec.name],
                ls=style,
                marker="o",
                ms=3,
                label=f"{spec.name}, coord {coord}",
            )
    axes[0].set_xscale("log", base=2)
    axes[0].set_yscale("log")
    axes[0].set_title("Finite-k exact test risk versus sample size")
    axes[0].set_ylabel("median population risk")
    axes[1].set_xscale("log", base=2)
    axes[1].set_ylim(-0.03, 1.03)
    axes[1].set_title("Direction-wise recovery probability")
    axes[1].set_ylabel("fraction with alignment >= 0.8")
    for ax in axes:
        ax.set_xlabel("training samples (nested datasets)")
        ax.grid(alpha=0.20)
        ax.legend(fontsize=6.5, ncol=2)
    fig.tight_layout()
    fig.savefig(outdir / "sample_scaling_and_feature_recovery.png", bbox_inches="tight")
    plt.close(fig)


def write_report(summary: dict[str, Any], outdir: Path) -> None:
    lines = [
        "# Finite-k sample-size scaling panel\n\n",
        "This is a nested-dataset empirical panel for the repository's fixed `k=3` ResNet. ",
        "It tests finite-size smoothing and direction-wise recovery, but it cannot estimate the growing-rank exponent of arXiv:2605.14567.\n\n",
        "| optimizer | descriptive tail slope | final-n median risk |\n",
        "|---|---:|---:|\n",
    ]
    for spec in OPTIMIZERS:
        item = summary["optimizers"][spec.name]
        final = item["by_size"][-1]
        lines.append(
            f"| {spec.name} | {item['descriptive_tail_loglog_slope']:.4g} | {final['risk_median']:.5g} |\n"
        )
    lines.extend(
        [
            "\nThe fitted slope uses only the four largest sample sizes and is descriptive. ",
            "A theorem-facing scaling panel must let the latent rank grow with ambient dimension and reproduce the lifted `D ~ d^q` construction.\n",
        ]
    )
    (outdir / "REPORT.md").write_text("".join(lines), encoding="utf-8")


def parse_args() -> Config:
    parser = argparse.ArgumentParser(description=__doc__)
    defaults = Config()
    for name, value in asdict(defaults).items():
        parser.add_argument("--" + name.replace("_", "-"), type=type(value), default=value)
    return Config(**vars(parser.parse_args()))


def main() -> None:
    cfg = parse_args()
    outdir = Path(cfg.outdir)
    outdir.mkdir(parents=True, exist_ok=True)
    rows, summary = run_panel(cfg)
    write_csv(outdir / "sample_scaling_runs.csv", rows)
    (outdir / "summary.json").write_text(json.dumps(summary, indent=2), encoding="utf-8")
    plot_panel(summary, outdir)
    write_report(summary, outdir)
    print(f"Wrote {outdir}", flush=True)


if __name__ == "__main__":
    main()
