summaryrefslogtreecommitdiff
path: root/scripts
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-06-05 13:19:08 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-06-05 13:19:08 -0500
commitd97cbe275f25602986e024bdf243c13363d2003f (patch)
treedd4b8c9ba3fea8f55bda08fa3720d1f25895d1be /scripts
parent6d8be33a7fb9547c3f035bec0476eabcec87ecdd (diff)
Record dense soft capacity transition
Diffstat (limited to 'scripts')
-rw-r--r--scripts/plot_dense_phase_transition.py69
1 files changed, 69 insertions, 0 deletions
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()