summaryrefslogtreecommitdiff
path: root/scripts/actual_fa_initial_operator_moments.py
diff options
context:
space:
mode:
Diffstat (limited to 'scripts/actual_fa_initial_operator_moments.py')
-rw-r--r--scripts/actual_fa_initial_operator_moments.py293
1 files changed, 293 insertions, 0 deletions
diff --git a/scripts/actual_fa_initial_operator_moments.py b/scripts/actual_fa_initial_operator_moments.py
new file mode 100644
index 0000000..6ee4900
--- /dev/null
+++ b/scripts/actual_fa_initial_operator_moments.py
@@ -0,0 +1,293 @@
+#!/usr/bin/env python3
+"""Validate initial actual-FA operator moment predictions.
+
+The theorem tested here is conditional on a fixed forward initialization and
+training residual. With independent zero-mean feedback matrices, the hidden
+FA pseudo-gradient has zero conditional mean. Therefore the expected one-step
+FA/BP speed ratio equals the BP output-layer gradient-energy share.
+"""
+
+from __future__ import annotations
+
+import argparse
+import csv
+import sys
+from dataclasses import asdict, dataclass
+from pathlib import Path
+
+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 downstream_capacity_sweep as dcs # noqa: E402
+
+
+@dataclass(frozen=True)
+class MomentRow:
+ width: int
+ init_seed: int
+ feedback_samples: int
+ bp_speed: float
+ output_speed: float
+ predicted_mean_ratio: float
+ empirical_mean_ratio: float
+ empirical_std_ratio: float
+ predicted_mean_erosion: float
+ empirical_mean_erosion: float
+ empirical_std_erosion: float
+ mean_error: float
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser(
+ description="Actual FA initial operator moment validation."
+ )
+ parser.add_argument("--input-dim", type=int, default=16)
+ parser.add_argument("--output-dim", type=int, default=4)
+ parser.add_argument("--widths", type=int, nargs="+", default=[16, 24, 32, 48, 64, 96])
+ parser.add_argument("--train-samples", type=int, default=128)
+ parser.add_argument("--test-samples", type=int, default=512)
+ parser.add_argument("--init-seeds", type=int, default=8)
+ parser.add_argument("--feedback-samples", type=int, default=512)
+ parser.add_argument("--data-seed", type=int, default=123)
+ parser.add_argument(
+ "--feedback-scale",
+ choices=["relu", "fan-in", "unit"],
+ default="relu",
+ )
+ parser.add_argument("--device", choices=["cpu", "cuda"], default="cpu")
+ parser.add_argument("--torch-threads", type=int, default=8)
+ parser.add_argument(
+ "--outdir",
+ type=Path,
+ default=Path("outputs/actual_fa_initial_operator_moments"),
+ )
+ return parser.parse_args()
+
+
+def make_config(args: argparse.Namespace) -> dcs.RunConfig:
+ return dcs.RunConfig(
+ task="random",
+ input_dim=args.input_dim,
+ output_dim=args.output_dim,
+ teacher_rank=4,
+ teacher_width=64,
+ teacher_hidden_layers=2,
+ normalize_targets=False,
+ widths=args.widths,
+ train_samples=args.train_samples,
+ test_samples=args.test_samples,
+ probe_samples=8,
+ steps=1,
+ lr=1e-3,
+ optimizer="sgd",
+ init_seeds=args.init_seeds,
+ feedback_seeds=args.feedback_samples,
+ init_seed_offset=0,
+ feedback_seed_offset=0,
+ data_seed=args.data_seed,
+ noise_std=0.0,
+ feedback_scale=args.feedback_scale,
+ capacity_q=0.01,
+ jacobian_lambda_rel=1e-3,
+ skip_jacobian=True,
+ device=args.device,
+ torch_threads=args.torch_threads,
+ outdir=str(args.outdir),
+ plot=False,
+ )
+
+
+def flatten(grads: list[torch.Tensor]) -> torch.Tensor:
+ return torch.cat([grad.reshape(-1) for grad in grads])
+
+
+def squared_norm(grads: list[torch.Tensor]) -> float:
+ return float(sum(torch.sum(grad * grad).cpu() for grad in grads))
+
+
+def inner_product(left: list[torch.Tensor], right: list[torch.Tensor]) -> float:
+ return float(sum(torch.sum(a * b).cpu() for a, b in zip(left, right)))
+
+
+def run_one(
+ config: dcs.RunConfig,
+ x: torch.Tensor,
+ y: torch.Tensor,
+ width: int,
+ init_seed: int,
+ feedback_samples: int,
+) -> tuple[MomentRow, np.ndarray]:
+ weights = dcs.initialize_weights(config, width, init_seed)
+ bp_grads = dcs.gradients(weights, x, y, feedback=None)
+ bp_speed = squared_norm(bp_grads)
+ output_speed = squared_norm([bp_grads[-1]])
+ predicted_mean_ratio = output_speed / bp_speed
+ predicted_mean_erosion = 1.0 - predicted_mean_ratio
+
+ ratios: list[float] = []
+ for feedback_seed in range(feedback_samples):
+ feedback = dcs.init_feedback(config, width, feedback_seed)
+ fa_grads = dcs.gradients(weights, x, y, feedback=feedback)
+ ratios.append(inner_product(bp_grads, fa_grads) / bp_speed)
+
+ ratio_array = np.array(ratios, dtype=np.float64)
+ erosion_array = 1.0 - ratio_array
+ row = MomentRow(
+ width=width,
+ init_seed=init_seed,
+ feedback_samples=feedback_samples,
+ bp_speed=bp_speed,
+ output_speed=output_speed,
+ predicted_mean_ratio=predicted_mean_ratio,
+ empirical_mean_ratio=float(np.mean(ratio_array)),
+ empirical_std_ratio=float(np.std(ratio_array, ddof=1)),
+ predicted_mean_erosion=predicted_mean_erosion,
+ empirical_mean_erosion=float(np.mean(erosion_array)),
+ empirical_std_erosion=float(np.std(erosion_array, ddof=1)),
+ mean_error=float(np.mean(erosion_array) - predicted_mean_erosion),
+ )
+ return row, erosion_array
+
+
+def write_rows(path: Path, rows: list[MomentRow]) -> None:
+ path.parent.mkdir(parents=True, exist_ok=True)
+ with path.open("w", newline="") as handle:
+ writer = csv.DictWriter(handle, fieldnames=list(MomentRow.__annotations__.keys()))
+ writer.writeheader()
+ for row in rows:
+ writer.writerow(asdict(row))
+
+
+def plot_summary(rows: list[MomentRow], outdir: Path) -> list[Path]:
+ outdir.mkdir(parents=True, exist_ok=True)
+ paths: list[Path] = []
+
+ scatter = outdir / "predicted_vs_empirical_initial_erosion_mean.png"
+ fig, ax = plt.subplots(figsize=(6.0, 5.2))
+ widths = sorted({row.width for row in rows})
+ for width in widths:
+ subset = [row for row in rows if row.width == width]
+ ax.scatter(
+ [row.predicted_mean_erosion for row in subset],
+ [row.empirical_mean_erosion for row in subset],
+ s=32,
+ alpha=0.65,
+ label=f"w={width}",
+ )
+ values = [
+ value
+ for row in rows
+ for value in (row.predicted_mean_erosion, row.empirical_mean_erosion)
+ ]
+ lo, hi = min(values), max(values)
+ pad = 0.03 * (hi - lo + 1e-12)
+ ax.plot([lo - pad, hi + pad], [lo - pad, hi + pad], color="black", linewidth=1)
+ ax.set_xlabel("theory mean erosion: 1 - output BP speed share")
+ ax.set_ylabel("empirical mean erosion over feedback seeds")
+ ax.set_title("Actual FA initial operator moment")
+ ax.legend(fontsize=8, ncols=2)
+ fig.tight_layout()
+ fig.savefig(scatter, dpi=180)
+ plt.close(fig)
+ paths.append(scatter)
+
+ error = outdir / "initial_erosion_mean_error_by_width.png"
+ fig, ax = plt.subplots(figsize=(7.0, 4.6))
+ positions = []
+ means = []
+ ses = []
+ for width in widths:
+ errors = np.array([row.mean_error for row in rows if row.width == width])
+ positions.append(width)
+ means.append(float(np.mean(errors)))
+ ses.append(float(np.std(errors, ddof=1) / np.sqrt(len(errors))) if len(errors) > 1 else 0.0)
+ ax.errorbar(positions, means, yerr=ses, marker="o", linewidth=1.8, capsize=4)
+ ax.axhline(0.0, color="black", linewidth=1)
+ ax.set_xlabel("width")
+ ax.set_ylabel("empirical mean - theory mean")
+ ax.set_title("No-fit moment error")
+ fig.tight_layout()
+ fig.savefig(error, dpi=180)
+ plt.close(fig)
+ paths.append(error)
+
+ return paths
+
+
+def plot_histograms(samples_by_key: dict[tuple[int, int], np.ndarray], rows: list[MomentRow], outdir: Path) -> Path:
+ selected: list[MomentRow] = []
+ for width in sorted({row.width for row in rows}):
+ width_rows = [row for row in rows if row.width == width]
+ selected.append(width_rows[0])
+ selected = selected[:6]
+
+ cols = 3
+ rows_n = int(np.ceil(len(selected) / cols))
+ fig, axes = plt.subplots(rows_n, cols, figsize=(12.0, 3.3 * rows_n), squeeze=False)
+ for ax in axes.ravel():
+ ax.axis("off")
+
+ for ax, row in zip(axes.ravel(), selected):
+ ax.axis("on")
+ values = samples_by_key[(row.width, row.init_seed)]
+ ax.hist(values, bins=36, density=True, color="tab:orange", alpha=0.55)
+ ax.axvline(row.predicted_mean_erosion, color="tab:blue", linewidth=2, label="theory mean")
+ ax.axvline(row.empirical_mean_erosion, color="tab:orange", linewidth=2, linestyle="--", label="empirical mean")
+ ax.set_title(f"width={row.width}, init={row.init_seed}")
+ ax.set_xlabel("one-step erosion")
+ ax.set_ylabel("density")
+ ax.legend(fontsize=8)
+
+ fig.suptitle("Actual FA initial erosion distributions over random feedback", y=1.02)
+ fig.tight_layout()
+ path = outdir / "initial_erosion_histograms.png"
+ fig.savefig(path, dpi=180, bbox_inches="tight")
+ plt.close(fig)
+ return path
+
+
+def main() -> None:
+ args = parse_args()
+ torch.set_num_threads(args.torch_threads)
+ config = make_config(args)
+ x_train, y_train, *_ = dcs.make_data(config)
+
+ all_rows: list[MomentRow] = []
+ samples_by_key: dict[tuple[int, int], np.ndarray] = {}
+ for width in args.widths:
+ for init_seed in range(args.init_seeds):
+ row, samples = run_one(
+ config,
+ x_train,
+ y_train,
+ width,
+ init_seed,
+ args.feedback_samples,
+ )
+ all_rows.append(row)
+ samples_by_key[(width, init_seed)] = samples
+ print(
+ f"width={width} init={init_seed}: "
+ f"theory={row.predicted_mean_erosion:.6f} "
+ f"empirical={row.empirical_mean_erosion:.6f} "
+ f"error={row.mean_error:.3g}"
+ )
+
+ args.outdir.mkdir(parents=True, exist_ok=True)
+ csv_path = args.outdir / "initial_operator_moment_rows.csv"
+ write_rows(csv_path, all_rows)
+ paths = plot_summary(all_rows, args.outdir)
+ paths.append(plot_histograms(samples_by_key, all_rows, args.outdir))
+
+ print(f"rows: {csv_path}")
+ for path in paths:
+ print(f"plot: {path}")
+
+
+if __name__ == "__main__":
+ main()