1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
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()
|