#!/usr/bin/env python3 """Mechanics checks for the post-estimator Dillavou update bias.""" 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 DillavouBiasProfile, DillavouUpdateCorrector, attach_dillavou_to_rain_estimator, ) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--author-root", type=Path, required=True) parser.add_argument( "--profile-json", type=Path, default=ROOT / "results/physical_bias/p0_state_dependence.json") return parser.parse_args() 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) cost = SquaredError(energy.layers()[-1]) 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 main() -> None: args = parse_args() torch.manual_seed(20260807) # The exact paper model has a fixed B_i after the estimator. Changing the # clean signal or parameter state must not change that field. clean_a = [torch.randn(11, 7), torch.randn(7)] clean_b = [torch.randn_like(value) for value in clean_a] parameters_a = [torch.randn_like(value) for value in clean_a] parameters_b = [value + 0.3 for value in parameters_a] raw = DillavouUpdateCorrector( mode="raw", bias_ratio=0.2, seed=41) measured_a = raw.apply(clean_a, parameters_a) measured_b = raw.apply(clean_b, parameters_b) bias_a = [value - clean for value, clean in zip(measured_a, clean_a)] bias_b = [value - clean for value, clean in zip(measured_b, clean_b)] fixed_relative_error = max( float((first - second).norm() / first.norm().clamp_min(1e-30)) for first, second in zip(bias_a, bias_b) ) assert fixed_relative_error < 2e-6, fixed_relative_error constant = DillavouUpdateCorrector( mode="constant", bias_ratio=0.2, predictor_rate=1.0, calibration_steps=1, neutral_cadence=0, seed=41) innovation = DillavouUpdateCorrector( mode="innovation", bias_ratio=0.2, predictor_rate=1.0, calibration_steps=1, neutral_cadence=0, seed=41) corrected_constant = constant.apply(clean_a, parameters_a) corrected_innovation = innovation.apply(clean_a, parameters_a) constant_error = max( float((actual - target).norm() / target.norm().clamp_min(1e-30)) for actual, target in zip(corrected_constant, clean_a) ) innovation_error = max( float((actual - target).norm() / target.norm().clamp_min(1e-30)) for actual, target in zip(corrected_innovation, clean_a) ) assert constant_error < 2e-7, constant_error assert innovation_error < 2e-7, innovation_error constant.apply(clean_b, parameters_b) innovation.apply(clean_b, parameters_b) assert constant.debiaser.neutral_observations == 1 assert innovation.debiaser.neutral_observations == 1 # Build the state-dependent shape only from the committed analysis of the # released physical traces. With matched neutral observations, an affine # predictor must generalize across local parameter states better than an # intercept-only predictor. import json profile = DillavouBiasProfile.from_state_dependence_report( json.loads(args.profile_json.read_text()), source=str(args.profile_json.resolve()), ) profile_constant = DillavouUpdateCorrector( mode="constant", bias_ratio=0.2, predictor_rate=0.2, calibration_steps=1, neutral_cadence=1, empirical_profile=profile, seed=67) profile_innovation = DillavouUpdateCorrector( mode="innovation", bias_ratio=0.2, predictor_rate=0.2, calibration_steps=1, neutral_cadence=1, empirical_profile=profile, seed=67) parameter_scale = [ value.square().mean().sqrt().clamp_min(1e-6) for value in parameters_a ] for _ in range(12): for displacement in torch.linspace(-1.0, 1.0, 21): state = [ value + displacement * scale for value, scale in zip(parameters_a, parameter_scale) ] profile_constant.apply(clean_a, state) profile_innovation.apply(clean_a, state) assert ( profile_constant.debiaser.neutral_observations == profile_innovation.debiaser.neutral_observations ) profile_constant.neutral_cadence = 0 profile_innovation.neutral_cadence = 0 held_state = [ value - 0.55 * scale for value, scale in zip(parameters_a, parameter_scale) ] held_constant = profile_constant.apply(clean_b, held_state) held_innovation = profile_innovation.apply(clean_b, held_state) held_constant_error = sum( float((actual - target).square().sum()) for actual, target in zip(held_constant, clean_b) ) held_innovation_error = sum( float((actual - target).square().sum()) for actual, target in zip(held_innovation, clean_b) ) assert held_innovation_error < 0.05 * held_constant_error, ( held_innovation_error, held_constant_error) # Two distinct local neutral states identify an exactly affine field when # the online sufficient-statistics predictor is selected. profile_ols = DillavouUpdateCorrector( mode="innovation", bias_ratio=0.2, predictor_rate=1.0, calibration_steps=2, neutral_cadence=0, empirical_profile=profile, predictor_kind="ols", seed=67) for displacement in (-0.25, 0.25): state = [ value + displacement * scale for value, scale in zip(parameters_a, parameter_scale) ] profile_ols.apply(clean_a, state) ols_state = [ value + 0.7 * scale for value, scale in zip(parameters_a, parameter_scale) ] held_ols = profile_ols.apply(clean_b, ols_state) held_ols_error = sum( float((actual - target).square().sum()) for actual, target in zip(held_ols, clean_b) ) assert held_ols_error < 1e-10, held_ols_error assert profile_ols.debiaser.neutral_observations == 2 # Integration check: the corruption is attached after Rain's hand-written # local EP estimator and introduces no autograd graph. energy, network, cost, augmented, minimizer, estimator = build_estimator( args.author_root) x = torch.randn(8, 4) labels = torch.arange(8) % 3 network.set_input(x, reset=True) cost.set_target(labels) augmented.nudging = 0.0 minimizer.compute_equilibrium() free = [layer.state.clone() for layer in minimizer._layers] clean = [value.clone() for value in estimator.compute_gradient()] for layer, state in zip(minimizer._layers, free): layer.state = state.clone() integrated = DillavouUpdateCorrector( mode="raw", bias_ratio=0.2, seed=53) attach_dillavou_to_rain_estimator(estimator, integrated) measured = estimator.compute_gradient() assert all(not value.requires_grad for value in measured) assert integrated.last_diagnostics["bias_model"] == ( "dillavou_constant_update") observed_ratio = integrated.last_diagnostics["bias_to_clean_update_rms"] assert abs(observed_ratio - 0.2) < 2e-6, observed_ratio print({ "fixed_bias_relative_error_after_state_change": fixed_relative_error, "constant_calibration_relative_error": constant_error, "innovation_relative_error": innovation_error, "integrated_bias_to_clean_update_rms": observed_ratio, "neutral_observations": constant.debiaser.neutral_observations, "released_profile_normalized_offsets": profile.normalized_offsets, "released_profile_normalized_state_variations": ( profile.normalized_state_variations), "released_profile_heldout_mse_ratio_affine_over_constant": ( held_innovation_error / held_constant_error), "released_profile_two_probe_ols_mse": held_ols_error, "autodiff_used_for_learning": False, }) if __name__ == "__main__": main()