From e65dd2e2460b1da48c83a930304cf9b269fc4447 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Mon, 27 Jul 2026 17:47:13 -0500 Subject: Add AAAI depth experiments and diagnostic figures --- scripts/cnn_initialization_validation.py | 320 +++++++++++++++++++++++++++++++ 1 file changed, 320 insertions(+) create mode 100644 scripts/cnn_initialization_validation.py (limited to 'scripts/cnn_initialization_validation.py') diff --git a/scripts/cnn_initialization_validation.py b/scripts/cnn_initialization_validation.py new file mode 100644 index 0000000..dcf4f92 --- /dev/null +++ b/scripts/cnn_initialization_validation.py @@ -0,0 +1,320 @@ +#!/usr/bin/env python3 +"""Cross-entropy CNN validation of the exact FA/DFA initialization cost. + +The model has two 3x3 convolutional ReLU layers, global average pooling, and a +linear classifier. The experiment runs the real convolutional FA and DFA +backward rules for every feedback draw and compares their mean first-order +loss-decrease deficit with the hidden-parameter share from one BP backward +pass. This directly tests that the initialization result is not a +squared-loss MLP identity. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +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 +import torch.nn.functional as functional + + +Tensor = torch.Tensor + + +@dataclass(frozen=True) +class CNNRow: + architecture: str + loss: str + rule: str + init_seed: int + feedback_draws: int + bp_speed: float + output_speed: float + predicted_deficit: float + empirical_deficit: float + empirical_stderr: float + empirical_std: float + calibration_error: float + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--data-root", type=Path, default=Path("data")) + parser.add_argument("--train-samples", type=int, default=128) + parser.add_argument("--image-size", type=int, default=8) + parser.add_argument("--channels", type=int, nargs=2, default=[8, 12]) + parser.add_argument("--init-seeds", type=int, default=4) + parser.add_argument("--feedback-draws", type=int, default=512) + parser.add_argument("--torch-threads", type=int, default=16) + parser.add_argument("--self-test", action="store_true") + parser.add_argument( + "--outdir", type=Path, default=Path("outputs/cnn_initialization_validation") + ) + return parser.parse_args() + + +def load_mnist(root: Path, samples: int, image_size: int) -> tuple[Tensor, Tensor]: + from torchvision import datasets + + dataset = datasets.MNIST(root=str(root), train=True, download=True) + generator = torch.Generator().manual_seed(2027) + indices = torch.randperm(len(dataset), generator=generator)[:samples] + images = dataset.data[indices].to(torch.float64).unsqueeze(1) / 255.0 + images = functional.interpolate( + images, size=(image_size, image_size), mode="bilinear", align_corners=False + ) + images = (images - images.mean()) / (images.std() + 1e-12) + labels = dataset.targets[indices] + return images, labels + + +def initialize_weights(channels: tuple[int, int], classes: int, seed: int) -> list[Tensor]: + c1, c2 = channels + generator = torch.Generator().manual_seed(seed) + w1 = torch.randn(c1, 1, 3, 3, generator=generator, dtype=torch.float64) * math.sqrt(2 / 9) + w2 = torch.randn(c2, c1, 3, 3, generator=generator, dtype=torch.float64) * math.sqrt( + 2 / (9 * c1) + ) + w3 = torch.randn(classes, c2, generator=generator, dtype=torch.float64) / math.sqrt(c2) + return [w1, w2, w3] + + +def forward(weights: list[Tensor], x: Tensor) -> tuple[Tensor, tuple[Tensor, ...]]: + w1, w2, w3 = weights + z1 = functional.conv2d(x, w1, padding=1) + h1 = torch.relu(z1) + z2 = functional.conv2d(h1, w2, padding=1) + h2 = torch.relu(z2) + pooled = h2.mean(dim=(2, 3)) + logits = pooled @ w3.T + return logits, (z1, h1, z2, h2, pooled) + + +def autograd_gradients(weights: list[Tensor], x: Tensor, labels: Tensor) -> list[Tensor]: + leaves = [weight.clone().detach().requires_grad_(True) for weight in weights] + logits, _ = forward(leaves, x) + loss = functional.cross_entropy(logits, labels) + return list(torch.autograd.grad(loss, leaves)) + + +def init_feedback( + weights: list[Tensor], classes: int, seed: int, rule: str +) -> list[Tensor]: + w1, w2, _w3 = weights + c1, c2 = w1.shape[0], w2.shape[0] + generator = torch.Generator().manual_seed(seed) + if rule == "fa": + conv_map = torch.randn(w2.shape, generator=generator, dtype=torch.float64) * math.sqrt( + 2 / (9 * c1) + ) + output_map = torch.randn( + classes, c2, generator=generator, dtype=torch.float64 + ) * math.sqrt(2 / c2) + return [conv_map, output_map] + if rule == "dfa": + direct_first = torch.randn( + classes, c1, generator=generator, dtype=torch.float64 + ) * math.sqrt(2 / c1) + direct_second = torch.randn( + classes, c2, generator=generator, dtype=torch.float64 + ) * math.sqrt(2 / c2) + return [direct_first, direct_second] + raise ValueError("rule must be fa or dfa") + + +def manual_gradients( + weights: list[Tensor], + x: Tensor, + labels: Tensor, + rule: str, + feedback: list[Tensor] | None = None, +) -> list[Tensor]: + if rule not in {"bp", "fa", "dfa"}: + raise ValueError("rule must be bp, fa, or dfa") + if rule != "bp" and feedback is None: + raise ValueError("FA/DFA require feedback maps") + + w1, w2, w3 = weights + logits, (z1, h1, z2, h2, pooled) = forward(weights, x) + batch, height, width = h2.shape[0], h2.shape[2], h2.shape[3] + probabilities = torch.softmax(logits, dim=1) + probabilities[torch.arange(batch), labels] -= 1.0 + delta_output = probabilities / batch + grad_w3 = delta_output.T @ pooled + + if rule == "bp": + delta_pooled = delta_output @ w3 + delta_z2 = delta_pooled[:, :, None, None].expand_as(h2) / (height * width) + delta_z2 = delta_z2 * (z2 > 0) + delta_h1 = functional.conv_transpose2d(delta_z2, w2, padding=1) + delta_z1 = delta_h1 * (z1 > 0) + elif rule == "fa": + assert feedback is not None + conv_map, output_map = feedback + delta_pooled = delta_output @ output_map + delta_z2 = delta_pooled[:, :, None, None].expand_as(h2) / (height * width) + delta_z2 = delta_z2 * (z2 > 0) + delta_h1 = functional.conv_transpose2d(delta_z2, conv_map, padding=1) + delta_z1 = delta_h1 * (z1 > 0) + else: + assert feedback is not None + direct_first, direct_second = feedback + direct_z2 = delta_output @ direct_second + delta_z2 = direct_z2[:, :, None, None].expand_as(h2) / (height * width) + delta_z2 = delta_z2 * (z2 > 0) + direct_z1 = delta_output @ direct_first + delta_z1 = direct_z1[:, :, None, None].expand_as(z1) / (height * width) + delta_z1 = delta_z1 * (z1 > 0) + + grad_w2 = torch.nn.grad.conv2d_weight(h1, w2.shape, delta_z2, padding=1) + grad_w1 = torch.nn.grad.conv2d_weight(x, w1.shape, delta_z1, padding=1) + return [grad_w1, grad_w2, grad_w3] + + +def squared_norm(grads: list[Tensor]) -> float: + return float(sum(torch.sum(grad * grad) for grad in grads)) + + +def inner_product(left: list[Tensor], right: list[Tensor]) -> float: + return float(sum(torch.sum(a * b) for a, b in zip(left, right))) + + +def self_test() -> None: + generator = torch.Generator().manual_seed(31) + x = torch.randn(5, 1, 6, 6, generator=generator, dtype=torch.float64) + labels = torch.randint(0, 3, (5,), generator=generator) + weights = initialize_weights((4, 5), classes=3, seed=41) + automatic = autograd_gradients(weights, x, labels) + manual = manual_gradients(weights, x, labels, rule="bp") + error = max(float((a - b).abs().max()) for a, b in zip(automatic, manual)) + print(f"CNN BP manual/autograd max error: {error:.3e}") + assert error < 1e-11 + for rule in ("fa", "dfa"): + feedback = init_feedback(weights, classes=3, seed=51, rule=rule) + rule_grads = manual_gradients(weights, x, labels, rule=rule, feedback=feedback) + output_error = float((rule_grads[-1] - manual[-1]).abs().max()) + print(f"CNN {rule.upper()} output-gradient error: {output_error:.3e}") + assert output_error < 1e-12 + print("CNN initialization self-test PASSED") + + +def run(args: argparse.Namespace, x: Tensor, labels: Tensor) -> list[CNNRow]: + rows: list[CNNRow] = [] + classes = 10 + channels = (args.channels[0], args.channels[1]) + for init_index in range(args.init_seeds): + init_seed = 10_000 + init_index + weights = initialize_weights(channels, classes, init_seed) + bp = manual_gradients(weights, x, labels, rule="bp") + bp_speed = squared_norm(bp) + output_speed = squared_norm([bp[-1]]) + prediction = 1.0 - output_speed / bp_speed + for rule_index, rule in enumerate(("fa", "dfa")): + deficits: list[float] = [] + for draw in range(args.feedback_draws): + feedback = init_feedback( + weights, + classes, + seed=100_000 + 10_000 * init_index + 2 * draw + rule_index, + rule=rule, + ) + rule_grads = manual_gradients(weights, x, labels, rule, feedback) + deficits.append(1.0 - inner_product(bp, rule_grads) / bp_speed) + values = np.asarray(deficits) + empirical = float(values.mean()) + row = CNNRow( + architecture=f"Conv({channels[0]},{channels[1]})-GAP-linear", + loss="cross_entropy", + rule=rule.upper(), + init_seed=init_seed, + feedback_draws=args.feedback_draws, + bp_speed=bp_speed, + output_speed=output_speed, + predicted_deficit=prediction, + empirical_deficit=empirical, + empirical_stderr=float(values.std(ddof=1) / math.sqrt(len(values))), + empirical_std=float(values.std(ddof=1)), + calibration_error=empirical - prediction, + ) + rows.append(row) + print( + f"[CNN] {rule.upper()} init={init_index}: prediction={prediction:.5f}, " + f"measured={empirical:.5f} +/- {row.empirical_stderr:.5f}", + flush=True, + ) + return rows + + +def write_rows(path: Path, rows: list[CNNRow]) -> None: + with path.open("w", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=list(CNNRow.__annotations__)) + writer.writeheader() + for row in rows: + writer.writerow(asdict(row)) + + +def plot(rows: list[CNNRow], outdir: Path) -> None: + fig, ax = plt.subplots(figsize=(5.5, 5.0), dpi=180) + for rule, marker, color in (("FA", "o", "#2f6f9f"), ("DFA", "s", "#c65f16")): + subset = [row for row in rows if row.rule == rule] + ax.errorbar( + [row.predicted_deficit for row in subset], + [row.empirical_deficit for row in subset], + yerr=[2 * row.empirical_stderr for row in subset], + fmt=marker, + color=color, + capsize=2, + label=rule, + ) + values = [value for row in rows for value in (row.predicted_deficit, row.empirical_deficit)] + lo, hi = min(values), max(values) + pad = 0.05 * (hi - lo + 1e-12) + ax.plot([lo - pad, hi + pad], [lo - pad, hi + pad], color="black", lw=1) + ax.set_xlabel("exact expected initial deficit") + ax.set_ylabel("measured mean initial deficit") + ax.set_title("Cross-entropy CNN: FA and DFA initialization cost") + ax.grid(alpha=0.16) + ax.legend() + fig.tight_layout() + fig.savefig(outdir / "cnn_initialization_calibration.png", bbox_inches="tight") + plt.close(fig) + + +def main() -> None: + args = parse_args() + torch.set_num_threads(args.torch_threads) + if args.self_test: + self_test() + return + args.outdir.mkdir(parents=True, exist_ok=True) + x, labels = load_mnist(args.data_root, args.train_samples, args.image_size) + self_test() + rows = run(args, x, labels) + write_rows(args.outdir / "cnn_initialization_rows.csv", rows) + plot(rows, args.outdir) + payload = { + "config": {key: str(value) if isinstance(value, Path) else value for key, value in vars(args).items()}, + "rows": len(rows), + "max_abs_calibration_error": max(abs(row.calibration_error) for row in rows), + "max_standardized_error": max( + abs(row.calibration_error) / row.empirical_stderr for row in rows + ), + "rules": sorted({row.rule for row in rows}), + "architecture": rows[0].architecture, + "loss": rows[0].loss, + } + (args.outdir / "summary.json").write_text(json.dumps(payload, indent=2) + "\n") + print(f"summary: {args.outdir / 'summary.json'}") + + +if __name__ == "__main__": + main() -- cgit v1.2.3