summaryrefslogtreecommitdiff
path: root/scripts/test_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/test_feedback_rules.py
parent82a49011de15287583d8cfec11ac5cca7efee747 (diff)
Add AAAI depth experiments and diagnostic figures
Diffstat (limited to 'scripts/test_feedback_rules.py')
-rw-r--r--scripts/test_feedback_rules.py54
1 files changed, 54 insertions, 0 deletions
diff --git a/scripts/test_feedback_rules.py b/scripts/test_feedback_rules.py
new file mode 100644
index 0000000..9957315
--- /dev/null
+++ b/scripts/test_feedback_rules.py
@@ -0,0 +1,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()