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
|
#!/usr/bin/env python3
"""Numerical tests for the shared BP/FA/DFA implementation."""
from __future__ import annotations
import torch
import feedback_rules as fr
def main() -> None:
torch.set_num_threads(2)
dims = [5, 7, 6, 3]
weights = fr.initialize_mlp(dims, seed=11)
generator = torch.Generator().manual_seed(7)
x = torch.randn(9, dims[0], generator=generator, dtype=torch.float64)
y = torch.randn(9, dims[-1], generator=generator, dtype=torch.float64)
manual = fr.gradients(weights, x, y)
auto = fr.autograd_bp_gradients(weights, x, y)
bp_error = max(float((left - right).abs().max()) for left, right in zip(manual, auto))
fa = fr.init_feedback(dims, seed=21, rule="fa")
dfa = fr.init_feedback(dims, seed=22, rule="dfa")
fa_grads = fr.gradients(weights, x, y, rule="fa", feedback=fa)
dfa_grads = fr.gradients(weights, x, y, rule="dfa", feedback=dfa)
output_fa_error = float((manual[-1] - fa_grads[-1]).abs().max())
output_dfa_error = float((manual[-1] - dfa_grads[-1]).abs().max())
bp_speed = fr.squared_norm(manual)
output_share = fr.squared_norm([manual[-1]]) / bp_speed
for rule in ("fa", "dfa"):
ratios = []
for draw in range(4096):
feedback = fr.init_feedback(dims, seed=100_000 + draw, rule=rule)
rule_grads = fr.gradients(weights, x, y, rule=rule, feedback=feedback)
ratios.append(fr.inner_product(manual, rule_grads) / bp_speed)
mean_ratio = sum(ratios) / len(ratios)
stderr = torch.tensor(ratios, dtype=torch.float64).std(unbiased=True).item() / len(ratios) ** 0.5
error = abs(mean_ratio - output_share)
print(
f"{rule.upper()}: mean ratio={mean_ratio:.6f}, prediction={output_share:.6f}, "
f"error={error:.3g}, stderr={stderr:.3g}"
)
assert error <= 4.5 * stderr + 1e-4
print(f"BP manual/autograd max error: {bp_error:.3e}")
print(f"FA/DFA output-gradient errors: {output_fa_error:.3e}, {output_dfa_error:.3e}")
assert max(bp_error, output_fa_error, output_dfa_error) < 1e-11
print("feedback_rules self-test PASSED")
if __name__ == "__main__":
main()
|