#!/usr/bin/env python3 """Shared BP, FA, and DFA utilities for arbitrary-depth MLP experiments. The older experiment scripts predate the DFA validation and expose only the layerwise-FA rule. This module keeps the new paper-critical experiments on a single, explicitly tested implementation. All losses use ``0.5 * mean(sum((prediction - target)**2, dim=1))``. """ from __future__ import annotations import math from collections.abc import Sequence import torch Tensor = torch.Tensor Rule = str def initialize_mlp(dims: Sequence[int], seed: int, device: str = "cpu") -> list[Tensor]: """He-initialize hidden layers and fan-in initialize the linear output.""" if len(dims) < 3: raise ValueError("dims must contain input, at least one hidden layer, and output") generator = torch.Generator(device=device) generator.manual_seed(seed) weights: list[Tensor] = [] for layer, (fan_in, fan_out) in enumerate(zip(dims[:-1], dims[1:])): scale = math.sqrt(2.0 / fan_in) if layer < len(dims) - 2 else 1.0 / math.sqrt(fan_in) weights.append( torch.randn( fan_out, fan_in, generator=generator, dtype=torch.float64, device=device, ) * scale ) return weights def clone_weights(weights: Sequence[Tensor]) -> list[Tensor]: return [weight.clone() for weight in weights] def forward(weights: Sequence[Tensor], x: Tensor) -> tuple[list[Tensor], list[Tensor]]: activations = [x] preacts: list[Tensor] = [] current = x for layer, weight in enumerate(weights): current = current @ weight.T preacts.append(current) if layer < len(weights) - 1: current = torch.relu(current) activations.append(current) return activations, preacts def predict(weights: Sequence[Tensor], x: Tensor) -> Tensor: return forward(weights, x)[0][-1] def mse(weights: Sequence[Tensor], x: Tensor, y: Tensor) -> float: error = predict(weights, x) - y return float(0.5 * torch.mean(torch.sum(error * error, dim=1))) def _feedback_scale(rows: int, mode: str) -> float: if mode == "relu": return math.sqrt(2.0 / rows) if mode == "fan-in": return math.sqrt(1.0 / rows) if mode == "unit": return 1.0 raise ValueError(f"unknown feedback scale: {mode}") def init_feedback( dims: Sequence[int], seed: int, rule: Rule, mode: str = "relu", device: str = "cpu", ) -> list[Tensor]: """Initialize independent zero-mean feedback maps for FA or DFA. FA map ``l`` has shape ``(dims[l+1], dims[l+2])`` and replaces the transpose action of the next forward matrix. DFA map ``l`` has shape ``(dims[l+1], dims[-1])`` and maps output error directly to hidden layer ``l``. The output layer never receives a feedback matrix. """ if rule not in {"fa", "dfa"}: raise ValueError("feedback rule must be 'fa' or 'dfa'") generator = torch.Generator(device=device) generator.manual_seed(seed) shapes = ( [(dims[layer + 1], dims[layer + 2]) for layer in range(len(dims) - 2)] if rule == "fa" else [(dims[layer + 1], dims[-1]) for layer in range(len(dims) - 2)] ) return [ torch.randn(rows, cols, generator=generator, dtype=torch.float64, device=device) * _feedback_scale(rows, mode) for rows, cols in shapes ] def gradients( weights: Sequence[Tensor], x: Tensor, y: Tensor, rule: Rule = "bp", feedback: Sequence[Tensor] | None = None, ) -> list[Tensor]: """Return BP gradients or FA/DFA pseudo-gradients for an MLP.""" if rule not in {"bp", "fa", "dfa"}: raise ValueError("rule must be one of: bp, fa, dfa") if rule == "bp" and feedback is not None: raise ValueError("BP does not take feedback maps") if rule != "bp" and (feedback is None or len(feedback) != len(weights) - 1): raise ValueError(f"{rule.upper()} needs one feedback map per hidden layer") activations, preacts = forward(weights, x) batch = x.shape[0] deltas: list[Tensor] = [torch.empty(0, dtype=x.dtype, device=x.device) for _ in weights] deltas[-1] = (activations[-1] - y) / batch if rule == "dfa": assert feedback is not None output_delta = deltas[-1] for layer in range(len(weights) - 2, -1, -1): deltas[layer] = (output_delta @ feedback[layer].T) * (preacts[layer] > 0) else: for layer in range(len(weights) - 2, -1, -1): if rule == "bp": back = deltas[layer + 1] @ weights[layer + 1] else: assert feedback is not None back = deltas[layer + 1] @ feedback[layer].T deltas[layer] = back * (preacts[layer] > 0) return [delta.T @ activations[layer] for layer, delta in enumerate(deltas)] def train( weights0: Sequence[Tensor], x: Tensor, y: Tensor, lr: float, steps: int, rule: Rule = "bp", feedback: Sequence[Tensor] | None = None, ) -> list[Tensor]: weights = clone_weights(weights0) for _ in range(steps): grads = gradients(weights, x, y, rule=rule, feedback=feedback) weights = [weight - lr * grad for weight, grad in zip(weights, grads)] if not all(bool(torch.isfinite(weight).all()) for weight in weights): break return weights def squared_norm(grads: Sequence[Tensor]) -> float: return float(sum(torch.sum(grad * grad) for grad in grads)) def inner_product(left: Sequence[Tensor], right: Sequence[Tensor]) -> float: return float(sum(torch.sum(a * b) for a, b in zip(left, right))) def first_order_statistics( bp_grads: Sequence[Tensor], rule_grads: Sequence[Tensor] ) -> tuple[float, float, float, float]: """Return BP speed, output speed, rule/BP decrease ratio, and deficit.""" bp_speed = squared_norm(bp_grads) output_speed = squared_norm([bp_grads[-1]]) ratio = inner_product(bp_grads, rule_grads) / bp_speed return bp_speed, output_speed, ratio, 1.0 - ratio def autograd_bp_gradients(weights: Sequence[Tensor], x: Tensor, y: Tensor) -> list[Tensor]: leaves = [weight.clone().detach().requires_grad_(True) for weight in weights] error = predict(leaves, x) - y loss = 0.5 * torch.mean(torch.sum(error * error, dim=1)) return list(torch.autograd.grad(loss, leaves))