summaryrefslogtreecommitdiff
path: root/scripts/closed_form_soft_ramp.py
blob: ac71988d676b8f112d150bc84218600a9891770e (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
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
#!/usr/bin/env python3
"""Closed-form soft-ramp gap law: derive and validate against the dense sweep.

Lazy (frozen) tangent operators, squared loss L = ||r||^2/(2N). BP residual
decays as r_t = (I - eta K_BP/N)^t r_0 with K_BP = J J^T. FA uses the symmetric
part S_FA = (K_FA + K_FA^T)/2 of the non-symmetric tangent operator
K_FA = J J_tilde^T. The actual-FA initial-moment theorem gives the residual-
direction scalar burden

    rho = output_share = r^T K_out r / r^T K_BP r = 1 - E_B[e_0]  in (0,1),

and the lazy scalar model S_FA ~= rho K_BP yields the closed form

    gap_T = (1/2N) sum_i c_i^2 [ (1 - eta rho lam_i/N)^{2T}
                                - (1 - eta lam_i/N)^{2T} ],

with K_BP v_i = lam_i v_i and c_i = v_i^T r_0. Each bracket is >= 0 and analytic
in (rho, lam_i, T): the gap is a smooth functional of the spectrum, so it is a
soft ramp, with a hard threshold only in the degenerate rho -> 0 limit.

In the BP-converged regime the gap is dominated by the slowest surviving FA mode,
    gap_T ~= (c_min^2/2N) exp(-2 eta rho lam_min T/N),
so log gap_T ~= const - (2 eta T/N) rho lam_min: the ramp is driven by lam_min(w)
growing smoothly with capacity.

This script computes the no-fit closed-form gap at initialization for each width
in the dense T=30000 sweep and compares it to the empirical ramp. The dense
run's config and measured gaps are loaded from its output directory, and the
closed form is evaluated on the *same data and the same init seed* the sweep
actually used (matched comparison), plus extra init seeds on the same data to
show the init-to-init variability of the no-fit prediction (band).
"""

from __future__ import annotations

import argparse
import csv
import json
import sys
from pathlib import Path

import matplotlib.pyplot as plt
import numpy as np
import torch

SCRIPT_DIR = Path(__file__).resolve().parent
if str(SCRIPT_DIR) not in sys.path:
    sys.path.insert(0, str(SCRIPT_DIR))

import downstream_capacity_sweep as dcs  # noqa: E402


def parse_args() -> argparse.Namespace:
    p = argparse.ArgumentParser(description="Closed-form soft-ramp gap law.")
    p.add_argument(
        "--dense-dir",
        type=Path,
        default=Path("outputs/phase_transition_dense_T30000_352traj"),
        help="Dense sweep output dir (summary.json + width_summary.csv).",
    )
    p.add_argument("--extra-init-seeds", type=int, default=4,
                   help="Extra init seeds (same data) for the variability band.")
    p.add_argument("--torch-threads", type=int, default=8)
    p.add_argument("--outdir", type=Path, default=Path("outputs/closed_form_soft_ramp"))
    return p.parse_args()


def load_dense_run(dense_dir: Path) -> tuple[dict, list[dict]]:
    """Load the dense sweep's config and per-width measured gaps."""
    summary = json.loads((dense_dir / "summary.json").read_text())
    dense_config = summary["config"]
    rows = []
    with (dense_dir / "width_summary.csv").open() as fh:
        for row in csv.DictReader(fh):
            rows.append({
                "width": int(row["width"]),
                "margin": int(row["fa_capacity_margin"]),
                "measured_gap": float(row["train_gap_mean"]),
                "measured_gap_std": float(row["train_gap_std"]),
                "runs_fa": int(row["runs_fa"]),
            })
    rows.sort(key=lambda r: r["width"])
    return dense_config, rows


def make_config(dense_config: dict, torch_threads: int, outdir: Path) -> dcs.RunConfig:
    # Mirror the dense run's data-generating config exactly so make_data
    # reproduces the very same x_train/y_train the sweep trained on.
    return dcs.RunConfig(
        task=dense_config["task"],
        input_dim=dense_config["input_dim"],
        output_dim=dense_config["output_dim"],
        teacher_rank=dense_config["teacher_rank"],
        teacher_width=dense_config["teacher_width"],
        teacher_hidden_layers=dense_config["teacher_hidden_layers"],
        normalize_targets=dense_config["normalize_targets"],
        widths=dense_config["widths"],
        train_samples=dense_config["train_samples"],
        test_samples=dense_config["test_samples"],
        probe_samples=dense_config["probe_samples"],
        steps=dense_config["steps"],
        lr=dense_config["lr"],
        optimizer=dense_config.get("optimizer", "sgd"),
        init_seeds=dense_config["init_seeds"],
        feedback_seeds=dense_config["feedback_seeds"],
        init_seed_offset=dense_config["init_seed_offset"],
        feedback_seed_offset=dense_config["feedback_seed_offset"],
        data_seed=dense_config["data_seed"],
        noise_std=dense_config["noise_std"],
        feedback_scale=dense_config["feedback_scale"],
        capacity_q=dense_config["capacity_q"],
        jacobian_lambda_rel=dense_config["jacobian_lambda_rel"],
        skip_jacobian=True,
        device="cpu",
        torch_threads=torch_threads,
        outdir=str(outdir),
        plot=False,
    )


def forward(weights: list[torch.Tensor], x: torch.Tensor) -> torch.Tensor:
    a = x
    for layer, w in enumerate(weights):
        z = a @ w.T
        a = torch.relu(z) if layer < len(weights) - 1 else z
    return a


def jacobian(weights: list[torch.Tensor], x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, int]:
    leaves = [w.clone().requires_grad_(True) for w in weights]
    pred = forward(leaves, x).reshape(-1)
    n_params = sum(w.numel() for w in leaves)
    jac = torch.zeros(pred.numel(), n_params, dtype=torch.float64)
    for i in range(pred.numel()):
        grads = torch.autograd.grad(pred[i], leaves, retain_graph=True)
        jac[i] = torch.cat([g.reshape(-1) for g in grads])
    return jac, pred.detach(), leaves[-1].numel()


def closed_form_gap(config, x, y, width: int, init_seed: int) -> dict:
    weights = dcs.initialize_weights(config, width, init_seed)
    jac, pred, n_out = jacobian(weights, x)
    r0 = pred - y.reshape(-1)
    k_bp = jac @ jac.T
    jac_out = jac[:, jac.shape[1] - n_out:]
    speed_bp = float((jac.T @ r0).pow(2).sum())
    rho = float((jac_out.T @ r0).pow(2).sum() / speed_bp)
    lam, vecs = torch.linalg.eigh(k_bp)
    coeff = vecs.T @ r0
    eta, n, t = config.lr, config.train_samples, config.steps
    a_bp = 1.0 - eta * lam / n
    a_fa = 1.0 - eta * rho * lam / n
    # only modes the discrete recursion does not blow up (|a|<=1); the real SGD
    # run was stable, so unstable modes carry negligible residual mass.
    bp = torch.where(a_bp.abs() <= 1, a_bp.pow(2 * t), torch.zeros_like(lam))
    fa = torch.where(a_fa.abs() <= 1, a_fa.pow(2 * t), torch.zeros_like(lam))
    gap = float((coeff.pow(2) * (fa - bp)).sum() / (2 * n))
    pos = lam[lam > 1e-9]
    return {"rho": rho, "lam_min": float(pos.min()), "lam_max": float(lam.max()),
            "pred_gap": gap}


def main() -> None:
    args = parse_args()
    torch.set_num_threads(args.torch_threads)
    dense_config, dense_rows = load_dense_run(args.dense_dir)
    config = make_config(dense_config, args.torch_threads, args.outdir)
    args.outdir.mkdir(parents=True, exist_ok=True)

    # The sweep's init seed is 10_000 + init_seed_offset + init_index
    # (downstream_capacity_sweep.main); the dense run used init_seeds=1.
    matched_seed = 10_000 + dense_config["init_seed_offset"]
    band_seeds = [matched_seed + k for k in range(1, args.extra_init_seeds + 1)]
    x, y, *_ = dcs.make_data(config)

    rows, seed_rows = [], []
    print(f"matched init seed = {matched_seed}, band seeds = {band_seeds}")
    print(f"{'w':>3}{'margin':>7}{'rho':>7}{'lam_min':>9}{'pred_gap':>11}{'measured':>11}")
    for dense_row in dense_rows:
        width, margin = dense_row["width"], dense_row["margin"]
        meas, meas_std = dense_row["measured_gap"], dense_row["measured_gap_std"]
        d = closed_form_gap(config, x, y, width, matched_seed)
        band = []
        for seed in band_seeds:
            extra = closed_form_gap(config, x, y, width, seed)
            band.append(extra["pred_gap"])
            seed_rows.append({"width": width, "margin": margin, "init_seed": seed,
                              "rho": extra["rho"], "lam_min": extra["lam_min"],
                              "pred_gap": extra["pred_gap"]})
        seed_rows.append({"width": width, "margin": margin, "init_seed": matched_seed,
                          "rho": d["rho"], "lam_min": d["lam_min"],
                          "pred_gap": d["pred_gap"]})
        all_preds = band + [d["pred_gap"]]
        width_seed_rows = [r for r in seed_rows if r["width"] == width]
        rows.append({"width": width, "margin": margin, "rho": d["rho"],
                     "lam_min": d["lam_min"], "lam_max": d["lam_max"],
                     "pred_gap": d["pred_gap"], "measured_gap": meas,
                     "measured_gap_std": meas_std,
                     "pred_gap_band_lo": min(all_preds),
                     "pred_gap_band_hi": max(all_preds),
                     "pred_gap_geo_mean": float(np.exp(np.mean(
                         [np.log(r["pred_gap"]) for r in width_seed_rows]))),
                     "rho_seed_mean": float(np.mean(
                         [r["rho"] for r in width_seed_rows])),
                     "lam_min_seed_mean": float(np.mean(
                         [r["lam_min"] for r in width_seed_rows]))})
        print(f"{width:>3}{margin:>7}{d['rho']:>7.3f}{d['lam_min']:>9.2f}"
              f"{d['pred_gap']:>11.5f}{meas:>11.5f}", flush=True)

    csv_path = args.outdir / "closed_form_vs_measured.csv"
    with csv_path.open("w", newline="") as fh:
        w = csv.DictWriter(fh, fieldnames=list(rows[0].keys()))
        w.writeheader()
        w.writerows(rows)
    seed_csv_path = args.outdir / "closed_form_per_seed.csv"
    with seed_csv_path.open("w", newline="") as fh:
        w = csv.DictWriter(fh, fieldnames=list(seed_rows[0].keys()))
        w.writeheader()
        w.writerows(seed_rows)

    margins = np.array([r["margin"] for r in rows])
    measured = np.array([r["measured_gap"] for r in rows])
    measured_std = np.array([r["measured_gap_std"] for r in rows])
    pred = np.array([r["pred_gap"] for r in rows])
    band_lo = np.array([r["pred_gap_band_lo"] for r in rows])
    band_hi = np.array([r["pred_gap_band_hi"] for r in rows])
    lam_min = np.array([r["lam_min"] for r in rows])
    pred_geo = np.array([r["pred_gap_geo_mean"] for r in rows])
    lam_min_mean = np.array([r["lam_min_seed_mean"] for r in rows])
    lm = np.log(measured)
    fit = np.polyfit(margins, lm, 1)
    corr_pred = np.corrcoef(np.log(np.clip(pred, 1e-12, None)), lm)[0, 1]
    corr_geo = np.corrcoef(np.log(np.clip(pred_geo, 1e-12, None)), lm)[0, 1]
    corr_lam = np.corrcoef(lam_min, -lm)[0, 1]
    corr_lam_mean = np.corrcoef(lam_min_mean, -lm)[0, 1]
    r2 = np.corrcoef(margins, lm)[0, 1] ** 2

    fig, axes = plt.subplots(1, 2, figsize=(12.0, 4.8), dpi=180)
    axes[0].errorbar(margins, measured, yerr=measured_std, fmt="o-", color="#174f91",
                     capsize=2.5, label="empirical (T=30000, 32 feedback seeds)")
    axes[0].semilogy(margins, pred_geo, "s-", color="#c65f16",
                     label=f"closed form, {len(band_seeds) + 1}-init geo mean (no fit)")
    axes[0].semilogy(margins, pred, "^:", color="#7a4a13", alpha=0.7, markersize=4,
                     label="closed form, matched init")
    axes[0].fill_between(margins, np.clip(band_lo, 1e-12, None), band_hi,
                         color="#c65f16", alpha=0.18,
                         label="init-seed range")
    axes[0].set_yscale("log")
    axes[0].axvline(0, color="grey", linewidth=1, linestyle=":")
    axes[0].set_xlabel("hard FA capacity margin")
    axes[0].set_ylabel("FA/BP train gap")
    axes[0].set_title(f"Soft ramp: closed form vs empirical\ncorr(log,log)={corr_geo:.3f} (init mean), no kink at margin 0")
    axes[0].legend(fontsize=8)
    axes[0].grid(alpha=0.18, which="both")

    axes[1].plot(lam_min_mean, -lm, "o", color="#174f91")
    coef = np.polyfit(lam_min_mean, -lm, 1)
    grid = np.linspace(lam_min_mean.min(), lam_min_mean.max(), 50)
    axes[1].plot(grid, np.polyval(coef, grid), "--", color="#c65f16")
    axes[1].set_xlabel("smallest BP-NTK eigenvalue, init mean  lam_min(w)")
    axes[1].set_ylabel("-log(empirical gap)")
    axes[1].set_title(f"Ramp law: log gap ~ const - c.rho.lam_min\ncorr={corr_lam_mean:.3f}")
    axes[1].grid(alpha=0.18)
    fig.tight_layout()
    fig_path = args.outdir / "closed_form_vs_measured_ramp.png"
    fig.savefig(fig_path)
    plt.close(fig)

    print(f"\nmeasured log(gap) ~ {fit[0]:.4f}*margin + {fit[1]:.2f}  (R^2={r2:.3f})")
    print(f"corr(log pred, log measured) = {corr_pred:.3f}  (matched init)")
    print(f"corr(log pred, log measured) = {corr_geo:.3f}  ({len(band_seeds) + 1}-init geo mean)")
    print(f"corr(lam_min, -log measured) = {corr_lam:.3f}  (matched init)")
    print(f"corr(lam_min, -log measured) = {corr_lam_mean:.3f}  (init mean)")
    print(f"rows: {csv_path}")
    print(f"per-seed rows: {seed_csv_path}")
    print(f"plot: {fig_path}")


if __name__ == "__main__":
    main()