summaryrefslogtreecommitdiff
path: root/scripts/plot_dense_phase_transition.py
blob: 4b6b4b6bba32f5b7fd93acc5f8a7572ea61e349c (plain)
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()