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/feedback_rules.py | 187 ++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 187 insertions(+) create mode 100644 scripts/feedback_rules.py (limited to 'scripts/feedback_rules.py') diff --git a/scripts/feedback_rules.py b/scripts/feedback_rules.py new file mode 100644 index 0000000..b658d1d --- /dev/null +++ b/scripts/feedback_rules.py @@ -0,0 +1,187 @@ +#!/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)) -- cgit v1.2.3