diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-05-29 12:52:29 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-05-29 12:52:29 -0500 |
| commit | 385a3c0fc97c517bc7361bbb86ca6a134ca198c0 (patch) | |
| tree | 53ffa1de9f6a38ee2e08587d935b919f9058fbc1 /scripts/trajectory_gap_distribution.py | |
| parent | 96e201556ac94057a4a5c8864cf422ad43d72d58 (diff) | |
Add trajectory gap distribution bridge
Diffstat (limited to 'scripts/trajectory_gap_distribution.py')
| -rwxr-xr-x | scripts/trajectory_gap_distribution.py | 525 |
1 files changed, 525 insertions, 0 deletions
diff --git a/scripts/trajectory_gap_distribution.py b/scripts/trajectory_gap_distribution.py new file mode 100755 index 0000000..6551acd --- /dev/null +++ b/scripts/trajectory_gap_distribution.py @@ -0,0 +1,525 @@ +#!/usr/bin/env python3 +"""Match FA/BP trajectory gap distributions against a BP-path bridge predictor. + +This is a first distributional bridge for realized FA loss gaps. For each +feedback seed B, it compares: + +1. Empirical gap: train the actual FA trajectory and measure + L(theta_T^FA) - L(theta_T^BP). +2. Predicted bridge gap: integrate the FA/BP gradient mismatch along the fixed + BP trajectory, + + delta_T(B) = -eta sum_t [g_FA(theta_t^BP; B) - g_BP(theta_t^BP)], + + then evaluate L(theta_T^BP + delta_T(B)) - L(theta_T^BP). + +This is not a final theorem. It is the first test of whether a trajectory-level +bridge can predict a full gap distribution, not merely correlations. +""" + +from __future__ import annotations + +import argparse +import csv +import json +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 RunConfig: + input_dim: int + hidden_widths: list[int] + output_dim: int + samples: int + steps: int + lr: float + 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 GapRow: + feedback_seed: int + empirical_gap: float + predicted_bridge_gap: float + empirical_final_loss: float + predicted_bridge_loss: float + bp_final_loss: float + initial_hidden_gradient_cosine: float + final_hidden_gradient_cosine: float + initial_q_mean: float + final_q_mean: float + bridge_delta_norm: float + empirical_delta_norm: float + + +@dataclass(frozen=True) +class MatchSummary: + runs: int + bp_final_loss: float + empirical_mean: float + predicted_mean: float + empirical_std: float + predicted_std: float + empirical_q01: float + predicted_q01: float + empirical_q50: float + predicted_q50: float + empirical_q99: float + predicted_q99: float + ks_2sample_statistic: float + ks_2sample_pvalue: float + standardized_ks_statistic: float + standardized_ks_pvalue: float + moment_matched_ks_statistic: float + moment_matched_ks_pvalue: float + moment_matched_wasserstein_distance: float + wasserstein_distance: float + pearson_r: float + pearson_p: float + spearman_r: float + spearman_p: float + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="FA/BP trajectory gap distribution matching." + ) + parser.add_argument("--input-dim", type=int, default=16) + parser.add_argument("--hidden-widths", type=int, nargs="+", default=[16, 16]) + parser.add_argument("--output-dim", type=int, default=4) + parser.add_argument("--samples", type=int, default=256) + parser.add_argument("--steps", type=int, default=120) + parser.add_argument("--lr", type=float, default=0.02) + 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=5_000) + parser.add_argument("--feedback-runs", type=int, default=500) + 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_gap_distribution"), + ) + parser.add_argument("--plot", action="store_true") + return parser.parse_args() + + +def parse_config(args: argparse.Namespace) -> RunConfig: + return RunConfig( + input_dim=args.input_dim, + hidden_widths=args.hidden_widths, + output_dim=args.output_dim, + samples=args.samples, + steps=args.steps, + lr=args.lr, + 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 tm_config(config: RunConfig) -> tm.RunConfig: + return tm.RunConfig( + input_dim=config.input_dim, + hidden_widths=config.hidden_widths, + output_dim=config.output_dim, + samples=config.samples, + steps=config.steps, + lr=config.lr, + eval_every=config.steps, + data_seed=config.data_seed, + init_seed=config.init_seed, + feedback_seed_start=config.feedback_seed_start, + feedback_runs=config.feedback_runs, + feedback_init=config.feedback_init, + feedback_scale=config.feedback_scale, + noise_std=config.noise_std, + outdir=config.outdir, + plot=False, + ) + + +def validate_config(config: RunConfig) -> None: + run_config = tm_config(config) + tm.validate_config(run_config) + + +def zeros_like(weights: list[tm.Array]) -> list[tm.Array]: + return [np.zeros_like(weight) for weight in weights] + + +def copy_weights(weights: list[tm.Array]) -> list[tm.Array]: + return [weight.copy() for weight in weights] + + +def add_scaled(target: list[tm.Array], source: list[tm.Array], scale: float) -> None: + for target_array, source_array in zip(target, source): + target_array += scale * source_array + + +def add_weights(a: list[tm.Array], b: list[tm.Array]) -> list[tm.Array]: + return [left + right for left, right in zip(a, b)] + + +def weights_delta(a: list[tm.Array], b: list[tm.Array]) -> list[tm.Array]: + return [left - right for left, right in zip(a, b)] + + +def norm_weights(weights: list[tm.Array]) -> float: + return float(np.linalg.norm(tm.flatten(weights))) + + +def train_bp_path( + initial_weights: list[tm.Array], x: tm.Array, y: tm.Array, lr: float, steps: int +) -> tuple[list[list[tm.Array]], list[list[tm.Array]], list[tm.Array], float]: + weights = copy_weights(initial_weights) + path: list[list[tm.Array]] = [] + bp_grads: list[list[tm.Array]] = [] + + for _step in range(steps): + path.append(copy_weights(weights)) + grads, _ = tm.gradients(weights, x, y, feedback=None) + bp_grads.append(copy_weights(grads)) + tm.sgd_step(weights, grads, lr) + + bp_final_loss = tm.mse_loss(tm.predict(weights, x), y) + return path, bp_grads, weights, bp_final_loss + + +def bridge_delta( + bp_path: list[list[tm.Array]], + bp_grads: list[list[tm.Array]], + feedback: list[tm.Array], + x: tm.Array, + y: tm.Array, + lr: float, +) -> list[tm.Array]: + delta = zeros_like(bp_path[0]) + for weights, bp_grad in zip(bp_path, bp_grads): + fa_grad, _ = tm.gradients(weights, x, y, feedback=feedback) + mismatch = [fa - bp for fa, bp in zip(fa_grad, bp_grad)] + add_scaled(delta, mismatch, -lr) + return delta + + +def run_one_feedback( + config: RunConfig, + run_config: tm.RunConfig, + initial_weights: list[tm.Array], + bp_path: list[list[tm.Array]], + bp_grads: list[list[tm.Array]], + bp_final_weights: list[tm.Array], + bp_final_loss: float, + x: tm.Array, + y: tm.Array, + feedback_seed: int, +) -> GapRow: + widths = tm.layer_widths(run_config) + feedback = tm.init_feedback( + widths, feedback_seed, config.feedback_init, config.feedback_scale + ) + + delta = bridge_delta(bp_path, bp_grads, feedback, x, y, config.lr) + bridge_weights = add_weights(bp_final_weights, delta) + bridge_loss = tm.mse_loss(tm.predict(bridge_weights, x), y) + predicted_gap = bridge_loss - bp_final_loss + + fa_final_weights, trajectory, _layer_metrics = tm.train_fa( + initial_weights, feedback, feedback_seed, x, y, run_config + ) + empirical_final_loss = trajectory[-1].loss + empirical_gap = empirical_final_loss - bp_final_loss + + initial = trajectory[0] + final = trajectory[-1] + empirical_delta = weights_delta(fa_final_weights, bp_final_weights) + + return GapRow( + feedback_seed=feedback_seed, + empirical_gap=empirical_gap, + predicted_bridge_gap=predicted_gap, + empirical_final_loss=empirical_final_loss, + predicted_bridge_loss=bridge_loss, + bp_final_loss=bp_final_loss, + initial_hidden_gradient_cosine=float(initial.hidden_gradient_cosine), + final_hidden_gradient_cosine=float(final.hidden_gradient_cosine), + initial_q_mean=float(initial.q_mean), + final_q_mean=float(final.q_mean), + bridge_delta_norm=norm_weights(delta), + empirical_delta_norm=norm_weights(empirical_delta), + ) + + +def summarize(rows: list[GapRow], bp_final_loss: float) -> MatchSummary: + empirical = np.array([row.empirical_gap for row in rows], dtype=np.float64) + predicted = np.array([row.predicted_bridge_gap for row in rows], dtype=np.float64) + ks = stats.ks_2samp(predicted, empirical) + predicted_std = np.std(predicted, ddof=1) + empirical_std = np.std(empirical, ddof=1) + predicted_z = (predicted - np.mean(predicted)) / predicted_std + empirical_z = (empirical - np.mean(empirical)) / empirical_std + standardized_ks = stats.ks_2samp(predicted_z, empirical_z) + moment_matched = np.mean(empirical) + predicted_z * empirical_std + moment_matched_ks = stats.ks_2samp(moment_matched, empirical) + pearson = stats.pearsonr(predicted, empirical) + spearman = stats.spearmanr(predicted, empirical) + return MatchSummary( + runs=len(rows), + bp_final_loss=bp_final_loss, + empirical_mean=float(np.mean(empirical)), + predicted_mean=float(np.mean(predicted)), + empirical_std=float(empirical_std), + predicted_std=float(predicted_std), + empirical_q01=float(np.quantile(empirical, 0.01)), + predicted_q01=float(np.quantile(predicted, 0.01)), + empirical_q50=float(np.quantile(empirical, 0.50)), + predicted_q50=float(np.quantile(predicted, 0.50)), + empirical_q99=float(np.quantile(empirical, 0.99)), + predicted_q99=float(np.quantile(predicted, 0.99)), + ks_2sample_statistic=float(ks.statistic), + ks_2sample_pvalue=float(ks.pvalue), + standardized_ks_statistic=float(standardized_ks.statistic), + standardized_ks_pvalue=float(standardized_ks.pvalue), + moment_matched_ks_statistic=float(moment_matched_ks.statistic), + moment_matched_ks_pvalue=float(moment_matched_ks.pvalue), + moment_matched_wasserstein_distance=float( + stats.wasserstein_distance(moment_matched, empirical) + ), + wasserstein_distance=float(stats.wasserstein_distance(predicted, empirical)), + pearson_r=float(pearson.statistic), + pearson_p=float(pearson.pvalue), + spearman_r=float(spearman.statistic), + spearman_p=float(spearman.pvalue), + ) + + +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 write_outputs( + config: RunConfig, rows: list[GapRow], summary: MatchSummary, outdir: Path +) -> None: + outdir.mkdir(parents=True, exist_ok=True) + write_csv(outdir / "gap_distribution.csv", rows) + write_csv(outdir / "summary.csv", [summary]) + payload = { + "config": asdict(config), + "summary": asdict(summary), + "runs": [asdict(row) for row in rows], + } + (outdir / "summary.json").write_text( + json.dumps(payload, indent=2, sort_keys=True) + "\n" + ) + + +def save_plots(rows: list[GapRow], outdir: Path) -> list[Path]: + outdir.mkdir(parents=True, exist_ok=True) + paths: list[Path] = [] + empirical = np.array([row.empirical_gap for row in rows], dtype=np.float64) + predicted = np.array([row.predicted_bridge_gap for row in rows], dtype=np.float64) + predicted_z = (predicted - np.mean(predicted)) / np.std(predicted, ddof=1) + empirical_z = (empirical - np.mean(empirical)) / np.std(empirical, ddof=1) + moment_matched = np.mean(empirical) + predicted_z * np.std(empirical, ddof=1) + + hist_path = outdir / "gap_histogram_overlay.png" + bins = np.histogram_bin_edges(np.concatenate([empirical, predicted]), bins=60) + plt.figure(figsize=(7, 4.5)) + plt.hist(predicted, bins=bins, density=True, alpha=0.45, label="bridge predicted") + plt.hist(empirical, bins=bins, density=True, alpha=0.45, label="actual FA") + plt.xlabel("final loss gap to BP") + plt.ylabel("density") + plt.title("Predicted vs empirical FA/BP gap distribution") + plt.legend() + plt.tight_layout() + plt.savefig(hist_path, dpi=180) + plt.close() + paths.append(hist_path) + + matched_hist_path = outdir / "gap_histogram_moment_matched.png" + bins = np.histogram_bin_edges(np.concatenate([empirical, moment_matched]), bins=60) + plt.figure(figsize=(7, 4.5)) + plt.hist(moment_matched, bins=bins, density=True, alpha=0.45, label="bridge moment-matched") + plt.hist(empirical, bins=bins, density=True, alpha=0.45, label="actual FA") + plt.xlabel("final loss gap to BP") + plt.ylabel("density") + plt.title("Moment-matched bridge vs empirical gap distribution") + plt.legend() + plt.tight_layout() + plt.savefig(matched_hist_path, dpi=180) + plt.close() + paths.append(matched_hist_path) + + cdf_path = outdir / "gap_cdf_overlay.png" + plt.figure(figsize=(7, 4.5)) + for values, label in [(predicted, "bridge predicted"), (empirical, "actual FA")]: + sorted_values = np.sort(values) + probs = (np.arange(1, len(values) + 1) - 0.5) / len(values) + plt.plot(sorted_values, probs, label=label) + plt.xlabel("final loss gap to BP") + plt.ylabel("CDF") + plt.title("Gap distribution CDF") + plt.legend() + plt.tight_layout() + plt.savefig(cdf_path, dpi=180) + plt.close() + paths.append(cdf_path) + + qq_path = outdir / "gap_qq_plot.png" + probs = (np.arange(1, len(empirical) + 1) - 0.5) / len(empirical) + pred_q = np.quantile(predicted, probs) + emp_q = np.quantile(empirical, probs) + min_value = float(min(pred_q[0], emp_q[0])) + max_value = float(max(pred_q[-1], emp_q[-1])) + plt.figure(figsize=(5, 5)) + plt.scatter(pred_q, emp_q, s=8, alpha=0.35) + plt.plot([min_value, max_value], [min_value, max_value], color="black", linewidth=1) + plt.xlabel("bridge predicted quantile") + plt.ylabel("actual FA quantile") + plt.title("Gap distribution Q-Q") + plt.tight_layout() + plt.savefig(qq_path, dpi=180) + plt.close() + paths.append(qq_path) + + standardized_qq_path = outdir / "gap_standardized_qq_plot.png" + pred_q = np.quantile(predicted_z, probs) + emp_q = np.quantile(empirical_z, probs) + min_value = float(min(pred_q[0], emp_q[0])) + max_value = float(max(pred_q[-1], emp_q[-1])) + plt.figure(figsize=(5, 5)) + plt.scatter(pred_q, emp_q, s=8, alpha=0.35) + plt.plot([min_value, max_value], [min_value, max_value], color="black", linewidth=1) + plt.xlabel("standardized bridge quantile") + plt.ylabel("standardized actual FA quantile") + plt.title("Standardized gap distribution Q-Q") + plt.tight_layout() + plt.savefig(standardized_qq_path, dpi=180) + plt.close() + paths.append(standardized_qq_path) + + scatter_path = outdir / "paired_gap_scatter.png" + plt.figure(figsize=(5, 5)) + plt.scatter(predicted, empirical, s=12, alpha=0.45) + min_value = float(min(np.min(predicted), np.min(empirical))) + max_value = float(max(np.max(predicted), np.max(empirical))) + plt.plot([min_value, max_value], [min_value, max_value], color="black", linewidth=1) + plt.xlabel("bridge predicted gap") + plt.ylabel("actual FA gap") + plt.title("Paired feedback seeds") + plt.tight_layout() + plt.savefig(scatter_path, dpi=180) + plt.close() + paths.append(scatter_path) + + return paths + + +def main() -> None: + args = parse_args() + config = parse_config(args) + validate_config(config) + run_config = tm_config(config) + widths = tm.layer_widths(run_config) + x, y = tm.make_synthetic_regression(run_config) + initial_weights = tm.init_weights(widths, config.init_seed) + bp_path, bp_grads, bp_final_weights, bp_final_loss = train_bp_path( + initial_weights, x, y, config.lr, config.steps + ) + + rows: list[GapRow] = [] + for run_index in range(config.feedback_runs): + feedback_seed = config.feedback_seed_start + run_index + row = run_one_feedback( + config, + run_config, + initial_weights, + bp_path, + bp_grads, + bp_final_weights, + bp_final_loss, + x, + y, + feedback_seed, + ) + rows.append(row) + if (run_index + 1) % max(1, config.feedback_runs // 10) == 0: + print( + f"completed {run_index + 1}/{config.feedback_runs}: " + f"emp_gap={row.empirical_gap:.6g}, pred_gap={row.predicted_bridge_gap:.6g}" + ) + + summary = summarize(rows, bp_final_loss) + outdir = Path(config.outdir) + write_outputs(config, rows, summary, outdir) + plot_paths = save_plots(rows, outdir) if config.plot else [] + + print(f"widths: {widths}") + print(f"bp_final_loss: {bp_final_loss:.8g}") + print( + "empirical_gap: " + f"mean={summary.empirical_mean:.8g}, std={summary.empirical_std:.8g}" + ) + print( + "predicted_bridge_gap: " + f"mean={summary.predicted_mean:.8g}, std={summary.predicted_std:.8g}" + ) + print( + "distribution_match: " + f"KS={summary.ks_2sample_statistic:.8g}, " + f"W1={summary.wasserstein_distance:.8g}, " + f"Spearman={summary.spearman_r:.8g}" + ) + print( + "shape_match_after_standardization: " + f"KS={summary.standardized_ks_statistic:.8g}, " + f"p={summary.standardized_ks_pvalue:.8g}, " + f"moment_matched_W1={summary.moment_matched_wasserstein_distance:.8g}" + ) + print(f"gap_distribution: {outdir / 'gap_distribution.csv'}") + print(f"summary: {outdir / 'summary.csv'}") + for path in plot_paths: + print(f"plot: {path}") + + +if __name__ == "__main__": + main() |
