From daf0a6dc795235864d8279908320801f3c7b66b5 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Tue, 2 Jun 2026 15:13:04 -0500 Subject: Replace transition theory with tangent capacity --- scripts/fa_tangent_kernel_capacity.py | 400 ++++++++++++++++++++++++++++++++++ 1 file changed, 400 insertions(+) create mode 100644 scripts/fa_tangent_kernel_capacity.py (limited to 'scripts/fa_tangent_kernel_capacity.py') 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() -- cgit v1.2.3