diff options
Diffstat (limited to 'scripts/plot_phase_transition.py')
| -rw-r--r-- | scripts/plot_phase_transition.py | 245 |
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() |
