summaryrefslogtreecommitdiff
path: root/scripts/learning_rate_compensation.py
diff options
context:
space:
mode:
Diffstat (limited to 'scripts/learning_rate_compensation.py')
-rw-r--r--scripts/learning_rate_compensation.py488
1 files changed, 488 insertions, 0 deletions
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()