#!/usr/bin/env python3 """Run an ensemble of synthetic FA/BP MLP trajectories across architectures.""" 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 from scipy import stats import trajectory_mlp_fa as tm @dataclass(frozen=True) class EnsembleConfig: architectures: list[str] input_dim: int output_dim: int samples: int steps: int lr: float eval_every: int data_seed: int init_seed: int feedback_seed_start: int feedback_runs: int feedback_init: str feedback_scale: str noise_std: float outdir: str plot: bool @dataclass(frozen=True) class EnsembleSummaryRow: architecture: str hidden_widths: str hidden_layers: int hidden_mean_width: float parameter_count: int feedback_dimensions: str bp_final_loss: float feedback_seed: int fa_final_loss: float final_gap_to_bp: float initial_gradient_cosine: float final_gradient_cosine: float initial_hidden_gradient_cosine: float final_hidden_gradient_cosine: float initial_q_mean: float final_q_mean: float initial_capacity_nats: float final_capacity_nats: float @dataclass(frozen=True) class EnsembleTrajectoryRow: architecture: str feedback_seed: int run_type: str step: int loss: float gradient_cosine: float | None hidden_gradient_cosine: float | None q_mean: float | None q_min: float | None q_max: float | None @dataclass(frozen=True) class CorrelationRow: scope: str metric: str n: int pearson_r: float pearson_p: float spearman_r: float spearman_p: float @dataclass(frozen=True) class ArchitectureAggregateRow: architecture: str hidden_widths: str hidden_layers: int hidden_mean_width: float parameter_count: int feedback_runs: int bp_final_loss: float gap_mean: float gap_std: float gap_min: float gap_max: float final_hidden_cos_mean: float final_hidden_cos_std: float initial_capacity_mean: float final_capacity_mean: float final_q_mean_mean: float def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser( description="Run synthetic FA/BP trajectory ensembles across MLP architectures." ) parser.add_argument( "--architectures", nargs="+", default=["16,16", "24,24", "32,32", "24,24,24"], help="Hidden width specs, e.g. '16,16' '24,24,24'.", ) parser.add_argument("--input-dim", type=int, default=16) parser.add_argument("--output-dim", type=int, default=4) parser.add_argument("--samples", type=int, default=192) parser.add_argument("--steps", type=int, default=120) parser.add_argument("--lr", type=float, default=0.02) parser.add_argument("--eval-every", type=int, default=20) parser.add_argument("--data-seed", type=int, default=20) parser.add_argument("--init-seed", type=int, default=30) parser.add_argument("--feedback-seed-start", type=int, default=1_000) parser.add_argument("--feedback-runs", type=int, default=20) parser.add_argument( "--feedback-init", choices=["gaussian", "rademacher"], default="gaussian", ) parser.add_argument( "--feedback-scale", choices=["relu", "fan-in", "unit"], default="relu", ) parser.add_argument("--noise-std", type=float, default=0.01) parser.add_argument( "--outdir", type=Path, default=Path("outputs/trajectory_ensemble"), ) parser.add_argument("--plot", action="store_true") return parser.parse_args() def parse_hidden_widths(spec: str) -> list[int]: try: widths = [int(part) for part in spec.split(",") if part] except ValueError as exc: raise ValueError(f"Invalid architecture spec: {spec}") from exc if not widths or any(width < 1 for width in widths): raise ValueError(f"Architecture must contain positive widths: {spec}") return widths def make_config(args: argparse.Namespace) -> EnsembleConfig: return EnsembleConfig( architectures=args.architectures, input_dim=args.input_dim, output_dim=args.output_dim, samples=args.samples, steps=args.steps, lr=args.lr, eval_every=args.eval_every, data_seed=args.data_seed, init_seed=args.init_seed, feedback_seed_start=args.feedback_seed_start, feedback_runs=args.feedback_runs, feedback_init=args.feedback_init, feedback_scale=args.feedback_scale, noise_std=args.noise_std, outdir=str(args.outdir), plot=args.plot, ) def parameter_count(widths: list[int]) -> int: return sum(fan_in * fan_out for fan_in, fan_out in zip(widths[:-1], widths[1:])) def feedback_dimensions(widths: list[int]) -> list[int]: return [widths[index] * widths[index + 1] for index in range(1, len(widths) - 1)] def capacity_nats(q_values: list[float], dimensions: list[int]) -> float: total = 0.0 for q_value, dimension in zip(q_values, dimensions): q = min(max(q_value, 0.0), 1.0) beta_dist = stats.beta(0.5, (dimension - 1) / 2) total += float(-beta_dist.logsf(q)) return total def q_values_at_step( layer_metrics: list[tm.LayerMetricRow], step: int ) -> list[float]: values = [ row.q_alignment for row in layer_metrics if row.step == step and row.q_alignment is not None ] return [float(value) for value in values] 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 run_architecture( ensemble_config: EnsembleConfig, architecture: str, arch_index: int ) -> tuple[ list[EnsembleSummaryRow], list[EnsembleTrajectoryRow], ArchitectureAggregateRow, ]: hidden_widths = parse_hidden_widths(architecture) arch_name = "h" + "x".join(str(width) for width in hidden_widths) run_config = tm.RunConfig( input_dim=ensemble_config.input_dim, hidden_widths=hidden_widths, output_dim=ensemble_config.output_dim, samples=ensemble_config.samples, steps=ensemble_config.steps, lr=ensemble_config.lr, eval_every=ensemble_config.eval_every, data_seed=ensemble_config.data_seed + arch_index * 101, init_seed=ensemble_config.init_seed + arch_index * 101, feedback_seed_start=ensemble_config.feedback_seed_start + arch_index * 10_000, feedback_runs=ensemble_config.feedback_runs, feedback_init=ensemble_config.feedback_init, feedback_scale=ensemble_config.feedback_scale, noise_std=ensemble_config.noise_std, outdir=str(Path(ensemble_config.outdir) / arch_name), plot=False, ) tm.validate_config(run_config) widths = tm.layer_widths(run_config) dims = feedback_dimensions(widths) x, y = tm.make_synthetic_regression(run_config) initial_weights = tm.init_weights(widths, run_config.init_seed) _, bp_trajectory = tm.train_bp(initial_weights, x, y, run_config) bp_final_loss = bp_trajectory[-1].loss summary_rows: list[EnsembleSummaryRow] = [] trajectory_rows: list[EnsembleTrajectoryRow] = [] for row in bp_trajectory: trajectory_rows.append( EnsembleTrajectoryRow( architecture=arch_name, feedback_seed=-1, run_type="bp", step=row.step, loss=row.loss, gradient_cosine=row.gradient_cosine, hidden_gradient_cosine=row.hidden_gradient_cosine, q_mean=row.q_mean, q_min=row.q_min, q_max=row.q_max, ) ) for run_index in range(run_config.feedback_runs): feedback_seed = run_config.feedback_seed_start + run_index feedback = tm.init_feedback( widths, feedback_seed, run_config.feedback_init, run_config.feedback_scale ) _, trajectory, layer_metrics = tm.train_fa( initial_weights, feedback, feedback_seed, x, y, run_config ) first = trajectory[0] final = trajectory[-1] initial_q_values = q_values_at_step(layer_metrics, 0) final_q_values = q_values_at_step(layer_metrics, run_config.steps) initial_capacity = capacity_nats(initial_q_values, dims) final_capacity = capacity_nats(final_q_values, dims) summary_rows.append( EnsembleSummaryRow( architecture=arch_name, hidden_widths=",".join(str(width) for width in hidden_widths), hidden_layers=len(hidden_widths), hidden_mean_width=float(np.mean(hidden_widths)), parameter_count=parameter_count(widths), feedback_dimensions=",".join(str(dim) for dim in dims), bp_final_loss=bp_final_loss, feedback_seed=feedback_seed, fa_final_loss=final.loss, final_gap_to_bp=final.loss - bp_final_loss, initial_gradient_cosine=float(first.gradient_cosine), final_gradient_cosine=float(final.gradient_cosine), initial_hidden_gradient_cosine=float(first.hidden_gradient_cosine), final_hidden_gradient_cosine=float(final.hidden_gradient_cosine), initial_q_mean=float(first.q_mean), final_q_mean=float(final.q_mean), initial_capacity_nats=initial_capacity, final_capacity_nats=final_capacity, ) ) for row in trajectory: trajectory_rows.append( EnsembleTrajectoryRow( architecture=arch_name, feedback_seed=feedback_seed, run_type="fa", step=row.step, loss=row.loss, gradient_cosine=row.gradient_cosine, hidden_gradient_cosine=row.hidden_gradient_cosine, q_mean=row.q_mean, q_min=row.q_min, q_max=row.q_max, ) ) gaps = np.array([row.final_gap_to_bp for row in summary_rows]) final_hidden = np.array([row.final_hidden_gradient_cosine for row in summary_rows]) initial_cap = np.array([row.initial_capacity_nats for row in summary_rows]) final_cap = np.array([row.final_capacity_nats for row in summary_rows]) final_q = np.array([row.final_q_mean for row in summary_rows]) aggregate = ArchitectureAggregateRow( architecture=arch_name, hidden_widths=",".join(str(width) for width in hidden_widths), hidden_layers=len(hidden_widths), hidden_mean_width=float(np.mean(hidden_widths)), parameter_count=parameter_count(widths), feedback_runs=len(summary_rows), bp_final_loss=bp_final_loss, gap_mean=float(np.mean(gaps)), gap_std=float(np.std(gaps, ddof=1)) if len(gaps) > 1 else 0.0, gap_min=float(np.min(gaps)), gap_max=float(np.max(gaps)), final_hidden_cos_mean=float(np.mean(final_hidden)), final_hidden_cos_std=float(np.std(final_hidden, ddof=1)) if len(final_hidden) > 1 else 0.0, initial_capacity_mean=float(np.mean(initial_cap)), final_capacity_mean=float(np.mean(final_cap)), final_q_mean_mean=float(np.mean(final_q)), ) return summary_rows, trajectory_rows, aggregate def valid_correlation(x: np.ndarray, y: np.ndarray) -> bool: return len(x) >= 3 and np.std(x) > 0 and np.std(y) > 0 def correlation_row(scope: str, metric: str, x: np.ndarray, y: np.ndarray) -> CorrelationRow: mask = np.isfinite(x) & np.isfinite(y) x = x[mask] y = y[mask] if not valid_correlation(x, y): return CorrelationRow(scope, metric, len(x), math.nan, math.nan, math.nan, math.nan) pearson = stats.pearsonr(x, y) spearman = stats.spearmanr(x, y) return CorrelationRow( scope=scope, metric=metric, n=len(x), pearson_r=float(pearson.statistic), pearson_p=float(pearson.pvalue), spearman_r=float(spearman.statistic), spearman_p=float(spearman.pvalue), ) def compute_correlations(rows: list[EnsembleSummaryRow]) -> list[CorrelationRow]: correlations: list[CorrelationRow] = [] scopes = ["all", *sorted({row.architecture for row in rows})] metrics = [ "initial_gradient_cosine", "final_gradient_cosine", "initial_hidden_gradient_cosine", "final_hidden_gradient_cosine", "initial_q_mean", "final_q_mean", "initial_capacity_nats", "final_capacity_nats", ] for scope in scopes: scope_rows = rows if scope == "all" else [row for row in rows if row.architecture == scope] y = np.array([row.final_gap_to_bp for row in scope_rows], dtype=np.float64) for metric in metrics: x = np.array([getattr(row, metric) for row in scope_rows], dtype=np.float64) correlations.append(correlation_row(scope, metric, x, y)) return correlations def write_outputs( config: EnsembleConfig, summary_rows: list[EnsembleSummaryRow], trajectory_rows: list[EnsembleTrajectoryRow], aggregate_rows: list[ArchitectureAggregateRow], correlation_rows: list[CorrelationRow], outdir: Path, ) -> None: outdir.mkdir(parents=True, exist_ok=True) write_csv(outdir / "run_summary.csv", summary_rows) write_csv(outdir / "trajectories.csv", trajectory_rows) write_csv(outdir / "architecture_summary.csv", aggregate_rows) write_csv(outdir / "correlations.csv", correlation_rows) payload = { "config": asdict(config), "architecture_summary": [asdict(row) for row in aggregate_rows], "correlations": [asdict(row) for row in correlation_rows], } (outdir / "summary.json").write_text( json.dumps(payload, indent=2, sort_keys=True) + "\n" ) def save_plots( summary_rows: list[EnsembleSummaryRow], aggregate_rows: list[ArchitectureAggregateRow], outdir: Path, ) -> list[Path]: outdir.mkdir(parents=True, exist_ok=True) paths: list[Path] = [] gap_path = outdir / "gap_by_architecture.png" arch_names = [row.architecture for row in aggregate_rows] plt.figure(figsize=(7, 4.5)) plt.bar( arch_names, [row.gap_mean for row in aggregate_rows], yerr=[row.gap_std for row in aggregate_rows], capsize=4, ) plt.xlabel("architecture") plt.ylabel("FA final loss gap to BP") plt.title("Trajectory ensemble FA/BP gap") plt.tight_layout() plt.savefig(gap_path, dpi=180) plt.close() paths.append(gap_path) scatter_specs = [ ("final_hidden_gradient_cosine", "gap_vs_final_hidden_cosine.png"), ("final_q_mean", "gap_vs_final_q_mean.png"), ("final_capacity_nats", "gap_vs_final_capacity.png"), ] for metric, filename in scatter_specs: path = outdir / filename plt.figure(figsize=(6, 4.5)) for arch in arch_names: arch_rows = [row for row in summary_rows if row.architecture == arch] plt.scatter( [getattr(row, metric) for row in arch_rows], [row.final_gap_to_bp for row in arch_rows], label=arch, alpha=0.8, ) plt.xlabel(metric) plt.ylabel("FA final loss gap to BP") plt.title(f"Gap vs {metric}") plt.legend() plt.tight_layout() plt.savefig(path, dpi=180) plt.close() paths.append(path) return paths def main() -> None: args = parse_args() config = make_config(args) outdir = Path(config.outdir) all_summary_rows: list[EnsembleSummaryRow] = [] all_trajectory_rows: list[EnsembleTrajectoryRow] = [] aggregate_rows: list[ArchitectureAggregateRow] = [] for arch_index, architecture in enumerate(config.architectures): summary_rows, trajectory_rows, aggregate = run_architecture( config, architecture, arch_index ) all_summary_rows.extend(summary_rows) all_trajectory_rows.extend(trajectory_rows) aggregate_rows.append(aggregate) print( f"{aggregate.architecture}: " f"gap_mean={aggregate.gap_mean:.6g}, " f"gap_std={aggregate.gap_std:.6g}, " f"final_hidden_cos_mean={aggregate.final_hidden_cos_mean:.6g}, " f"final_capacity_mean={aggregate.final_capacity_mean:.6g}" ) correlation_rows = compute_correlations(all_summary_rows) write_outputs( config, all_summary_rows, all_trajectory_rows, aggregate_rows, correlation_rows, outdir, ) plot_paths = save_plots(all_summary_rows, aggregate_rows, outdir) if config.plot else [] all_gaps = np.array([row.final_gap_to_bp for row in all_summary_rows]) print(f"total_fa_runs: {len(all_summary_rows)}") print( "all_gap: " f"mean={np.mean(all_gaps):.8g}, " f"std={np.std(all_gaps, ddof=1):.8g}, " f"min={np.min(all_gaps):.8g}, " f"max={np.max(all_gaps):.8g}" ) print(f"run_summary: {outdir / 'run_summary.csv'}") print(f"architecture_summary: {outdir / 'architecture_summary.csv'}") print(f"correlations: {outdir / 'correlations.csv'}") for path in plot_paths: print(f"plot: {path}") if __name__ == "__main__": main()