diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-27 17:47:13 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-27 17:47:13 -0500 |
| commit | e65dd2e2460b1da48c83a930304cf9b269fc4447 (patch) | |
| tree | 94f583b0fb5e9989bd6c68f7f0e950c6ef91756d /scripts/aaai_depth_experiments.py | |
| parent | 82a49011de15287583d8cfec11ac5cca7efee747 (diff) | |
Add AAAI depth experiments and diagnostic figures
Diffstat (limited to 'scripts/aaai_depth_experiments.py')
| -rw-r--r-- | scripts/aaai_depth_experiments.py | 633 |
1 files changed, 633 insertions, 0 deletions
diff --git a/scripts/aaai_depth_experiments.py b/scripts/aaai_depth_experiments.py new file mode 100644 index 0000000..bb52f85 --- /dev/null +++ b/scripts/aaai_depth_experiments.py @@ -0,0 +1,633 @@ +#!/usr/bin/env python3 +"""AAAI depth experiments for fixed random feedback. + +Part ``init`` validates the exact expected initial first-order deficit for both +FA and DFA at hidden depths 1, 2, 3, 4, and 6. Part ``finite`` measures the +FA/BP finite-time loss gap and evaluates the frozen-initialization, linear +operator-velocity, and early-retangent predictors over depth and width. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +import sys +import time +from dataclasses import asdict, dataclass +from pathlib import Path + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import torch + +SCRIPT_DIR = Path(__file__).resolve().parent +if str(SCRIPT_DIR) not in sys.path: + sys.path.insert(0, str(SCRIPT_DIR)) + +import feedback_rules as fr # noqa: E402 +import real_data_validation as rdv # noqa: E402 + + +@dataclass(frozen=True) +class InitRow: + rule: str + depth: int + width: int + init_seed: int + feedback_draws: int + bp_speed: float + output_speed: float + predicted_deficit: float + empirical_deficit: float + empirical_stderr: float + empirical_std: float + calibration_error: float + + +@dataclass(frozen=True) +class FiniteRow: + depth: int + width: int + train_samples: int + init_seed: int + feedback_seed: int + lr: float + horizon: int + early_step: int + initial_deficit_prediction: float + empirical_bp_loss: float + empirical_fa_loss: float + empirical_gap: float + fixed_gap: float + velocity_gap: float + retangent_gap: float + fixed_error: float + velocity_error: float + retangent_error: float + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--part", choices=["init", "finite", "both", "summarize"], default="both" + ) + parser.add_argument("--depths", type=int, nargs="+", default=[1, 2, 3, 4, 6]) + parser.add_argument("--width", type=int, default=64, help="fixed width for initialization") + parser.add_argument("--finite-widths", type=int, nargs="+", default=[32, 64, 96]) + parser.add_argument("--input-dim", type=int, default=16) + parser.add_argument("--output-dim", type=int, default=4) + parser.add_argument("--train-samples", type=int, default=64) + parser.add_argument("--data-seed", type=int, default=2027) + parser.add_argument("--init-seeds", type=int, default=6) + parser.add_argument("--feedback-draws", type=int, default=256) + parser.add_argument("--finite-init-seeds", type=int, default=3) + parser.add_argument("--finite-feedback-seeds", type=int, default=4) + parser.add_argument("--lr", type=float, default=1e-3) + parser.add_argument("--horizon", type=int, default=50) + parser.add_argument("--early-step", type=int, default=20) + parser.add_argument("--feedback-scale", choices=["relu", "fan-in", "unit"], default="relu") + parser.add_argument("--torch-threads", type=int, default=16) + parser.add_argument("--outdir", type=Path, default=Path("outputs/aaai_depth_experiments")) + return parser.parse_args() + + +def make_data(args: argparse.Namespace) -> tuple[torch.Tensor, torch.Tensor]: + generator = torch.Generator().manual_seed(args.data_seed) + x = torch.randn(args.train_samples, args.input_dim, generator=generator, dtype=torch.float64) + y = torch.randn(args.train_samples, args.output_dim, generator=generator, dtype=torch.float64) + return x, y + + +def write_dataclass_csv(path: Path, rows: list[object]) -> None: + if not rows: + return + path.parent.mkdir(parents=True, exist_ok=True) + first = asdict(rows[0]) # type: ignore[arg-type] + with path.open("w", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=list(first)) + writer.writeheader() + for row in rows: + writer.writerow(asdict(row)) # type: ignore[arg-type] + + +def run_initialization(args: argparse.Namespace, x: torch.Tensor, y: torch.Tensor) -> list[InitRow]: + rows: list[InitRow] = [] + for depth in args.depths: + dims = [args.input_dim, *([args.width] * depth), args.output_dim] + for init_index in range(args.init_seeds): + init_seed = 10_000 + init_index + weights = fr.initialize_mlp(dims, seed=init_seed) + bp_grads = fr.gradients(weights, x, y) + bp_speed = fr.squared_norm(bp_grads) + output_speed = fr.squared_norm([bp_grads[-1]]) + predicted = 1.0 - output_speed / bp_speed + for rule_index, rule in enumerate(("fa", "dfa")): + deficits: list[float] = [] + for draw in range(args.feedback_draws): + feedback_seed = 100_000_000 + 10_000 * depth + 1_000 * init_index + 2 * draw + rule_index + feedback = fr.init_feedback( + dims, + seed=feedback_seed, + rule=rule, + mode=args.feedback_scale, + ) + rule_grads = fr.gradients(weights, x, y, rule=rule, feedback=feedback) + deficits.append(1.0 - fr.inner_product(bp_grads, rule_grads) / bp_speed) + values = np.asarray(deficits) + empirical = float(values.mean()) + row = InitRow( + rule=rule.upper(), + depth=depth, + width=args.width, + init_seed=init_seed, + feedback_draws=args.feedback_draws, + bp_speed=bp_speed, + output_speed=output_speed, + predicted_deficit=predicted, + empirical_deficit=empirical, + empirical_stderr=float(values.std(ddof=1) / math.sqrt(len(values))), + empirical_std=float(values.std(ddof=1)), + calibration_error=empirical - predicted, + ) + rows.append(row) + print( + f"[init] {rule.upper()} depth={depth} init={init_index}: " + f"prediction={predicted:.5f}, measured={empirical:.5f} " + f"+/- {row.empirical_stderr:.5f}", + flush=True, + ) + return rows + + +def train_with_snapshots( + weights0: list[torch.Tensor], + x: torch.Tensor, + y: torch.Tensor, + lr: float, + horizon: int, + feedback: list[torch.Tensor] | None, + early_step: int, +) -> tuple[list[torch.Tensor], list[torch.Tensor]]: + final, snapshots = rdv.train_with_snapshots( + weights0, + x, + y, + lr, + horizon, + feedback, + {early_step}, + ) + return final, snapshots[early_step] + + +def retangent_loss( + kernel_s: torch.Tensor, + residual_s: torch.Tensor, + lr: float, + samples: int, + horizon: int, + early_step: int, +) -> float: + return rdv.rollout_loss( + kernel_s, + None, + residual_s, + lr, + samples, + horizon - early_step, + early_step, + ) + + +def run_finite_time(args: argparse.Namespace, x: torch.Tensor, y: torch.Tensor) -> list[FiniteRow]: + rows: list[FiniteRow] = [] + n = x.shape[0] + for depth in args.depths: + for width in args.finite_widths: + dims = [args.input_dim, *([width] * depth), args.output_dim] + for init_index in range(args.finite_init_seeds): + block_start = time.time() + init_seed = 20_000 + init_index + weights0 = fr.initialize_mlp(dims, seed=init_seed) + r0 = (fr.predict(weights0, x) - y).reshape(-1) + bp_grads = fr.gradients(weights0, x, y) + initial_prediction = 1.0 - fr.squared_norm([bp_grads[-1]]) / fr.squared_norm(bp_grads) + + k_bp0 = rdv.tangent_kernel(weights0, x, None) + bp_final, bp_early = train_with_snapshots( + weights0, x, y, args.lr, args.horizon, None, args.early_step + ) + k_bps = rdv.tangent_kernel(bp_early, x, None) + bp_emp = fr.mse(bp_final, x, y) + bp_fixed = rdv.rollout_loss( + k_bp0, None, r0, args.lr, n, args.horizon, args.early_step + ) + bp_velocity = rdv.rollout_loss( + k_bp0, + k_bps - k_bp0, + r0, + args.lr, + n, + args.horizon, + args.early_step, + ) + bp_residual_s = (fr.predict(bp_early, x) - y).reshape(-1) + bp_retangent = retangent_loss( + k_bps, bp_residual_s, args.lr, n, args.horizon, args.early_step + ) + + for feedback_index in range(args.finite_feedback_seeds): + feedback_seed = 200_000 + 1_000 * init_index + feedback_index + feedback = fr.init_feedback( + dims, + seed=feedback_seed, + rule="fa", + mode=args.feedback_scale, + ) + k_fa0 = rdv.tangent_kernel(weights0, x, feedback) + fa_final, fa_early = train_with_snapshots( + weights0, + x, + y, + args.lr, + args.horizon, + feedback, + args.early_step, + ) + k_fas = rdv.tangent_kernel(fa_early, x, feedback) + fa_emp = fr.mse(fa_final, x, y) + fa_fixed = rdv.rollout_loss( + k_fa0, None, r0, args.lr, n, args.horizon, args.early_step + ) + fa_velocity = rdv.rollout_loss( + k_fa0, + k_fas - k_fa0, + r0, + args.lr, + n, + args.horizon, + args.early_step, + ) + fa_residual_s = (fr.predict(fa_early, x) - y).reshape(-1) + fa_retangent = retangent_loss( + k_fas, fa_residual_s, args.lr, n, args.horizon, args.early_step + ) + + empirical_gap = fa_emp - bp_emp + fixed_gap = fa_fixed - bp_fixed + velocity_gap = fa_velocity - bp_velocity + retangent_gap = fa_retangent - bp_retangent + rows.append( + FiniteRow( + depth=depth, + width=width, + train_samples=n, + init_seed=init_seed, + feedback_seed=feedback_seed, + lr=args.lr, + horizon=args.horizon, + early_step=args.early_step, + initial_deficit_prediction=initial_prediction, + empirical_bp_loss=bp_emp, + empirical_fa_loss=fa_emp, + empirical_gap=empirical_gap, + fixed_gap=fixed_gap, + velocity_gap=velocity_gap, + retangent_gap=retangent_gap, + fixed_error=fixed_gap - empirical_gap, + velocity_error=velocity_gap - empirical_gap, + retangent_error=retangent_gap - empirical_gap, + ) + ) + print( + f"[finite] depth={depth} width={width} init={init_index}: " + f"{args.finite_feedback_seeds} feedbacks in {time.time() - block_start:.1f}s", + flush=True, + ) + return rows + + +def safe_corr(left: np.ndarray, right: np.ndarray) -> float: + if len(left) < 2 or float(np.std(left) * np.std(right)) == 0.0: + return math.nan + return float(np.corrcoef(left, right)[0, 1]) + + +def summarize_finite(rows: list[FiniteRow]) -> list[dict[str, float | int | str]]: + metrics: list[dict[str, float | int | str]] = [] + groups = sorted({(row.depth, row.width) for row in rows}) + for depth, width in groups: + subset = [row for row in rows if row.depth == depth and row.width == width] + empirical = np.asarray([row.empirical_gap for row in subset]) + for predictor, field in ( + ("frozen K(0)", "fixed_gap"), + ("linear velocity", "velocity_gap"), + ("early retangent", "retangent_gap"), + ): + predicted = np.asarray([getattr(row, field) for row in subset]) + metrics.append( + { + "depth": depth, + "width": width, + "predictor": predictor, + "rows": len(subset), + "empirical_gap_mean": float(empirical.mean()), + "empirical_gap_std": float(empirical.std(ddof=1)), + "mae": float(np.mean(np.abs(predicted - empirical))), + "bias": float(np.mean(predicted - empirical)), + "corr": safe_corr(predicted, empirical), + } + ) + return metrics + + +def write_dict_csv(path: Path, rows: list[dict[str, object]]) -> None: + if not rows: + return + with path.open("w", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=list(rows[0])) + writer.writeheader() + writer.writerows(rows) + + +def plot_initial(rows: list[InitRow], outdir: Path) -> None: + fig, ax = plt.subplots(figsize=(5.7, 5.2), dpi=180) + colors = plt.cm.viridis(np.linspace(0.1, 0.9, len(sorted({row.depth for row in rows})))) + for color, depth in zip(colors, sorted({row.depth for row in rows})): + for rule, marker in (("FA", "o"), ("DFA", "s")): + subset = [row for row in rows if row.depth == depth and row.rule == rule] + ax.errorbar( + [row.predicted_deficit for row in subset], + [row.empirical_deficit for row in subset], + yerr=[2 * row.empirical_stderr for row in subset], + fmt=marker, + ms=4.5, + color=color, + alpha=0.78, + capsize=1.5, + lw=0.8, + label=f"depth {depth}, {rule}", + ) + values = [value for row in rows for value in (row.predicted_deficit, row.empirical_deficit)] + lo, hi = min(values), max(values) + pad = 0.04 * (hi - lo + 1e-12) + ax.plot([lo - pad, hi + pad], [lo - pad, hi + pad], color="black", lw=1) + ax.set_xlabel("exact expected initial deficit") + ax.set_ylabel("measured mean initial deficit") + ax.set_title("FA and DFA share the exact initialization cost") + ax.grid(alpha=0.16) + ax.legend(fontsize=6.5, ncols=2) + fig.tight_layout() + fig.savefig(outdir / "initial_deficit_calibration.png", bbox_inches="tight") + plt.close(fig) + + +def plot_finite(rows: list[FiniteRow], metrics: list[dict[str, object]], outdir: Path) -> None: + fig, ax = plt.subplots(figsize=(6.4, 4.7), dpi=180) + for width, marker in zip(sorted({row.width for row in rows}), ("o", "s", "^", "D")): + depths = sorted({row.depth for row in rows if row.width == width}) + means = [np.mean([row.empirical_gap for row in rows if row.width == width and row.depth == d]) for d in depths] + sems = [] + for depth in depths: + cell = [row for row in rows if row.width == width and row.depth == depth] + init_means = [ + np.mean([row.empirical_gap for row in cell if row.init_seed == init_seed]) + for init_seed in sorted({row.init_seed for row in cell}) + ] + sems.append(float(np.std(init_means, ddof=1) / math.sqrt(len(init_means)))) + ax.errorbar(depths, means, yerr=sems, marker=marker, capsize=3, label=f"width {width}") + ax.set_xlabel("hidden-layer depth") + ax.set_ylabel("FA loss - BP loss at fixed horizon") + ax.set_title("Finite-time optimization cost across depth and width") + ax.grid(alpha=0.16) + ax.legend() + fig.tight_layout() + fig.savefig(outdir / "finite_gap_by_depth_width.png", bbox_inches="tight") + plt.close(fig) + + fig, ax = plt.subplots(figsize=(7.0, 4.8), dpi=180) + predictors = ["frozen K(0)", "linear velocity", "early retangent"] + for predictor, marker in zip(predictors, ("^", "o", "s")): + subset = [row for row in metrics if row["predictor"] == predictor and row["width"] == 64] + ax.plot( + [int(row["depth"]) for row in subset], + [float(row["mae"]) for row in subset], + marker=marker, + label=predictor, + ) + ax.set_xlabel("hidden-layer depth") + ax.set_ylabel("gap prediction MAE (width 64)") + ax.set_title("Operator drift, not initialization mismatch, controls later error") + ax.set_yscale("log") + ax.grid(alpha=0.16, which="both") + ax.legend() + fig.tight_layout() + fig.savefig(outdir / "predictor_mae_by_depth.png", bbox_inches="tight") + plt.close(fig) + + +def summarize_mismatch_vs_cost( + init_rows: list[dict[str, str]], + finite_rows: list[dict[str, str]], + width: int, + output_dim: int, + outdir: Path, +) -> list[dict[str, object]]: + """Contrast exponentially compounding matrix mismatch with measured costs.""" + depths = sorted({int(row["depth"]) for row in init_rows}) + records: list[dict[str, object]] = [] + for depth in depths: + init_subset = [row for row in init_rows if int(row["depth"]) == depth] + finite_subset = [ + row + for row in finite_rows + if int(row["depth"]) == depth and int(row["width"]) == width + ] + # One width x output block plus depth-1 width x width FA blocks. + joint_alignment_proxy = (1.0 / (width * output_dim)) * ( + 1.0 / (width * width) + ) ** (depth - 1) + records.append( + { + "depth": depth, + "width": width, + "joint_alignment_proxy": joint_alignment_proxy, + "negative_log10_joint_alignment": -math.log10(joint_alignment_proxy), + "initial_deficit_mean": float( + np.mean([float(row["empirical_deficit"]) for row in init_subset]) + ), + "finite_gap_mean": float( + np.mean([float(row["empirical_gap"]) for row in finite_subset]) + ), + } + ) + + write_dict_csv(outdir / "mismatch_vs_cost_by_depth.csv", records) + fig, ax = plt.subplots(figsize=(7.0, 4.8), dpi=180) + ax.plot( + [int(row["depth"]) for row in records], + [float(row["joint_alignment_proxy"]) for row in records], + "o--", + color="black", + label=r"joint matrix-alignment proxy $\prod_l 1/D_l$", + ) + ax.set_yscale("log") + ax.set_xlabel("hidden-layer depth") + ax.set_ylabel("joint squared-alignment proxy") + ax.grid(alpha=0.16, which="both") + ax2 = ax.twinx() + ax2.plot( + [int(row["depth"]) for row in records], + [float(row["initial_deficit_mean"]) for row in records], + "s-", + color="#2f6f9f", + label="measured initial deficit", + ) + ax2.plot( + [int(row["depth"]) for row in records], + [float(row["finite_gap_mean"]) for row in records], + "^-", + color="#c65f16", + label="measured finite-time gap", + ) + ax2.set_ylabel("optimization cost") + handles1, labels1 = ax.get_legend_handles_labels() + handles2, labels2 = ax2.get_legend_handles_labels() + ax2.legend(handles1 + handles2, labels1 + labels2, fontsize=8, loc="center right") + ax.set_title("Matrix mismatch compounds exponentially; optimization cost does not") + fig.tight_layout() + fig.savefig(outdir / "mismatch_vs_cost_by_depth.png", bbox_inches="tight") + plt.close(fig) + return records + + +def main() -> None: + args = parse_args() + if args.early_step <= 0 or args.early_step >= args.horizon: + raise ValueError("early-step must lie strictly between 0 and horizon") + torch.set_num_threads(args.torch_threads) + args.outdir.mkdir(parents=True, exist_ok=True) + x, y = make_data(args) + summary: dict[str, object] = { + "config": {key: str(value) if isinstance(value, Path) else value for key, value in vars(args).items()} + } + + if args.part == "summarize": + with (args.outdir / "initialization_rows.csv").open() as handle: + existing_init = list(csv.DictReader(handle)) + with (args.outdir / "finite_time_rows.csv").open() as handle: + existing_finite = list(csv.DictReader(handle)) + init_errors = np.asarray([float(row["calibration_error"]) for row in existing_init]) + init_stderr = np.asarray([float(row["empirical_stderr"]) for row in existing_init]) + summary["initialization"] = { + "rows": len(existing_init), + "max_abs_calibration_error": float(np.max(np.abs(init_errors))), + "max_standardized_error": float(np.max(np.abs(init_errors) / init_stderr)), + "rules": sorted({row["rule"] for row in existing_init}), + "depths": sorted({int(row["depth"]) for row in existing_init}), + } + empirical = np.asarray([float(row["empirical_gap"]) for row in existing_finite]) + summary["finite_time"] = { + "rows": len(existing_finite), + "depths": sorted({int(row["depth"]) for row in existing_finite}), + "widths": sorted({int(row["width"]) for row in existing_finite}), + "overall": { + predictor: { + "mae": float(np.mean(np.abs(predicted - empirical))), + "corr": safe_corr(predicted, empirical), + } + for predictor, predicted in ( + ( + "frozen", + np.asarray([float(row["fixed_gap"]) for row in existing_finite]), + ), + ( + "linear_velocity", + np.asarray([float(row["velocity_gap"]) for row in existing_finite]), + ), + ( + "early_retangent", + np.asarray([float(row["retangent_gap"]) for row in existing_finite]), + ), + ) + }, + } + mismatch_records = summarize_mismatch_vs_cost( + existing_init, + existing_finite, + width=args.width, + output_dim=args.output_dim, + outdir=args.outdir, + ) + summary["mismatch_vs_cost"] = { + "joint_alignment_proxy_depth_1": mismatch_records[0]["joint_alignment_proxy"], + "joint_alignment_proxy_depth_6": mismatch_records[-1]["joint_alignment_proxy"], + "orders_of_magnitude_drop": float( + math.log10( + float(mismatch_records[0]["joint_alignment_proxy"]) + / float(mismatch_records[-1]["joint_alignment_proxy"]) + ) + ), + } + (args.outdir / "summary.json").write_text(json.dumps(summary, indent=2) + "\n") + print(f"summary: {args.outdir / 'summary.json'}") + return + + if args.part in {"init", "both"}: + init_rows = run_initialization(args, x, y) + write_dataclass_csv(args.outdir / "initialization_rows.csv", init_rows) + plot_initial(init_rows, args.outdir) + summary["initialization"] = { + "rows": len(init_rows), + "max_abs_calibration_error": max(abs(row.calibration_error) for row in init_rows), + "max_standardized_error": max( + abs(row.calibration_error) / row.empirical_stderr for row in init_rows + ), + "rules": sorted({row.rule for row in init_rows}), + "depths": sorted({row.depth for row in init_rows}), + } + + if args.part in {"finite", "both"}: + finite_rows = run_finite_time(args, x, y) + write_dataclass_csv(args.outdir / "finite_time_rows.csv", finite_rows) + metrics = summarize_finite(finite_rows) + write_dict_csv(args.outdir / "finite_time_metrics.csv", metrics) # type: ignore[arg-type] + plot_finite(finite_rows, metrics, args.outdir) + summary["finite_time"] = { + "rows": len(finite_rows), + "depths": sorted({row.depth for row in finite_rows}), + "widths": sorted({row.width for row in finite_rows}), + "overall": { + predictor: { + "mae": float( + np.mean( + [ + abs(getattr(row, field) - row.empirical_gap) + for row in finite_rows + ] + ) + ), + "corr": safe_corr( + np.asarray([getattr(row, field) for row in finite_rows]), + np.asarray([row.empirical_gap for row in finite_rows]), + ), + } + for predictor, field in ( + ("frozen", "fixed_gap"), + ("linear_velocity", "velocity_gap"), + ("early_retangent", "retangent_gap"), + ) + }, + } + + (args.outdir / "summary.json").write_text(json.dumps(summary, indent=2) + "\n") + print(f"summary: {args.outdir / 'summary.json'}") + + +if __name__ == "__main__": + main() |
