#!/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()