#!/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()