diff options
Diffstat (limited to 'scripts/actual_fa_initial_operator_moments.py')
| -rw-r--r-- | scripts/actual_fa_initial_operator_moments.py | 293 |
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() |
