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