summaryrefslogtreecommitdiff
path: root/experiments/rain_ep_adapter_smoke.py
blob: 425cc26748a08b42f4a027d1d0c9cd76bafc810d (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
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
#!/usr/bin/env python3
"""Integration smoke test against the pinned Rain EP implementation."""

from __future__ import annotations

import argparse
from pathlib import Path
import sys

import torch

ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))

from sdil.rain_ep_adapter import (  # noqa: E402
    RainGradientCorrector,
    attach_to_rain_estimator,
)


RAIN_REVISION = "6b253fd8a5d267535f58ab79992256ef10031ceb"


def build_estimator(author_root: Path):
    sys.path.insert(0, str(author_root))
    from model.function.cost import SquaredError
    from model.function.network import Network
    from model.hopfield.minimizer import FixedPointMinimizer
    from model.hopfield.network import DeepHopfieldEnergy
    from training.sgd import AugmentedFunction, EquilibriumProp

    energy = DeepHopfieldEnergy([(4,), (7,), (3,)], [0.5, 0.5])
    energy.set_device("cpu")
    network = Network(energy)
    output = energy.layers()[-1]
    cost = SquaredError(output)
    augmented = AugmentedFunction(energy, cost)
    minimizer = FixedPointMinimizer(augmented, network.free_layers())
    minimizer.mode = "asynchronous"
    minimizer.num_iterations = 12
    estimator = EquilibriumProp(
        energy.params(), energy.layers(), augmented, cost, minimizer)
    estimator.variant = "positive"
    estimator.nudging = 0.25
    return energy, network, cost, augmented, minimizer, estimator


def free_state(network, cost, minimizer, augmented, x, labels):
    network.set_input(x, reset=True)
    cost.set_target(labels)
    augmented.nudging = 0.0
    minimizer.compute_equilibrium()
    return [layer.state.clone() for layer in minimizer._layers]


def restore(minimizer, states):
    for layer, state in zip(minimizer._layers, states):
        layer.state = state.clone()


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser()
    parser.add_argument("--author-root", type=Path, required=True)
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    torch.manual_seed(20260806)
    energy, network, cost, augmented, minimizer, estimator = build_estimator(
        args.author_root)
    x = torch.randn(8, 4)
    labels = torch.arange(8) % 3
    states = free_state(network, cost, minimizer, augmented, x, labels)
    clean = [value.clone() for value in estimator.compute_gradient()]
    assert all(not value.requires_grad for value in clean)

    restore(minimizer, states)
    oracle_corrector = RainGradientCorrector(
        mode="oracle", bias_ratio=0.5, seed=19)
    attach_to_rain_estimator(estimator, oracle_corrector)
    oracle = estimator.compute_gradient()
    assert all(torch.equal(a, b) for a, b in zip(clean, oracle))

    # Exercise the shared corrector on a sequence of local states.  The
    # structured field is exactly affine in its fixed local basis; innovation
    # should learn it while a constant filter retains state-dependent error.
    template_states = [torch.randn_like(value) for value in clean]
    innovation = RainGradientCorrector(
        mode="innovation", bias_ratio=0.5, predictor_rate=0.2, seed=31)
    constant = RainGradientCorrector(
        mode="constant", bias_ratio=0.5, predictor_rate=0.2, seed=31)
    zero = [torch.zeros_like(value) for value in clean]
    local_sequence = [
        [scale * value for value in template_states]
        for scale in torch.linspace(-1.2, 1.2, 50)
    ]
    generator = torch.Generator().manual_seed(1988)
    for local in local_sequence:
        innovation.observe_neutral(local)
    for _ in range(20):
        for index in torch.randperm(len(local_sequence), generator=generator):
            local = local_sequence[int(index)]
            innovation.apply(zero, local)
            constant.apply(zero, local)
    held = [1.45 * value for value in template_states]
    innovation.apply(zero, held)
    constant.apply(zero, held)
    innovation_error = innovation.last_diagnostics["residual_bias_rms"]
    constant_error = constant.last_diagnostics["residual_bias_rms"]
    assert innovation_error < 0.25 * constant_error, (
        innovation_error, constant_error)
    assert (
        innovation.debiaser.neutral_observations
        == constant.debiaser.neutral_observations + len(local_sequence)
    )
    print({
        "rain_revision_expected": RAIN_REVISION,
        "parameter_tensors": len(clean),
        "oracle_matches_clean_bitwise": True,
        "innovation_residual_bias_rms": innovation_error,
        "constant_residual_bias_rms": constant_error,
        "innovation_neutral_observations": innovation.debiaser.neutral_observations,
        "constant_neutral_observations": constant.debiaser.neutral_observations,
        "requires_grad": False,
    })


if __name__ == "__main__":
    main()