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/test_feedback_rules.py | 54 ++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 54 insertions(+) create mode 100644 scripts/test_feedback_rules.py (limited to 'scripts/test_feedback_rules.py') 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() -- cgit v1.2.3