diff options
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/rain_ep_adapter_smoke.py | 124 |
1 files changed, 124 insertions, 0 deletions
diff --git a/experiments/rain_ep_adapter_smoke.py b/experiments/rain_ep_adapter_smoke.py new file mode 100644 index 0000000..8ad73b7 --- /dev/null +++ b/experiments/rain_ep_adapter_smoke.py @@ -0,0 +1,124 @@ +#!/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 _ 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 + 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, + "matched_neutral_observations": innovation.debiaser.neutral_observations, + "requires_grad": False, + }) + + +if __name__ == "__main__": + main() |
