#!/usr/bin/env python3 """Plot dense long-training phase-transition trajectories.""" from __future__ import annotations import argparse from pathlib import Path import matplotlib.pyplot as plt import pandas as pd def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument( "--run-dir", type=Path, default=Path("outputs/phase_transition_dense_T30000_352traj"), ) return parser.parse_args() def main() -> None: args = parse_args() run_dir = args.run_dir runs = pd.read_csv(run_dir / "runs.csv") summary = pd.read_csv(run_dir / "width_summary.csv").sort_values( "fa_capacity_margin", ascending=False ) fig, ax = plt.subplots(figsize=(8.8, 5.4), dpi=180) fa = runs[runs.run_type == "fa"].copy() fa["_tid"] = fa.init_seed.astype(str) + ":" + fa.feedback_seed.astype(str) for _tid, group in fa.groupby("_tid"): group = group.sort_values("fa_capacity_margin", ascending=False) ax.plot( group.fa_capacity_margin, group.train_gap_to_bp.clip(lower=1e-7), color="#174f91", alpha=0.12, linewidth=0.8, ) ax.errorbar( summary.fa_capacity_margin, summary.train_gap_mean.clip(lower=1e-7), yerr=summary.train_gap_std, color="#0b4d92", marker="o", linewidth=2.2, capsize=3, label="mean +/- sd", ) ax.axvline(0, color="black", linestyle="--", linewidth=1.1) ax.set_yscale("log") ax.set_xlim(150, -210) ax.set_xlabel("hard FA capacity margin P - K_FA - N*out") ax.set_ylabel("FA train MSE - BP train MSE (log scale)") ax.set_title("Dense long-training FA/BP gap vs capacity margin") ax.grid(alpha=0.18, which="both") ax.legend(frameon=True) fig.tight_layout() path = run_dir / "dense_T30000_gap_logscale.png" fig.savefig(path) plt.close(fig) print(f"plot: {path}") if __name__ == "__main__": main()