#!/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()