#!/usr/bin/env python3 """Downstream BP/FA capacity sweep on controlled capacity tasks.""" from __future__ import annotations import argparse import csv import json import math from dataclasses import asdict, dataclass from pathlib import Path import matplotlib.pyplot as plt import numpy as np import torch from scipy import stats Tensor = torch.Tensor @dataclass(frozen=True) class RunConfig: task: str input_dim: int output_dim: int teacher_rank: int teacher_width: int teacher_hidden_layers: int normalize_targets: bool widths: list[int] train_samples: int test_samples: int probe_samples: int steps: int lr: float optimizer: str init_seeds: int feedback_seeds: int init_seed_offset: int feedback_seed_offset: int data_seed: int noise_std: float feedback_scale: str capacity_q: float jacobian_lambda_rel: float skip_jacobian: bool device: str torch_threads: int outdir: str plot: bool @dataclass(frozen=True) class RunRow: width: int parameter_count: int init_seed: int feedback_seed: int run_type: str train_mse: float test_mse: float bp_train_mse: float bp_test_mse: float train_gap_to_bp: float test_gap_to_bp: float task_dimension: int fa_constraint_rank: int bp_capacity_margin: int fa_capacity_margin: int d_eff: float hard_rank: int jacobian_lambda: float fa_burden_nats: float redundancy_score: float @dataclass(frozen=True) class WidthSummaryRow: width: int parameter_count: int runs_bp: int runs_fa: int bp_test_mean: float bp_test_std: float bp_train_mean: float bp_train_std: float fa_test_mean: float fa_test_std: float fa_train_mean: float fa_train_std: float gap_mean: float gap_std: float train_gap_mean: float train_gap_std: float task_dimension: int fa_constraint_rank: int bp_capacity_margin: int fa_capacity_margin: int d_eff_mean: float hard_rank_mean: float fa_burden_nats: float redundancy_score_mean: float def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser( description="Capacity shrinkage experiment: BP vs FA on controlled tasks." ) parser.add_argument("--input-dim", type=int, default=16) parser.add_argument("--output-dim", type=int, default=4) parser.add_argument("--task", choices=["linear", "mlp", "random"], default="linear") parser.add_argument("--teacher-rank", type=int, default=4) parser.add_argument("--teacher-width", type=int, default=64) parser.add_argument("--teacher-hidden-layers", type=int, default=2) parser.add_argument("--normalize-targets", action="store_true") parser.add_argument("--widths", type=int, nargs="+", default=[2, 4, 8, 16, 32]) parser.add_argument("--train-samples", type=int, default=256) parser.add_argument("--test-samples", type=int, default=1024) parser.add_argument("--probe-samples", type=int, default=64) parser.add_argument("--steps", type=int, default=800) parser.add_argument("--lr", type=float, default=0.03) parser.add_argument("--optimizer", choices=["sgd", "adam"], default="sgd") parser.add_argument("--init-seeds", type=int, default=3) parser.add_argument("--feedback-seeds", type=int, default=5) parser.add_argument("--init-seed-offset", type=int, default=0) parser.add_argument("--feedback-seed-offset", type=int, default=0) parser.add_argument("--data-seed", type=int, default=0) parser.add_argument("--noise-std", type=float, default=0.0) parser.add_argument( "--feedback-scale", choices=["relu", "fan-in", "unit"], default="relu", ) parser.add_argument("--capacity-q", type=float, default=0.01) parser.add_argument("--jacobian-lambda-rel", type=float, default=1e-3) parser.add_argument("--skip-jacobian", action="store_true") parser.add_argument("--device", choices=["cpu", "cuda"], default="cpu") parser.add_argument("--torch-threads", type=int, default=0) parser.add_argument( "--outdir", type=Path, default=Path("outputs/downstream_capacity_sweep"), ) parser.add_argument("--plot", action="store_true") return parser.parse_args() def parse_config(args: argparse.Namespace) -> RunConfig: return RunConfig( task=args.task, input_dim=args.input_dim, output_dim=args.output_dim, teacher_rank=args.teacher_rank, teacher_width=args.teacher_width, teacher_hidden_layers=args.teacher_hidden_layers, normalize_targets=args.normalize_targets, widths=args.widths, train_samples=args.train_samples, test_samples=args.test_samples, probe_samples=args.probe_samples, steps=args.steps, lr=args.lr, optimizer=args.optimizer, init_seeds=args.init_seeds, feedback_seeds=args.feedback_seeds, init_seed_offset=args.init_seed_offset, feedback_seed_offset=args.feedback_seed_offset, data_seed=args.data_seed, noise_std=args.noise_std, feedback_scale=args.feedback_scale, capacity_q=args.capacity_q, jacobian_lambda_rel=args.jacobian_lambda_rel, skip_jacobian=args.skip_jacobian, device=args.device, torch_threads=args.torch_threads, outdir=str(args.outdir), plot=args.plot, ) def validate_config(config: RunConfig) -> None: if config.input_dim < 1 or config.output_dim < 1: raise ValueError("Input and output dimensions must be positive.") if not 1 <= config.teacher_rank <= min(config.input_dim, config.output_dim): raise ValueError("teacher-rank must be in [1, min(input_dim, output_dim)].") if config.teacher_width < 1 or config.teacher_hidden_layers < 1: raise ValueError("Teacher MLP dimensions must be positive.") if any(width < 1 for width in config.widths): raise ValueError("All widths must be positive.") if config.train_samples < 1 or config.test_samples < 1 or config.probe_samples < 1: raise ValueError("Sample counts must be positive.") if config.steps < 1: raise ValueError("--steps must be positive.") if config.lr <= 0: raise ValueError("--lr must be positive.") if config.init_seeds < 1 or config.feedback_seeds < 1: raise ValueError("Seed counts must be positive.") if not 0 <= config.capacity_q <= 1: raise ValueError("--capacity-q must be in [0, 1].") if config.jacobian_lambda_rel <= 0: raise ValueError("--jacobian-lambda-rel must be positive.") def make_low_rank_teacher(config: RunConfig, rng: np.random.Generator) -> np.ndarray: u, _ = np.linalg.qr(rng.standard_normal((config.output_dim, config.teacher_rank))) v, _ = np.linalg.qr(rng.standard_normal((config.input_dim, config.teacher_rank))) singular = np.linspace(1.0, 0.4, config.teacher_rank) return (u * singular) @ v.T def numpy_teacher_mlp(config: RunConfig, rng: np.random.Generator) -> list[np.ndarray]: dims = [ config.input_dim, *([config.teacher_width] * config.teacher_hidden_layers), config.output_dim, ] weights: list[np.ndarray] = [] for layer, (fan_in, fan_out) in enumerate(zip(dims[:-1], dims[1:])): scale = math.sqrt(2.0 / fan_in) if layer < len(dims) - 2 else 1.0 / math.sqrt(fan_in) weights.append(rng.standard_normal((fan_out, fan_in)) * scale) return weights def numpy_teacher_predict(weights: list[np.ndarray], x: np.ndarray) -> np.ndarray: current = x for layer, weight in enumerate(weights): current = current @ weight.T if layer < len(weights) - 1: current = np.maximum(current, 0.0) return current def make_data(config: RunConfig) -> tuple[Tensor, Tensor, Tensor, Tensor, Tensor, object]: rng = np.random.default_rng(config.data_seed) teacher: object if config.task == "linear": teacher = make_low_rank_teacher(config, rng) elif config.task == "mlp": teacher = numpy_teacher_mlp(config, rng) else: teacher = None def sample_x(count: int) -> np.ndarray: x = rng.standard_normal((count, config.input_dim)) return x.astype(np.float64) def teacher_y(x: np.ndarray) -> np.ndarray: if config.task == "linear": y = x @ teacher.T # type: ignore[union-attr] elif config.task == "mlp": y = numpy_teacher_predict(teacher, x) # type: ignore[arg-type] else: y = rng.standard_normal((x.shape[0], config.output_dim)) if config.noise_std > 0: y = y + rng.standard_normal(y.shape) * config.noise_std return y.astype(np.float64) x_train = sample_x(config.train_samples) y_train = teacher_y(x_train) x_test = sample_x(config.test_samples) y_test = teacher_y(x_test) x_probe = sample_x(config.probe_samples) if config.normalize_targets: mean = y_train.mean(axis=0, keepdims=True) std = y_train.std(axis=0, keepdims=True) + 1e-8 y_train = (y_train - mean) / std y_test = (y_test - mean) / std device = torch.device(config.device) return ( torch.tensor(x_train, dtype=torch.float64, device=device), torch.tensor(y_train, dtype=torch.float64, device=device), torch.tensor(x_test, dtype=torch.float64, device=device), torch.tensor(y_test, dtype=torch.float64, device=device), torch.tensor(x_probe, dtype=torch.float64, device=device), teacher, ) def initialize_weights(config: RunConfig, width: int, seed: int) -> list[Tensor]: generator = torch.Generator(device=config.device) generator.manual_seed(seed) dims = [config.input_dim, width, width, config.output_dim] weights: list[Tensor] = [] for layer, (fan_in, fan_out) in enumerate(zip(dims[:-1], dims[1:])): scale = math.sqrt(2.0 / fan_in) if layer < len(dims) - 2 else 1.0 / math.sqrt(fan_in) weight = torch.randn( fan_out, fan_in, generator=generator, dtype=torch.float64, device=config.device, ) * scale weights.append(weight) return weights def clone_weights(weights: list[Tensor]) -> list[Tensor]: return [weight.clone() for weight in weights] def forward(weights: list[Tensor], x: Tensor) -> tuple[list[Tensor], list[Tensor]]: activations = [x] preacts: list[Tensor] = [] current = x for layer, weight in enumerate(weights): preact = current @ weight.T preacts.append(preact) current = torch.relu(preact) if layer < len(weights) - 1 else preact activations.append(current) return activations, preacts def predict(weights: list[Tensor], x: Tensor) -> Tensor: return forward(weights, x)[0][-1] def mse(weights: list[Tensor], x: Tensor, y: Tensor) -> float: with torch.no_grad(): error = predict(weights, x) - y return float(0.5 * torch.mean(torch.sum(error * error, dim=1)).cpu()) def gradients(weights: list[Tensor], x: Tensor, y: Tensor, feedback: list[Tensor] | None) -> list[Tensor]: activations, preacts = forward(weights, x) pred = activations[-1] batch = x.shape[0] deltas: list[Tensor] = [torch.empty(0, dtype=torch.float64, device=x.device) for _ in weights] deltas[-1] = (pred - y) / batch 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) return [delta.T @ activations[layer] for layer, delta in enumerate(deltas)] def train( weights_init: list[Tensor], x: Tensor, y: Tensor, lr: float, steps: int, feedback: list[Tensor] | None, optimizer: str, ) -> list[Tensor]: weights = clone_weights(weights_init) first_moment = [torch.zeros_like(weight) for weight in weights] second_moment = [torch.zeros_like(weight) for weight in weights] beta1 = 0.9 beta2 = 0.999 eps = 1e-8 for step in range(1, steps + 1): grads = gradients(weights, x, y, feedback) for index, grad in enumerate(grads): if optimizer == "sgd": update = grad elif optimizer == "adam": first_moment[index] = beta1 * first_moment[index] + (1.0 - beta1) * grad second_moment[index] = beta2 * second_moment[index] + (1.0 - beta2) * (grad * grad) first_hat = first_moment[index] / (1.0 - beta1**step) second_hat = second_moment[index] / (1.0 - beta2**step) update = first_hat / (torch.sqrt(second_hat) + eps) else: raise ValueError(f"Unknown optimizer: {optimizer}") weights[index] = weights[index] - lr * update return weights def feedback_scale(rows: int, mode: str) -> float: if mode == "relu": return math.sqrt(2.0 / rows) if mode == "fan-in": return math.sqrt(1.0 / rows) if mode == "unit": return 1.0 raise ValueError(f"Unknown feedback scale: {mode}") def init_feedback(config: RunConfig, width: int, seed: int) -> list[Tensor]: generator = torch.Generator(device=config.device) generator.manual_seed(seed) shapes = [(width, width), (width, config.output_dim)] feedback: list[Tensor] = [] for rows, cols in shapes: matrix = torch.randn( rows, cols, generator=generator, dtype=torch.float64, device=config.device, ) * feedback_scale(rows, config.feedback_scale) feedback.append(matrix) return feedback def parameter_count(config: RunConfig, width: int) -> int: dims = [config.input_dim, width, width, config.output_dim] return sum(fan_in * fan_out for fan_in, fan_out in zip(dims[:-1], dims[1:])) def task_dimension(config: RunConfig) -> int: return config.train_samples * config.output_dim def fa_constraint_rank(config: RunConfig, width: int) -> int: feedback_target_dims = [width * width, width * config.output_dim] return sum(max(dim - 1, 0) for dim in feedback_target_dims) def fa_burden_nats(config: RunConfig, width: int) -> float: dims = [width * width, width * config.output_dim] total = 0.0 for dim in dims: total += float(-stats.beta(0.5, (dim - 1) / 2).logsf(config.capacity_q)) return total def flatten_params(params: list[Tensor]) -> Tensor: return torch.cat([param.reshape(-1) for param in params]) def jacobian_effective_dimension( weights: list[Tensor], x_probe: Tensor, lambda_rel: float, ) -> tuple[float, int, float]: params = [weight.clone().detach().requires_grad_(True) for weight in weights] outputs = predict(params, x_probe).reshape(-1) rows: list[Tensor] = [] for output in outputs: grads = torch.autograd.grad(output, params, retain_graph=True) rows.append(flatten_params(list(grads)).detach()) jacobian = torch.stack(rows, dim=0) kernel = jacobian @ jacobian.T eigenvalues = torch.linalg.eigvalsh(kernel).clamp_min(0.0) max_eval = float(torch.max(eigenvalues).cpu()) lam = max(lambda_rel * max_eval, 1e-12) d_eff = float(torch.sum(eigenvalues / (eigenvalues + lam)).cpu()) hard_rank = int(torch.count_nonzero(eigenvalues > (1e-8 * max_eval)).cpu()) if max_eval > 0 else 0 return d_eff, hard_rank, lam def write_csv(path: Path, rows: list[object]) -> None: if not rows: return path.parent.mkdir(parents=True, exist_ok=True) with path.open("w", newline="") as handle: first = asdict(rows[0]) # type: ignore[arg-type] writer = csv.DictWriter(handle, fieldnames=list(first.keys())) writer.writeheader() for row in rows: writer.writerow(asdict(row)) # type: ignore[arg-type] def summarize_widths(rows: list[RunRow]) -> list[WidthSummaryRow]: summaries: list[WidthSummaryRow] = [] for width in sorted({row.width for row in rows}): width_rows = [row for row in rows if row.width == width] bp_rows = [row for row in width_rows if row.run_type == "bp"] fa_rows = [row for row in width_rows if row.run_type == "fa"] bp_test = np.array([row.test_mse for row in bp_rows], dtype=np.float64) bp_train = np.array([row.train_mse for row in bp_rows], dtype=np.float64) fa_test = np.array([row.test_mse for row in fa_rows], dtype=np.float64) fa_train = np.array([row.train_mse for row in fa_rows], dtype=np.float64) gaps = np.array([row.test_gap_to_bp for row in fa_rows], dtype=np.float64) train_gaps = np.array([row.train_gap_to_bp for row in fa_rows], dtype=np.float64) d_eff = np.array([row.d_eff for row in bp_rows], dtype=np.float64) hard_rank = np.array([row.hard_rank for row in bp_rows], dtype=np.float64) burden = fa_rows[0].fa_burden_nats if fa_rows else bp_rows[0].fa_burden_nats redundancy = np.array([row.redundancy_score for row in bp_rows], dtype=np.float64) first = width_rows[0] summaries.append( WidthSummaryRow( width=width, parameter_count=first.parameter_count, runs_bp=len(bp_rows), runs_fa=len(fa_rows), bp_test_mean=float(np.mean(bp_test)), bp_test_std=float(np.std(bp_test, ddof=1)) if len(bp_test) > 1 else 0.0, bp_train_mean=float(np.mean(bp_train)), bp_train_std=float(np.std(bp_train, ddof=1)) if len(bp_train) > 1 else 0.0, fa_test_mean=float(np.mean(fa_test)), fa_test_std=float(np.std(fa_test, ddof=1)) if len(fa_test) > 1 else 0.0, fa_train_mean=float(np.mean(fa_train)), fa_train_std=float(np.std(fa_train, ddof=1)) if len(fa_train) > 1 else 0.0, gap_mean=float(np.mean(gaps)), gap_std=float(np.std(gaps, ddof=1)) if len(gaps) > 1 else 0.0, train_gap_mean=float(np.mean(train_gaps)), train_gap_std=float(np.std(train_gaps, ddof=1)) if len(train_gaps) > 1 else 0.0, task_dimension=first.task_dimension, fa_constraint_rank=first.fa_constraint_rank, bp_capacity_margin=first.bp_capacity_margin, fa_capacity_margin=first.fa_capacity_margin, d_eff_mean=float(np.mean(d_eff)), hard_rank_mean=float(np.mean(hard_rank)), fa_burden_nats=float(burden), redundancy_score_mean=float(np.mean(redundancy)), ) ) return summaries def add_capacity_thresholds(summaries: list[WidthSummaryRow]) -> None: bp_candidates = [row.width for row in summaries if row.bp_capacity_margin >= 0] fa_candidates = [row.width for row in summaries if row.fa_capacity_margin >= 0] if bp_candidates: plt.axvline( min(bp_candidates), color="tab:blue", linestyle="--", linewidth=1, alpha=0.65, label="BP P>=N*out", ) if fa_candidates: plt.axvline( min(fa_candidates), color="tab:orange", linestyle="--", linewidth=1, alpha=0.65, label="FA P-K>=N*out", ) def save_plots(summaries: list[WidthSummaryRow], outdir: Path) -> list[Path]: outdir.mkdir(parents=True, exist_ok=True) paths: list[Path] = [] widths = np.array([row.width for row in summaries]) perf_path = outdir / "bp_fa_test_mse_vs_width.png" plt.figure(figsize=(7, 4.5)) plt.errorbar( widths, [row.bp_test_mean for row in summaries], yerr=[row.bp_test_std for row in summaries], marker="o", label="BP", ) plt.errorbar( widths, [row.fa_test_mean for row in summaries], yerr=[row.fa_test_std for row in summaries], marker="o", label="FA", ) plt.xscale("log", base=2) plt.yscale("log") plt.xlabel("hidden width") plt.ylabel("test MSE") plt.title("Downstream capacity sweep") add_capacity_thresholds(summaries) plt.legend() plt.tight_layout() plt.savefig(perf_path, dpi=180) plt.close() paths.append(perf_path) train_path = outdir / "bp_fa_train_mse_vs_width.png" plt.figure(figsize=(7, 4.5)) plt.errorbar( widths, [row.bp_train_mean for row in summaries], yerr=[row.bp_train_std for row in summaries], marker="o", label="BP", ) plt.errorbar( widths, [row.fa_train_mean for row in summaries], yerr=[row.fa_train_std for row in summaries], marker="o", label="FA", ) plt.xscale("log", base=2) plt.yscale("log") plt.xlabel("hidden width") plt.ylabel("train MSE") plt.title("Memorization capacity sweep") add_capacity_thresholds(summaries) plt.legend() plt.tight_layout() plt.savefig(train_path, dpi=180) plt.close() paths.append(train_path) gap_path = outdir / "gap_vs_width.png" plt.figure(figsize=(7, 4.5)) plt.errorbar( widths, [row.gap_mean for row in summaries], yerr=[row.gap_std for row in summaries], marker="o", ) plt.axhline(0.0, color="black", linewidth=1) plt.xscale("log", base=2) plt.xlabel("hidden width") plt.ylabel("FA test MSE - BP test MSE") plt.title("FA/BP gap vs width") plt.tight_layout() plt.savefig(gap_path, dpi=180) plt.close() paths.append(gap_path) train_gap_path = outdir / "train_gap_vs_width.png" plt.figure(figsize=(7, 4.5)) plt.errorbar( widths, [row.train_gap_mean for row in summaries], yerr=[row.train_gap_std for row in summaries], marker="o", ) plt.axhline(0.0, color="black", linewidth=1) plt.xscale("log", base=2) plt.xlabel("hidden width") plt.ylabel("FA train MSE - BP train MSE") plt.title("FA/BP memorization gap vs width") add_capacity_thresholds(summaries) plt.tight_layout() plt.savefig(train_gap_path, dpi=180) plt.close() paths.append(train_gap_path) fa_margin_path = outdir / "train_gap_vs_fa_capacity_margin.png" plt.figure(figsize=(7, 4.5)) plt.errorbar( [row.fa_capacity_margin for row in summaries], [row.train_gap_mean for row in summaries], yerr=[row.train_gap_std for row in summaries], marker="o", linestyle="-", ) plt.axvline(0.0, color="black", linewidth=1) plt.axhline(0.0, color="black", linewidth=1) plt.xlabel("P - K_FA - N*out") plt.ylabel("FA train MSE - BP train MSE") plt.title("Memorization gap vs predicted FA margin") plt.tight_layout() plt.savefig(fa_margin_path, dpi=180) plt.close() paths.append(fa_margin_path) red_path = outdir / "gap_vs_redundancy_score.png" plt.figure(figsize=(7, 4.5)) plt.errorbar( [row.redundancy_score_mean for row in summaries], [row.gap_mean for row in summaries], yerr=[row.gap_std for row in summaries], marker="o", linestyle="-", ) plt.axvline(0.0, color="black", linewidth=1) plt.axhline(0.0, color="black", linewidth=1) plt.xlabel("P - d_eff - k_FA_proxy") plt.ylabel("FA test MSE - BP test MSE") plt.title("Gap vs effective redundancy proxy") plt.tight_layout() plt.savefig(red_path, dpi=180) plt.close() paths.append(red_path) dims_path = outdir / "capacity_terms_vs_width.png" plt.figure(figsize=(7, 4.5)) plt.plot(widths, [row.parameter_count for row in summaries], marker="o", label="P") plt.plot( widths, [row.parameter_count - row.fa_constraint_rank for row in summaries], marker="o", label="P - K_FA", ) plt.plot(widths, [row.task_dimension for row in summaries], marker="o", label="N*out") plt.plot(widths, [row.d_eff_mean for row in summaries], marker="o", label="d_eff") plt.plot(widths, [row.fa_burden_nats for row in summaries], marker="o", label="k_FA proxy") plt.xscale("log", base=2) plt.yscale("log") plt.xlabel("hidden width") plt.ylabel("quantity") plt.title("Capacity terms") plt.legend() plt.tight_layout() plt.savefig(dims_path, dpi=180) plt.close() paths.append(dims_path) return paths def save_distribution_plots( rows: list[RunRow], summaries: list[WidthSummaryRow], outdir: Path, ) -> list[Path]: outdir.mkdir(parents=True, exist_ok=True) paths: list[Path] = [] rng = np.random.default_rng(0) train_path = outdir / "train_mse_trajectory_distribution.png" plt.figure(figsize=(7.5, 4.8)) for run_type, color, marker, offset in [ ("bp", "tab:blue", "o", -0.035), ("fa", "tab:orange", "x", 0.035), ]: subset = [row for row in rows if row.run_type == run_type] x_values = np.array([row.width for row in subset], dtype=np.float64) jitter = np.exp(rng.normal(offset, 0.018, size=len(subset))) y_values = np.array([max(row.train_mse, 1e-10) for row in subset]) plt.scatter( x_values * jitter, y_values, s=14, alpha=0.25, color=color, marker=marker, label=f"{run_type.upper()} trajectories", ) summary_widths = np.array([row.width for row in summaries]) plt.plot( summary_widths, [max(row.bp_train_mean, 1e-10) for row in summaries], color="tab:blue", linewidth=2, label="BP mean", ) plt.plot( summary_widths, [max(row.fa_train_mean, 1e-10) for row in summaries], color="tab:orange", linewidth=2, label="FA mean", ) plt.xscale("log", base=2) plt.yscale("log") plt.xlabel("hidden width") plt.ylabel("train MSE") plt.title("Trajectory distribution with capacity thresholds") add_capacity_thresholds(summaries) plt.legend() plt.tight_layout() plt.savefig(train_path, dpi=180) plt.close() paths.append(train_path) margin_path = outdir / "train_gap_trajectory_distribution_vs_fa_margin.png" fa_rows = [row for row in rows if row.run_type == "fa"] plt.figure(figsize=(7.5, 4.8)) plt.scatter( [row.fa_capacity_margin for row in fa_rows], [row.train_gap_to_bp for row in fa_rows], s=14, alpha=0.28, color="tab:orange", label="FA trajectories", ) plt.plot( [row.fa_capacity_margin for row in summaries], [row.train_gap_mean for row in summaries], color="black", linewidth=2, marker="o", label="mean", ) plt.axvline(0.0, color="black", linestyle="--", linewidth=1, label="FA margin=0") plt.axhline(0.0, color="black", linewidth=1) plt.xlabel("P - K_FA - N*out") plt.ylabel("FA train MSE - BP train MSE") plt.title("Trajectory gap vs predicted FA margin") plt.legend() plt.tight_layout() plt.savefig(margin_path, dpi=180) plt.close() paths.append(margin_path) return paths def main() -> None: args = parse_args() config = parse_config(args) validate_config(config) if config.torch_threads > 0: torch.set_num_threads(config.torch_threads) outdir = Path(config.outdir) x_train, y_train, x_test, y_test, x_probe, _teacher = make_data(config) rows: list[RunRow] = [] task_dim = task_dimension(config) for width in config.widths: p_count = parameter_count(config, width) constraint_rank = fa_constraint_rank(config, width) bp_margin = p_count - task_dim fa_margin = p_count - constraint_rank - task_dim burden = fa_burden_nats(config, width) print( f"width={width}: P={p_count}, K_FA={constraint_rank}, " f"N*out={task_dim}, BP_margin={bp_margin}, FA_margin={fa_margin}, " f"burden={burden:.6g}", flush=True, ) for init_index in range(config.init_seeds): init_seed = 10_000 + config.init_seed_offset + init_index initial_weights = initialize_weights(config, width, init_seed) bp_weights = train( initial_weights, x_train, y_train, config.lr, config.steps, feedback=None, optimizer=config.optimizer, ) bp_train = mse(bp_weights, x_train, y_train) bp_test = mse(bp_weights, x_test, y_test) if config.skip_jacobian: d_eff, hard_rank, lam = 0.0, 0, 0.0 else: d_eff, hard_rank, lam = jacobian_effective_dimension( bp_weights, x_probe, config.jacobian_lambda_rel ) redundancy = p_count - d_eff - burden rows.append( RunRow( width=width, parameter_count=p_count, init_seed=init_seed, feedback_seed=-1, run_type="bp", train_mse=bp_train, test_mse=bp_test, bp_train_mse=bp_train, bp_test_mse=bp_test, train_gap_to_bp=0.0, test_gap_to_bp=0.0, task_dimension=task_dim, fa_constraint_rank=constraint_rank, bp_capacity_margin=bp_margin, fa_capacity_margin=fa_margin, d_eff=d_eff, hard_rank=hard_rank, jacobian_lambda=lam, fa_burden_nats=burden, redundancy_score=redundancy, ) ) for feedback_index in range(config.feedback_seeds): feedback_seed = ( 100_000 + config.feedback_seed_offset + init_index * 1000 + feedback_index ) feedback = init_feedback(config, width, feedback_seed) fa_weights = train( initial_weights, x_train, y_train, config.lr, config.steps, feedback=feedback, optimizer=config.optimizer, ) fa_train = mse(fa_weights, x_train, y_train) fa_test = mse(fa_weights, x_test, y_test) rows.append( RunRow( width=width, parameter_count=p_count, init_seed=init_seed, feedback_seed=feedback_seed, run_type="fa", train_mse=fa_train, test_mse=fa_test, bp_train_mse=bp_train, bp_test_mse=bp_test, train_gap_to_bp=fa_train - bp_train, test_gap_to_bp=fa_test - bp_test, task_dimension=task_dim, fa_constraint_rank=constraint_rank, bp_capacity_margin=bp_margin, fa_capacity_margin=fa_margin, d_eff=d_eff, hard_rank=hard_rank, jacobian_lambda=lam, fa_burden_nats=burden, redundancy_score=redundancy, ) ) partial_summaries = summarize_widths(rows) outdir.mkdir(parents=True, exist_ok=True) write_csv(outdir / "runs.partial.csv", rows) write_csv(outdir / "width_summary.partial.csv", partial_summaries) summaries = summarize_widths(rows) outdir.mkdir(parents=True, exist_ok=True) write_csv(outdir / "runs.csv", rows) write_csv(outdir / "width_summary.csv", summaries) payload = { "config": asdict(config), "width_summary": [asdict(row) for row in summaries], } (outdir / "summary.json").write_text( json.dumps(payload, indent=2, sort_keys=True) + "\n" ) plot_paths = [] if config.plot: plot_paths.extend(save_plots(summaries, outdir)) plot_paths.extend(save_distribution_plots(rows, summaries, outdir)) for row in summaries: print( f"width={row.width}: " f"BP_train={row.bp_train_mean:.6g}, FA_train={row.fa_train_mean:.6g}, " f"train_gap={row.train_gap_mean:.6g}, " f"BP_test={row.bp_test_mean:.6g}, FA_test={row.fa_test_mean:.6g}, " f"test_gap={row.gap_mean:.6g}, " f"BP_margin={row.bp_capacity_margin}, FA_margin={row.fa_capacity_margin}, " f"d_eff={row.d_eff_mean:.3f}, R={row.redundancy_score_mean:.3f}" ) print(f"runs: {outdir / 'runs.csv'}") print(f"summary: {outdir / 'width_summary.csv'}") for path in plot_paths: print(f"plot: {path}") if __name__ == "__main__": main()