summaryrefslogtreecommitdiff
path: root/scripts/feedback_rules.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-27 17:47:13 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-27 17:47:13 -0500
commite65dd2e2460b1da48c83a930304cf9b269fc4447 (patch)
tree94f583b0fb5e9989bd6c68f7f0e950c6ef91756d /scripts/feedback_rules.py
parent82a49011de15287583d8cfec11ac5cca7efee747 (diff)
Add AAAI depth experiments and diagnostic figures
Diffstat (limited to 'scripts/feedback_rules.py')
-rw-r--r--scripts/feedback_rules.py187
1 files changed, 187 insertions, 0 deletions
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))