#!/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()