diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-05-28 23:19:05 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-05-28 23:19:05 -0500 |
| commit | 66dc2d38ad8187e111843b911d29b9fc23283aa0 (patch) | |
| tree | bbd5ffb3f6499b2c122dab7d6d4b7ca9683effc6 /scripts/trajectory_ensemble.py | |
| parent | 0a3cc65f54d4ffc90b0e73f1f8713820352b27bf (diff) | |
Add trajectory ensemble validation
Diffstat (limited to 'scripts/trajectory_ensemble.py')
| -rwxr-xr-x | scripts/trajectory_ensemble.py | 528 |
1 files changed, 528 insertions, 0 deletions
diff --git a/scripts/trajectory_ensemble.py b/scripts/trajectory_ensemble.py new file mode 100755 index 0000000..e14600c --- /dev/null +++ b/scripts/trajectory_ensemble.py @@ -0,0 +1,528 @@ +#!/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() |
