1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
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))
|