summaryrefslogtreecommitdiff
path: root/scripts/fa_tangent_kernel_capacity.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-06-02 15:13:04 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-06-02 15:13:04 -0500
commitdaf0a6dc795235864d8279908320801f3c7b66b5 (patch)
treeebe59a8768889cf5430c6a4ec142a2e3a7693a1b /scripts/fa_tangent_kernel_capacity.py
parentb436e6c033c2c3b37d2ee26efe0dbe28381117dc (diff)
Replace transition theory with tangent capacity
Diffstat (limited to 'scripts/fa_tangent_kernel_capacity.py')
-rw-r--r--scripts/fa_tangent_kernel_capacity.py400
1 files changed, 400 insertions, 0 deletions
diff --git a/scripts/fa_tangent_kernel_capacity.py b/scripts/fa_tangent_kernel_capacity.py
new file mode 100644
index 0000000..ca779cf
--- /dev/null
+++ b/scripts/fa_tangent_kernel_capacity.py
@@ -0,0 +1,400 @@
+#!/usr/bin/env python3
+"""Measure BP/FA tangent-kernel capacity for existing downstream sweeps."""
+
+from __future__ import annotations
+
+import argparse
+import csv
+import json
+import math
+import sys
+from dataclasses import dataclass
+from pathlib import Path
+
+import matplotlib.pyplot as plt
+import numpy as np
+import torch
+from scipy.sparse.linalg import expm_multiply
+
+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
+
+
+Tensor = torch.Tensor
+
+
+@dataclass(frozen=True)
+class SourceRun:
+ width: int
+ parameter_count: int
+ init_seed: int
+ feedback_seed: int
+ run_type: str
+ train_mse: float
+ bp_train_mse: float
+ train_gap_to_bp: float
+ fa_capacity_margin: int
+
+
+@dataclass(frozen=True)
+class KernelRow:
+ width: int
+ init_seed: int
+ feedback_seed: int
+ fa_capacity_margin: int
+ empirical_train_gap: float
+ predicted_bp_loss: float
+ predicted_fa_loss: float
+ predicted_train_gap: float
+ bp_rank: int
+ fa_rank: int
+ bp_d_eff: float
+ fa_d_eff: float
+ bp_lambda: float
+ fa_lambda: float
+ fa_sym_min: float
+ fa_sym_neg_mass: float
+ fa_trace: float
+ bp_trace: float
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser(
+ description="Compute full-training-set BP/FA tangent-kernel capacity."
+ )
+ parser.add_argument(
+ "--source-outdir",
+ type=Path,
+ default=Path("outputs/downstream_capacity_random_main_fast"),
+ )
+ parser.add_argument(
+ "--outdir",
+ type=Path,
+ default=Path("outputs/fa_tangent_kernel_capacity"),
+ )
+ parser.add_argument("--max-fa-per-width", type=int, default=0)
+ parser.add_argument("--max-widths", type=int, default=0)
+ parser.add_argument("--sigma-rel", type=float, default=1e-3)
+ parser.add_argument("--rank-tol-rel", type=float, default=1e-8)
+ parser.add_argument(
+ "--prediction",
+ choices=["continuous"],
+ default="continuous",
+ help="continuous uses exp(-(lr*steps/N)K) on the initial residual.",
+ )
+ parser.add_argument("--plot", action="store_true")
+ return parser.parse_args()
+
+
+def load_config(source_outdir: Path) -> dcs.RunConfig:
+ payload = json.loads((source_outdir / "summary.json").read_text())
+ raw = payload["config"]
+ raw["outdir"] = str(source_outdir)
+ raw.setdefault("init_seed_offset", 0)
+ raw.setdefault("feedback_seed_offset", 0)
+ raw.setdefault("skip_jacobian", False)
+ return dcs.RunConfig(**raw)
+
+
+def load_runs(source_outdir: Path) -> list[SourceRun]:
+ rows: list[SourceRun] = []
+ with (source_outdir / "runs.csv").open(newline="") as handle:
+ reader = csv.DictReader(handle)
+ for raw in reader:
+ rows.append(
+ SourceRun(
+ width=int(raw["width"]),
+ parameter_count=int(raw["parameter_count"]),
+ init_seed=int(raw["init_seed"]),
+ feedback_seed=int(raw["feedback_seed"]),
+ run_type=raw["run_type"],
+ train_mse=float(raw["train_mse"]),
+ bp_train_mse=float(raw["bp_train_mse"]),
+ train_gap_to_bp=float(raw["train_gap_to_bp"]),
+ fa_capacity_margin=int(raw["fa_capacity_margin"]),
+ )
+ )
+ return rows
+
+
+def flatten_grads(grads: list[Tensor]) -> Tensor:
+ return torch.cat([grad.reshape(-1) for grad in grads])
+
+
+def pseudo_jacobian(
+ weights: list[Tensor],
+ x: Tensor,
+ feedback: list[Tensor] | None,
+) -> Tensor:
+ activations, preacts = dcs.forward(weights, x)
+ outputs = activations[-1]
+ sample_count, output_dim = outputs.shape
+ rows: list[Tensor] = []
+
+ for index in range(sample_count * output_dim):
+ delta_out = torch.zeros_like(outputs)
+ delta_out.reshape(-1)[index] = 1.0
+ deltas: list[Tensor] = [
+ torch.empty(0, dtype=torch.float64, device=x.device) for _ in weights
+ ]
+ deltas[-1] = delta_out
+ for layer in range(len(weights) - 2, -1, -1):
+ if feedback is None:
+ back = deltas[layer + 1] @ weights[layer + 1]
+ else:
+ back = deltas[layer + 1] @ feedback[layer].T
+ deltas[layer] = back * (preacts[layer] > 0)
+ grads = [delta.T @ activations[layer] for layer, delta in enumerate(deltas)]
+ rows.append(flatten_grads(grads).detach())
+
+ return torch.stack(rows, dim=0)
+
+
+def symmetric_capacity(
+ kernel: np.ndarray,
+ sigma_rel: float,
+ rank_tol_rel: float,
+) -> tuple[int, float, float, float, float]:
+ sym = 0.5 * (kernel + kernel.T)
+ eigenvalues = np.linalg.eigvalsh(sym)
+ positive = np.clip(eigenvalues, 0.0, None)
+ max_eval = float(np.max(positive)) if positive.size else 0.0
+ lam = max(sigma_rel * max_eval, 1e-12)
+ d_eff = float(np.sum(positive / (positive + lam)))
+ rank = int(np.count_nonzero(positive > rank_tol_rel * max_eval)) if max_eval > 0 else 0
+ neg_mass = float(np.sum(np.clip(-eigenvalues, 0.0, None)))
+ min_eval = float(np.min(eigenvalues)) if eigenvalues.size else 0.0
+ trace = float(np.trace(kernel))
+ return rank, d_eff, lam, min_eval, neg_mass, trace
+
+
+def predict_loss(
+ kernel: np.ndarray,
+ residual: np.ndarray,
+ lr: float,
+ steps: int,
+ samples: int,
+) -> float:
+ scale = -(lr * steps / samples)
+ evolved = expm_multiply(scale * kernel, residual)
+ return 0.5 * float(np.dot(evolved, evolved)) / samples
+
+
+def select_fa_rows(
+ rows: list[SourceRun],
+ max_fa_per_width: int,
+ max_widths: int,
+) -> list[SourceRun]:
+ fa_rows = [row for row in rows if row.run_type == "fa"]
+ widths = sorted({row.width for row in fa_rows})
+ if max_widths > 0:
+ widths = widths[:max_widths]
+ selected: list[SourceRun] = []
+ for width in widths:
+ width_rows = [row for row in fa_rows if row.width == width]
+ width_rows.sort(key=lambda row: (row.init_seed, row.feedback_seed))
+ if max_fa_per_width > 0:
+ width_rows = width_rows[:max_fa_per_width]
+ selected.extend(width_rows)
+ return selected
+
+
+def write_kernel_rows(path: Path, rows: list[KernelRow]) -> None:
+ path.parent.mkdir(parents=True, exist_ok=True)
+ with path.open("w", newline="") as handle:
+ writer = csv.DictWriter(handle, fieldnames=list(KernelRow.__annotations__.keys()))
+ writer.writeheader()
+ for row in rows:
+ writer.writerow(row.__dict__)
+
+
+def plot_results(rows: list[KernelRow], outdir: Path) -> list[Path]:
+ paths: list[Path] = []
+ outdir.mkdir(parents=True, exist_ok=True)
+
+ scatter_path = outdir / "predicted_vs_empirical_train_gap.png"
+ plt.figure(figsize=(6.2, 5.2))
+ widths = sorted({row.width for row in rows})
+ for width in widths:
+ subset = [row for row in rows if row.width == width]
+ plt.scatter(
+ [row.predicted_train_gap for row in subset],
+ [row.empirical_train_gap for row in subset],
+ s=22,
+ alpha=0.55,
+ label=f"n={width}",
+ )
+ all_values = [
+ value
+ for row in rows
+ for value in (row.predicted_train_gap, row.empirical_train_gap)
+ if math.isfinite(value)
+ ]
+ if all_values:
+ lo = min(all_values)
+ hi = max(all_values)
+ plt.plot([lo, hi], [lo, hi], color="black", linewidth=1)
+ plt.xlabel("linearized FA/BP train-gap prediction")
+ plt.ylabel("empirical FA/BP train gap")
+ plt.title("Tangent-kernel prediction vs trajectory gap")
+ plt.legend(fontsize=8, ncols=2)
+ plt.tight_layout()
+ plt.savefig(scatter_path, dpi=180)
+ plt.close()
+ paths.append(scatter_path)
+
+ transition_path = outdir / "kernel_capacity_transition.png"
+ plt.figure(figsize=(7.5, 4.8))
+ for row in rows:
+ plt.plot(
+ [row.fa_capacity_margin, row.fa_capacity_margin],
+ [row.predicted_train_gap, row.empirical_train_gap],
+ color="0.7",
+ alpha=0.18,
+ linewidth=0.8,
+ )
+ plt.scatter(
+ [row.fa_capacity_margin for row in rows],
+ [row.predicted_train_gap for row in rows],
+ s=20,
+ alpha=0.55,
+ color="tab:blue",
+ label="kernel prediction",
+ )
+ plt.scatter(
+ [row.fa_capacity_margin for row in rows],
+ [row.empirical_train_gap for row in rows],
+ s=20,
+ alpha=0.35,
+ color="tab:orange",
+ label="trajectory",
+ )
+ plt.axvline(0.0, color="black", linewidth=1, linestyle="--")
+ plt.axhline(0.0, color="black", linewidth=1)
+ plt.xlabel("old hard FA margin, for reference")
+ plt.ylabel("FA train MSE - BP train MSE")
+ plt.title("Predicted and empirical transition")
+ plt.legend()
+ plt.tight_layout()
+ plt.savefig(transition_path, dpi=180)
+ plt.close()
+ paths.append(transition_path)
+
+ capacity_path = outdir / "fa_dynamic_effective_dimension_vs_margin.png"
+ plt.figure(figsize=(7.5, 4.8))
+ plt.scatter(
+ [row.fa_capacity_margin for row in rows],
+ [row.fa_d_eff for row in rows],
+ s=20,
+ alpha=0.55,
+ color="tab:green",
+ )
+ if rows:
+ task_dim = max(row.bp_rank for row in rows)
+ plt.axhline(task_dim, color="black", linewidth=1, linestyle="--", label="observed max BP rank")
+ plt.axvline(0.0, color="black", linewidth=1, linestyle="--")
+ plt.xlabel("old hard FA margin, for reference")
+ plt.ylabel("FA dynamic effective dimension")
+ plt.title("Measured FA tangent-kernel capacity")
+ plt.legend()
+ plt.tight_layout()
+ plt.savefig(capacity_path, dpi=180)
+ plt.close()
+ paths.append(capacity_path)
+
+ return paths
+
+
+def main() -> None:
+ args = parse_args()
+ config = load_config(args.source_outdir)
+ if config.optimizer != "sgd":
+ print(
+ "warning: source trajectories used optimizer="
+ f"{config.optimizer!r}; continuous tangent-kernel prediction is an SGD "
+ "linearization diagnostic, not a final-gap theorem for this run.",
+ flush=True,
+ )
+ if config.torch_threads > 0:
+ torch.set_num_threads(config.torch_threads)
+
+ source_rows = load_runs(args.source_outdir)
+ fa_rows = select_fa_rows(source_rows, args.max_fa_per_width, args.max_widths)
+ x_train, y_train, _x_test, _y_test, _x_probe, _teacher = dcs.make_data(config)
+
+ result_rows: list[KernelRow] = []
+ bp_cache: dict[tuple[int, int], tuple[np.ndarray, np.ndarray, float, int, float, float, float]] = {}
+
+ for index, row in enumerate(fa_rows, start=1):
+ print(
+ f"[{index}/{len(fa_rows)}] width={row.width} init={row.init_seed} "
+ f"feedback={row.feedback_seed}",
+ flush=True,
+ )
+ initial_weights = dcs.initialize_weights(config, row.width, row.init_seed)
+ feedback = dcs.init_feedback(config, row.width, row.feedback_seed)
+ cache_key = (row.width, row.init_seed)
+ if cache_key not in bp_cache:
+ with torch.no_grad():
+ initial_residual = (dcs.predict(initial_weights, x_train) - y_train).reshape(-1)
+ j_bp = pseudo_jacobian(initial_weights, x_train, feedback=None).cpu().numpy()
+ k_bp = j_bp @ j_bp.T
+ residual = initial_residual.detach().cpu().numpy()
+ bp_loss = predict_loss(k_bp, residual, config.lr, config.steps, config.train_samples)
+ bp_rank, bp_d_eff, bp_lam, _bp_min, _bp_neg, bp_trace = symmetric_capacity(
+ k_bp, args.sigma_rel, args.rank_tol_rel
+ )
+ bp_cache[cache_key] = (
+ j_bp,
+ residual,
+ bp_loss,
+ bp_rank,
+ bp_d_eff,
+ bp_lam,
+ bp_trace,
+ )
+
+ j_bp, residual, bp_loss, bp_rank, bp_d_eff, bp_lam, bp_trace = bp_cache[cache_key]
+ j_fa = pseudo_jacobian(initial_weights, x_train, feedback=feedback).cpu().numpy()
+ k_fa = j_bp @ j_fa.T
+ fa_loss = predict_loss(k_fa, residual, config.lr, config.steps, config.train_samples)
+ fa_rank, fa_d_eff, fa_lam, fa_min, fa_neg, fa_trace = symmetric_capacity(
+ k_fa, args.sigma_rel, args.rank_tol_rel
+ )
+ result_rows.append(
+ KernelRow(
+ width=row.width,
+ init_seed=row.init_seed,
+ feedback_seed=row.feedback_seed,
+ fa_capacity_margin=row.fa_capacity_margin,
+ empirical_train_gap=row.train_gap_to_bp,
+ predicted_bp_loss=bp_loss,
+ predicted_fa_loss=fa_loss,
+ predicted_train_gap=fa_loss - bp_loss,
+ bp_rank=bp_rank,
+ fa_rank=fa_rank,
+ bp_d_eff=bp_d_eff,
+ fa_d_eff=fa_d_eff,
+ bp_lambda=bp_lam,
+ fa_lambda=fa_lam,
+ fa_sym_min=fa_min,
+ fa_sym_neg_mass=fa_neg,
+ fa_trace=fa_trace,
+ bp_trace=bp_trace,
+ )
+ )
+
+ args.outdir.mkdir(parents=True, exist_ok=True)
+ csv_path = args.outdir / "kernel_capacity_rows.csv"
+ write_kernel_rows(csv_path, result_rows)
+ print(f"rows: {csv_path}")
+ if args.plot:
+ for path in plot_results(result_rows, args.outdir):
+ print(f"plot: {path}")
+
+
+if __name__ == "__main__":
+ main()