From e65dd2e2460b1da48c83a930304cf9b269fc4447 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Mon, 27 Jul 2026 17:47:13 -0500 Subject: Add AAAI depth experiments and diagnostic figures --- scripts/learning_rate_compensation.py | 488 ++++++++++++++++++++++++++++++++++ 1 file changed, 488 insertions(+) create mode 100644 scripts/learning_rate_compensation.py (limited to 'scripts/learning_rate_compensation.py') diff --git a/scripts/learning_rate_compensation.py b/scripts/learning_rate_compensation.py new file mode 100644 index 0000000..10e1a64 --- /dev/null +++ b/scripts/learning_rate_compensation.py @@ -0,0 +1,488 @@ +#!/usr/bin/env python3 +"""Learning-rate controls for the finite-time FA/BP optimization gap. + +For every depth, the script first tunes BP and FA learning rates on held-out +initialization/feedback seeds. It then evaluates three FA conditions against +BP at its tuned rate: + +1. the same learning rate as BP; +2. initialization compensation ``eta_FA = eta_BP / rho``, where ``rho`` is + the exact expected fraction of BP's first-order decrease retained by FA; +3. FA's independently tuned stable learning rate. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +import sys +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 + + +@dataclass(frozen=True) +class TuneRow: + depth: int + rule: str + lr: float + trials: int + stable_trials: int + mean_final_loss: float + median_final_loss: float + + +@dataclass(frozen=True) +class EvaluationRow: + depth: int + width: int + init_seed: int + feedback_seed: int + regime: str + steps: int + bp_lr: float + fa_lr: float + rho: float + initial_loss: float + bp_final_loss: float + fa_final_loss: float + gap_to_bp: float + stable: bool + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--depths", type=int, nargs="+", default=[1, 2, 3, 4, 6]) + parser.add_argument("--width", type=int, default=64) + 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=128) + parser.add_argument("--steps", type=int, default=200) + parser.add_argument("--base-lr", type=float, default=1e-3) + parser.add_argument( + "--lr-grid", + type=float, + nargs="+", + default=[ + 0.00025, + 0.0005, + 0.001, + 0.002, + 0.004, + 0.008, + 0.016, + 0.032, + 0.064, + 0.128, + 0.256, + ], + ) + parser.add_argument("--tune-init-seeds", type=int, default=2) + parser.add_argument("--tune-feedback-seeds", type=int, default=2) + parser.add_argument("--eval-init-seeds", type=int, default=4) + parser.add_argument("--eval-feedback-seeds", type=int, default=4) + parser.add_argument("--data-seed", type=int, default=4242) + parser.add_argument("--feedback-scale", choices=["relu", "fan-in", "unit"], default="relu") + parser.add_argument("--stability-factor", type=float, default=10.0) + parser.add_argument("--torch-threads", type=int, default=16) + parser.add_argument("--postprocess-only", action="store_true") + parser.add_argument("--outdir", type=Path, default=Path("outputs/learning_rate_compensation")) + 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 stable_loss( + weights: list[torch.Tensor], + x: torch.Tensor, + y: torch.Tensor, + initial_loss: float, + stability_factor: float, +) -> tuple[float, bool]: + loss = fr.mse(weights, x, y) + stable = math.isfinite(loss) and loss <= stability_factor * initial_loss + return loss, stable + + +def tune_depth( + args: argparse.Namespace, + depth: int, + x: torch.Tensor, + y: torch.Tensor, +) -> tuple[float, float, list[TuneRow]]: + dims = [args.input_dim, *([args.width] * depth), args.output_dim] + records: list[TuneRow] = [] + for rule in ("bp", "fa"): + for lr in args.lr_grid: + losses: list[float] = [] + stable_flags: list[bool] = [] + for init_index in range(args.tune_init_seeds): + init_seed = 700_000 + init_index + weights0 = fr.initialize_mlp(dims, init_seed) + initial_loss = fr.mse(weights0, x, y) + feedback_count = 1 if rule == "bp" else args.tune_feedback_seeds + for feedback_index in range(feedback_count): + feedback = None + if rule == "fa": + feedback = fr.init_feedback( + dims, + 710_000 + 1_000 * init_index + feedback_index, + rule="fa", + mode=args.feedback_scale, + ) + final = fr.train( + weights0, + x, + y, + lr, + args.steps, + rule=rule, + feedback=feedback, + ) + loss, stable = stable_loss( + final, x, y, initial_loss, args.stability_factor + ) + losses.append(loss if math.isfinite(loss) else math.inf) + stable_flags.append(stable) + stable_losses = [loss for loss, stable in zip(losses, stable_flags) if stable] + records.append( + TuneRow( + depth=depth, + rule=rule.upper(), + lr=lr, + trials=len(losses), + stable_trials=sum(stable_flags), + mean_final_loss=float(np.mean(stable_losses)) if stable_losses else math.inf, + median_final_loss=float(np.median(stable_losses)) if stable_losses else math.inf, + ) + ) + print( + f"[tune] depth={depth} {rule.upper()} lr={lr:g}: " + f"stable={sum(stable_flags)}/{len(losses)}, " + f"mean={records[-1].mean_final_loss:.5g}", + flush=True, + ) + + best: dict[str, float] = {} + for rule in ("BP", "FA"): + candidates = [ + row + for row in records + if row.rule == rule and row.stable_trials == row.trials and math.isfinite(row.mean_final_loss) + ] + if not candidates: + raise RuntimeError(f"no fully stable {rule} learning rate at depth {depth}") + best[rule] = min(candidates, key=lambda row: row.mean_final_loss).lr + print( + f"[tune] depth={depth}: selected BP={best['BP']:g}, FA={best['FA']:g}", + flush=True, + ) + return best["BP"], best["FA"], records + + +def evaluate_depth( + args: argparse.Namespace, + depth: int, + bp_lr: float, + fa_lr: float, + x: torch.Tensor, + y: torch.Tensor, +) -> list[EvaluationRow]: + dims = [args.input_dim, *([args.width] * depth), args.output_dim] + rows: list[EvaluationRow] = [] + for init_index in range(args.eval_init_seeds): + init_seed = 10_000 + init_index + weights0 = fr.initialize_mlp(dims, init_seed) + initial_loss = fr.mse(weights0, x, y) + bp_grads = fr.gradients(weights0, x, y) + rho = fr.squared_norm([bp_grads[-1]]) / fr.squared_norm(bp_grads) + bp_losses: dict[str, tuple[float, bool]] = {} + for name, rate in (("base", args.base_lr), ("tuned", bp_lr)): + bp_final = fr.train(weights0, x, y, rate, args.steps, rule="bp") + bp_losses[name] = stable_loss( + bp_final, x, y, initial_loss, args.stability_factor + ) + if not bp_losses[name][1]: + raise RuntimeError( + f"{name} BP rate became unstable at depth={depth}, init={init_seed}" + ) + + regimes = ( + ("same_lr", args.base_lr, args.base_lr, "base"), + ("initial_compensation", args.base_lr, args.base_lr / rho, "base"), + ("independently_tuned", bp_lr, fa_lr, "tuned"), + ) + for feedback_index in range(args.eval_feedback_seeds): + feedback_seed = 100_000 + 1_000 * init_index + feedback_index + feedback = fr.init_feedback( + dims, + feedback_seed, + rule="fa", + mode=args.feedback_scale, + ) + for regime, bp_rate, fa_rate, bp_key in regimes: + fa_final = fr.train( + weights0, + x, + y, + fa_rate, + args.steps, + rule="fa", + feedback=feedback, + ) + fa_final_loss, stable = stable_loss( + fa_final, x, y, initial_loss, args.stability_factor + ) + rows.append( + EvaluationRow( + depth=depth, + width=args.width, + init_seed=init_seed, + feedback_seed=feedback_seed, + regime=regime, + steps=args.steps, + bp_lr=bp_rate, + fa_lr=fa_rate, + rho=rho, + initial_loss=initial_loss, + bp_final_loss=bp_losses[bp_key][0], + fa_final_loss=fa_final_loss, + gap_to_bp=fa_final_loss - bp_losses[bp_key][0], + stable=stable, + ) + ) + print( + f"[eval] depth={depth} init={init_index}: rho={rho:.4f}, " + f"compensation={args.base_lr / rho:.5g}", + flush=True, + ) + return rows + + +def write_dataclass_csv(path: Path, rows: list[object]) -> None: + if not rows: + return + 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 summarize(rows: list[EvaluationRow]) -> list[dict[str, object]]: + result: list[dict[str, object]] = [] + for depth in sorted({row.depth for row in rows}): + for regime in ("same_lr", "initial_compensation", "independently_tuned"): + subset = [row for row in rows if row.depth == depth and row.regime == regime] + gaps = np.asarray([row.gap_to_bp for row in subset]) + fa_losses = np.asarray([row.fa_final_loss for row in subset]) + init_gap_means = np.asarray( + [ + np.mean( + [row.gap_to_bp for row in subset if row.init_seed == init_seed] + ) + for init_seed in sorted({row.init_seed for row in subset}) + ] + ) + result.append( + { + "depth": depth, + "regime": regime, + "rows": len(subset), + "stable_rows": sum(row.stable for row in subset), + "bp_lr": subset[0].bp_lr, + "fa_lr_mean": float(np.mean([row.fa_lr for row in subset])), + "rho_mean": float(np.mean([row.rho for row in subset])), + "bp_loss_mean": float(np.mean([row.bp_final_loss for row in subset])), + "fa_loss_mean": float(fa_losses.mean()), + "gap_mean": float(gaps.mean()), + "gap_std": float(init_gap_means.std(ddof=1)), + "gap_sem": float( + init_gap_means.std(ddof=1) / math.sqrt(len(init_gap_means)) + ), + } + ) + return result + + +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_results(summary: list[dict[str, object]], tune_rows: list[TuneRow], outdir: Path) -> None: + fig, ax = plt.subplots(figsize=(7.0, 4.8), dpi=180) + styles = { + "same_lr": ("o", "same learning rate"), + "initial_compensation": ("s", "initial-rate compensation"), + "independently_tuned": ("^", "independently tuned"), + } + for regime, (marker, label) in styles.items(): + subset = [row for row in summary if row["regime"] == regime] + ax.errorbar( + [int(row["depth"]) for row in subset], + [float(row["gap_mean"]) for row in subset], + yerr=[float(row["gap_sem"]) for row in subset], + marker=marker, + capsize=3, + label=label, + ) + ax.axhline(0, color="black", lw=1) + ax.set_xlabel("hidden-layer depth") + ax.set_ylabel("FA loss - BP loss at fixed steps") + ax.set_title("Finite-time gap under three learning-rate controls") + ax.grid(alpha=0.16) + ax.legend() + fig.tight_layout() + fig.savefig(outdir / "learning_rate_control_by_depth.png", bbox_inches="tight") + plt.close(fig) + + fig, axes = plt.subplots(1, 2, figsize=(10.5, 4.2), dpi=180, sharey=True) + for ax, rule in zip(axes, ("BP", "FA")): + for depth in sorted({row.depth for row in tune_rows}): + subset = [row for row in tune_rows if row.rule == rule and row.depth == depth] + ax.plot( + [row.lr for row in subset], + [row.mean_final_loss for row in subset], + marker="o", + ms=3, + label=f"depth {depth}", + ) + ax.set_xscale("log") + ax.set_yscale("log") + ax.set_xlabel("learning rate") + ax.set_title(rule) + ax.grid(alpha=0.16, which="both") + axes[0].set_ylabel("held-out tuning loss") + axes[1].legend(fontsize=7, ncols=2) + fig.tight_layout() + fig.savefig(outdir / "learning_rate_tuning_curves.png", bbox_inches="tight") + plt.close(fig) + + +def main() -> None: + args = parse_args() + torch.set_num_threads(args.torch_threads) + args.outdir.mkdir(parents=True, exist_ok=True) + if args.postprocess_only: + with (args.outdir / "tuning_rows.csv").open() as handle: + tune_dicts = list(csv.DictReader(handle)) + with (args.outdir / "evaluation_rows.csv").open() as handle: + eval_dicts = list(csv.DictReader(handle)) + tune_rows = [ + TuneRow( + depth=int(row["depth"]), + rule=row["rule"], + lr=float(row["lr"]), + trials=int(row["trials"]), + stable_trials=int(row["stable_trials"]), + mean_final_loss=float(row["mean_final_loss"]), + median_final_loss=float(row["median_final_loss"]), + ) + for row in tune_dicts + ] + eval_rows = [ + EvaluationRow( + depth=int(row["depth"]), + width=int(row["width"]), + init_seed=int(row["init_seed"]), + feedback_seed=int(row["feedback_seed"]), + regime=row["regime"], + steps=int(row["steps"]), + bp_lr=float(row["bp_lr"]), + fa_lr=float(row["fa_lr"]), + rho=float(row["rho"]), + initial_loss=float(row["initial_loss"]), + bp_final_loss=float(row["bp_final_loss"]), + fa_final_loss=float(row["fa_final_loss"]), + gap_to_bp=float(row["gap_to_bp"]), + stable=row["stable"] == "True", + ) + for row in eval_dicts + ] + summary_rows = summarize(eval_rows) + write_dict_csv(args.outdir / "regime_summary.csv", summary_rows) + plot_results(summary_rows, tune_rows, args.outdir) + summary_path = args.outdir / "summary.json" + payload = json.loads(summary_path.read_text()) + payload["regime_summary"] = summary_rows + payload["uncertainty"] = "SEM and SD are computed across forward-initialization means." + summary_path.write_text(json.dumps(payload, indent=2) + "\n") + print(f"summary: {summary_path}") + return + x, y = make_data(args) + all_tune: list[TuneRow] = [] + all_eval: list[EvaluationRow] = [] + selected: dict[str, dict[str, float]] = {} + for depth in args.depths: + bp_lr, fa_lr, tune_rows = tune_depth(args, depth, x, y) + all_tune.extend(tune_rows) + while True: + evaluation = evaluate_depth(args, depth, bp_lr, fa_lr, x, y) + tuned_rows = [row for row in evaluation if row.regime == "independently_tuned"] + if all(row.stable for row in tuned_rows): + break + lower_candidates = [ + row + for row in tune_rows + if row.rule == "FA" + and row.lr < fa_lr + and row.stable_trials == row.trials + and math.isfinite(row.mean_final_loss) + ] + if not lower_candidates: + raise RuntimeError( + f"no lower FA rate passed stability confirmation at depth {depth}" + ) + previous = fa_lr + fa_lr = min(lower_candidates, key=lambda row: row.mean_final_loss).lr + print( + f"[stability] depth={depth}: FA lr {previous:g} failed on held-out " + f"evaluation seeds; retrying {fa_lr:g}", + flush=True, + ) + selected[str(depth)] = {"bp_lr": bp_lr, "fa_lr": fa_lr} + all_eval.extend(evaluation) + + write_dataclass_csv(args.outdir / "tuning_rows.csv", all_tune) + write_dataclass_csv(args.outdir / "evaluation_rows.csv", all_eval) + summary_rows = summarize(all_eval) + write_dict_csv(args.outdir / "regime_summary.csv", summary_rows) + plot_results(summary_rows, all_tune, args.outdir) + payload = { + "config": {key: str(value) if isinstance(value, Path) else value for key, value in vars(args).items()}, + "selected_learning_rates": selected, + "rows": len(all_eval), + "regime_summary": summary_rows, + "uncertainty": "SEM and SD are computed across forward-initialization means.", + } + (args.outdir / "summary.json").write_text(json.dumps(payload, indent=2) + "\n") + print(f"summary: {args.outdir / 'summary.json'}") + + +if __name__ == "__main__": + main() -- cgit v1.2.3