summaryrefslogtreecommitdiff
path: root/scripts/cnn_initialization_validation.py
diff options
context:
space:
mode:
Diffstat (limited to 'scripts/cnn_initialization_validation.py')
-rw-r--r--scripts/cnn_initialization_validation.py320
1 files changed, 320 insertions, 0 deletions
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()