summaryrefslogtreecommitdiff
path: root/scripts/test_feedback_rules.py
blob: 99573153450c5c9fb38c1df0821100716dd71c70 (plain)
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()