summaryrefslogtreecommitdiff
path: root/scripts/plot_phase_transition.py
diff options
context:
space:
mode:
Diffstat (limited to 'scripts/plot_phase_transition.py')
-rw-r--r--scripts/plot_phase_transition.py245
1 files changed, 245 insertions, 0 deletions
diff --git a/scripts/plot_phase_transition.py b/scripts/plot_phase_transition.py
new file mode 100644
index 0000000..025e267
--- /dev/null
+++ b/scripts/plot_phase_transition.py
@@ -0,0 +1,245 @@
+#!/usr/bin/env python3
+"""Plot the capacity-exhaustion phase transition for BP/FA."""
+
+from __future__ import annotations
+
+import argparse
+from pathlib import Path
+
+import matplotlib.pyplot as plt
+import numpy as np
+import pandas as pd
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument(
+ "--run-dir",
+ type=Path,
+ default=Path("outputs/phase_transition_sgd_256_lr001_T3000"),
+ )
+ parser.add_argument(
+ "--outdir",
+ type=Path,
+ default=Path("outputs/phase_transition_sgd_256_lr001_T3000"),
+ )
+ return parser.parse_args()
+
+
+def summarize_runs(runs: pd.DataFrame, summary: pd.DataFrame) -> pd.DataFrame:
+ fa = runs[runs["run_type"] == "fa"].copy()
+ records = []
+ for margin, group in fa.groupby("fa_capacity_margin"):
+ gaps = group["train_gap_to_bp"].to_numpy(dtype=np.float64)
+ records.append(
+ {
+ "fa_capacity_margin": float(margin),
+ "gap_mean": float(gaps.mean()),
+ "gap_q05": float(np.quantile(gaps, 0.05)),
+ "gap_q25": float(np.quantile(gaps, 0.25)),
+ "gap_q75": float(np.quantile(gaps, 0.75)),
+ "gap_q95": float(np.quantile(gaps, 0.95)),
+ }
+ )
+ gap_summary = pd.DataFrame(records)
+ merged = summary.merge(gap_summary, on="fa_capacity_margin", how="left")
+ return merged.sort_values("fa_capacity_margin", ascending=False)
+
+
+def trajectory_id(frame: pd.DataFrame) -> pd.Series:
+ return frame["init_seed"].astype(str) + ":" + frame["feedback_seed"].astype(str)
+
+
+def add_regime_shading(ax: plt.Axes, summary: pd.DataFrame) -> None:
+ max_margin = float(summary["fa_capacity_margin"].max())
+ min_margin = float(summary["fa_capacity_margin"].min())
+ bp_bad_margin = float(
+ summary.loc[summary["bp_capacity_margin"] < 0, "fa_capacity_margin"].max()
+ )
+ ax.axvspan(max_margin + 80, 0, color="#e8f1fb", alpha=0.55, linewidth=0)
+ ax.axvspan(0, bp_bad_margin, color="#fff1d8", alpha=0.55, linewidth=0)
+ ax.axvspan(bp_bad_margin, min_margin - 80, color="#fde2e2", alpha=0.55, linewidth=0)
+ ax.axvline(0, color="black", linestyle="--", linewidth=1.1)
+ ax.axvline(bp_bad_margin, color="#8b1a1a", linestyle=":", linewidth=1.2)
+
+
+def plot_transition(run_dir: Path, outdir: Path) -> list[Path]:
+ runs = pd.read_csv(run_dir / "runs.csv")
+ summary = pd.read_csv(run_dir / "width_summary.csv")
+ merged = summarize_runs(runs, summary)
+ fa_runs = runs[runs["run_type"] == "fa"].copy()
+ fa_runs["_trajectory_id"] = trajectory_id(fa_runs)
+
+ outdir.mkdir(parents=True, exist_ok=True)
+ paths: list[Path] = []
+
+ fig, (ax_loss, ax_gap) = plt.subplots(
+ 2,
+ 1,
+ figsize=(11.2, 8.4),
+ dpi=180,
+ sharex=True,
+ height_ratios=[1.0, 1.12],
+ )
+
+ add_regime_shading(ax_loss, merged)
+ add_regime_shading(ax_gap, merged)
+
+ x = merged["fa_capacity_margin"].to_numpy(dtype=np.float64)
+ ax_loss.plot(
+ x,
+ merged["bp_train_mean"],
+ color="#174f91",
+ marker="o",
+ linewidth=2.1,
+ label="BP train loss",
+ )
+ ax_loss.plot(
+ x,
+ merged["fa_train_mean"],
+ color="#a84708",
+ marker="o",
+ linewidth=2.1,
+ label="FA train loss",
+ )
+ ax_loss.fill_between(
+ x,
+ np.maximum(merged["bp_train_mean"] - merged["bp_train_std"], 1e-8),
+ merged["bp_train_mean"] + merged["bp_train_std"],
+ color="#174f91",
+ alpha=0.13,
+ linewidth=0,
+ )
+ ax_loss.fill_between(
+ x,
+ np.maximum(merged["fa_train_mean"] - merged["fa_train_std"], 1e-8),
+ merged["fa_train_mean"] + merged["fa_train_std"],
+ color="#a84708",
+ alpha=0.13,
+ linewidth=0,
+ )
+ ax_loss.set_yscale("log")
+ ax_loss.set_ylabel("train MSE")
+ ax_loss.set_title("Capacity exhaustion opens the FA/BP train-gap")
+ ax_loss.legend(loc="upper left", frameon=True)
+ ax_loss.grid(alpha=0.18)
+
+ for _tid, group in fa_runs.groupby("_trajectory_id"):
+ group = group.sort_values("fa_capacity_margin", ascending=False)
+ ax_gap.plot(
+ group["fa_capacity_margin"],
+ group["train_gap_to_bp"],
+ color="#c65f16",
+ alpha=0.13,
+ linewidth=0.95,
+ )
+ ax_gap.fill_between(
+ x,
+ merged["gap_q05"],
+ merged["gap_q95"],
+ color="#c65f16",
+ alpha=0.13,
+ linewidth=0,
+ label="FA trajectories 5-95%",
+ )
+ ax_gap.fill_between(
+ x,
+ merged["gap_q25"],
+ merged["gap_q75"],
+ color="#c65f16",
+ alpha=0.24,
+ linewidth=0,
+ label="FA trajectories 25-75%",
+ )
+ ax_gap.plot(
+ x,
+ merged["train_gap_mean"],
+ color="#a84708",
+ marker="o",
+ linewidth=2.2,
+ label="mean gap",
+ )
+ ax_gap.axhline(0, color="black", linewidth=1.0)
+ ax_gap.set_ylabel("FA train MSE - BP train MSE")
+ ax_gap.set_xlabel("hard FA capacity margin P - K_FA - N*out (capacity decreases ->)")
+ ax_gap.grid(alpha=0.18)
+ ax_gap.legend(loc="upper left", frameon=True)
+
+ max_margin = float(merged["fa_capacity_margin"].max())
+ min_margin = float(merged["fa_capacity_margin"].min())
+ ax_gap.set_xlim(max_margin + 80, min_margin - 80)
+
+ ax_loss.text(
+ 0.18,
+ 0.92,
+ "redundant",
+ transform=ax_loss.transAxes,
+ ha="center",
+ va="center",
+ fontsize=10,
+ color="#174f91",
+ )
+ ax_loss.text(
+ 0.57,
+ 0.92,
+ "FA-deficient\nBP-capable",
+ transform=ax_loss.transAxes,
+ ha="center",
+ va="center",
+ fontsize=10,
+ color="#9a5b00",
+ )
+ ax_loss.text(
+ 0.90,
+ 0.92,
+ "both-deficient",
+ transform=ax_loss.transAxes,
+ ha="center",
+ va="center",
+ fontsize=10,
+ color="#8b1a1a",
+ )
+
+ fig.tight_layout()
+ path = outdir / "phase_transition_capacity_exhaustion.png"
+ fig.savefig(path)
+ plt.close(fig)
+ paths.append(path)
+
+ detail_fig, ax = plt.subplots(figsize=(8.4, 5.4), dpi=180)
+ add_regime_shading(ax, merged)
+ ax.errorbar(
+ x,
+ merged["train_gap_mean"],
+ yerr=merged["train_gap_std"],
+ color="#a84708",
+ marker="o",
+ linewidth=2.1,
+ capsize=3,
+ label="mean +/- sd",
+ )
+ ax.axhline(0, color="black", linewidth=1.0)
+ ax.set_xlim(max_margin + 80, min_margin - 80)
+ ax.set_xlabel("hard FA capacity margin P - K_FA - N*out (capacity decreases ->)")
+ ax.set_ylabel("FA train MSE - BP train MSE")
+ ax.set_title("FA/BP gap increases after redundant capacity is exhausted")
+ ax.legend(loc="upper left", frameon=True)
+ ax.grid(alpha=0.18)
+ detail_fig.tight_layout()
+ detail_path = outdir / "phase_transition_gap_only.png"
+ detail_fig.savefig(detail_path)
+ plt.close(detail_fig)
+ paths.append(detail_path)
+
+ return paths
+
+
+def main() -> None:
+ args = parse_args()
+ paths = plot_transition(args.run_dir, args.outdir)
+ for path in paths:
+ print(f"plot: {path}")
+
+
+if __name__ == "__main__":
+ main()