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