From d97cbe275f25602986e024bdf243c13363d2003f Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Fri, 5 Jun 2026 13:19:08 -0500 Subject: Record dense soft capacity transition --- scripts/plot_dense_phase_transition.py | 69 ++++++++++++++++++++++++++++++++++ 1 file changed, 69 insertions(+) create mode 100644 scripts/plot_dense_phase_transition.py (limited to 'scripts/plot_dense_phase_transition.py') diff --git a/scripts/plot_dense_phase_transition.py b/scripts/plot_dense_phase_transition.py new file mode 100644 index 0000000..4b6b4b6 --- /dev/null +++ b/scripts/plot_dense_phase_transition.py @@ -0,0 +1,69 @@ +#!/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() -- cgit v1.2.3