summaryrefslogtreecommitdiff
path: root/scripts
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-05-28 23:19:05 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-05-28 23:19:05 -0500
commit66dc2d38ad8187e111843b911d29b9fc23283aa0 (patch)
treebbd5ffb3f6499b2c122dab7d6d4b7ca9683effc6 /scripts
parent0a3cc65f54d4ffc90b0e73f1f8713820352b27bf (diff)
Add trajectory ensemble validation
Diffstat (limited to 'scripts')
-rw-r--r--scripts/README.md25
-rwxr-xr-xscripts/trajectory_ensemble.py528
2 files changed, 553 insertions, 0 deletions
diff --git a/scripts/README.md b/scripts/README.md
index 2d7a34c..37dc68b 100644
--- a/scripts/README.md
+++ b/scripts/README.md
@@ -134,3 +134,28 @@ Outputs are written under `outputs/trajectory_mlp_fa/`:
- `trajectories.csv`
- `layer_metrics.csv`
- diagnostic plots when `--plot` is set.
+
+## Synthetic Trajectory Ensemble
+
+Run:
+
+```bash
+python scripts/trajectory_ensemble.py --architectures 16,16 24,24 32,32 48,48 24,24,24 --samples 256 --steps 200 --lr 0.02 --eval-every 20 --feedback-runs 40 --data-seed 20 --init-seed 30 --feedback-seed-start 1000 --plot
+```
+
+This runs one BP baseline and many FA feedback seeds for each architecture,
+then aggregates:
+
+- final FA/BP loss gap;
+- full and hidden-only gradient cosine;
+- initial and final \(Q_l\) means;
+- initial and final log-volume capacity cost;
+- within-architecture and pooled correlations.
+
+Outputs are written under `outputs/trajectory_ensemble/`:
+
+- `run_summary.csv`
+- `architecture_summary.csv`
+- `correlations.csv`
+- `trajectories.csv`
+- diagnostic plots when `--plot` is set.
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()