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