diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-07 12:51:59 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-07 12:51:59 -0500 |
| commit | aea79bc3251f5e8a1ae798be97d773cf2639f46f (patch) | |
| tree | 3a1e6484a1a0ab5c934e15f34ef6d1467c94af08 /experiments | |
| parent | ffaa5c82b643b4313d288cdb9a4453254207f075 (diff) | |
feat: add Dillavou EP update-bias protocol
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/rain_ep_bias_train.py | 60 | ||||
| -rw-r--r-- | experiments/rain_ep_dillavou_smoke.py | 124 |
2 files changed, 173 insertions, 11 deletions
diff --git a/experiments/rain_ep_bias_train.py b/experiments/rain_ep_bias_train.py index db6f668..79e9d6b 100644 --- a/experiments/rain_ep_bias_train.py +++ b/experiments/rain_ep_bias_train.py @@ -19,8 +19,10 @@ ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) from sdil.rain_ep_adapter import ( # noqa: E402 + DillavouUpdateCorrector, RainGradientCorrector, RainLayerStateCorrector, + attach_dillavou_to_rain_estimator, attach_layer_to_rain_estimator, attach_to_rain_estimator, observe_rain_neutral, @@ -35,18 +37,23 @@ def parse_args() -> argparse.Namespace: parser.add_argument("--author-root", type=Path, required=True) parser.add_argument("--device", default="cuda") parser.add_argument( - "--adapter", choices=("parameter", "layer"), default="parameter") + "--adapter", choices=("parameter", "layer", "dillavou"), + default="parameter") parser.add_argument( "--network-protocol", choices=("conv28_screen", "comparative32"), default="conv28_screen") parser.add_argument( "--beta-policy", - choices=("fixed_positive", "fixed_negative", "random_sign"), + choices=("fixed_positive", "fixed_negative", "random_sign", "centered"), default="fixed_positive") parser.add_argument("--beta-seed", type=int, default=7100) + parser.add_argument("--beta-value", type=float, default=0.25) parser.add_argument( "--mode", choices=sorted(RainGradientCorrector.MODES), required=True) parser.add_argument("--bias-ratio", type=float, default=0.5) + parser.add_argument( + "--dillavou-drift-ratio", type=float, default=0.0, + help="zero is the exact fixed update-offset model from Dillavou et al.") parser.add_argument("--predictor-rate", type=float, default=0.1) parser.add_argument( "--neutral-cadence", type=int, default=1, @@ -103,11 +110,15 @@ def main() -> None: args = parse_args() if args.calibration_batches < 0: raise ValueError("calibration batches must be nonnegative") + if args.beta_value <= 0.0: + raise ValueError("beta value must be positive") if args.beta_policy != "fixed_positive" and not ( - args.adapter == "layer" and args.mode == "raw" + (args.adapter == "layer" and args.mode == "raw") + or args.adapter == "dillavou" ): raise ValueError( - "non-positive beta policies are raw layer-measurement baselines") + "non-positive beta policies require a raw layer baseline or " + "the post-estimator Dillavou adapter") author_root = args.author_root.resolve() author_revision = revision(author_root) if author_revision != PINNED_REVISION: @@ -185,8 +196,12 @@ def main() -> None: estimator = EquilibriumProp( energy.params(), energy.layers(), augmented, cost, training_minimizer) - estimator.variant = "positive" - estimator.nudging = 0.25 + estimator.variant = ( + "negative" if args.beta_policy == "fixed_negative" + else "centered" if args.beta_policy == "centered" + else "positive" + ) + estimator.nudging = args.beta_value if args.adapter == "parameter": corrector = RainGradientCorrector( mode=args.mode, @@ -195,7 +210,7 @@ def main() -> None: neutral_cadence=args.neutral_cadence, seed=args.seed + 1729) attach_to_rain_estimator(estimator, corrector) - else: + elif args.adapter == "layer": if args.calibration_batches: raise ValueError( "layer adapter calibrates inside existing free phases; " @@ -208,6 +223,20 @@ def main() -> None: bias_normalization=args.layer_bias_normalization, seed=args.seed + 1729) attach_layer_to_rain_estimator(estimator, corrector) + else: + if args.calibration_batches: + raise ValueError( + "Dillavou calibration probes the local update circuit and " + "does not require equilibrium batches") + corrector = DillavouUpdateCorrector( + mode=args.mode, + bias_ratio=args.bias_ratio, + predictor_rate=args.predictor_rate, + neutral_cadence=args.neutral_cadence, + drift_ratio=args.dillavou_drift_ratio, + seed=args.seed + 1729, + ) + attach_dillavou_to_rain_estimator(estimator, corrector) inference_minimizer = FixedPointMinimizer( energy, network.free_layers()) @@ -261,8 +290,9 @@ def main() -> None: if args.beta_policy == "fixed_positive": beta_sign_counts["positive"] += 1 elif args.beta_policy == "fixed_negative": - estimator._first_nudging = 0.0 - estimator._second_nudging = -estimator.nudging + beta_sign_counts["negative"] += 1 + elif args.beta_policy == "centered": + beta_sign_counts["positive"] += 1 beta_sign_counts["negative"] += 1 else: sign = 1 if int(torch.randint( @@ -316,9 +346,11 @@ def main() -> None: "algorithm": "equilibrium propagation", "beta_policy": args.beta_policy, "beta_seed": args.beta_seed, + "beta_value": args.beta_value, "adapter": args.adapter, "mode": args.mode, "bias_ratio": args.bias_ratio, + "dillavou_drift_ratio": args.dillavou_drift_ratio, "predictor_rate": args.predictor_rate, "neutral_cadence": args.neutral_cadence, "layer_calibration_steps": args.layer_calibration_steps, @@ -327,16 +359,22 @@ def main() -> None: "calibration_observations": calibration_observations, "calibration_seconds": calibration_seconds, "extra_equilibrium_phases_for_predictor": ( - 0 if args.adapter == "layer" else args.calibration_batches), + 0 if args.adapter in {"layer", "dillavou"} + else args.calibration_batches), "predictor_neutral_source": ( "existing_first_EP_phase" - if args.adapter == "layer" else "separate_free_equilibrium"), + if args.adapter == "layer" + else "instruction_off_local_update_probe" + if args.adapter == "dillavou" + else "separate_free_equilibrium"), "bias_ratio_normalization": ( ( "initial_free_layer_state_rms" if args.layer_bias_normalization == "first_state" else "experimenter_initial_clean_layer_state_difference_rms" ) if args.adapter == "layer" + else "initial_clean_local_update_rms_for_simulation_only" + if args.adapter == "dillavou" else "initial_local_parameter_state_rms"), "bias_ratio_normalization_visible_to_predictor": False, "epochs": args.epochs, diff --git a/experiments/rain_ep_dillavou_smoke.py b/experiments/rain_ep_dillavou_smoke.py new file mode 100644 index 0000000..d4a634a --- /dev/null +++ b/experiments/rain_ep_dillavou_smoke.py @@ -0,0 +1,124 @@ +#!/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 + DillavouUpdateCorrector, + attach_dillavou_to_rain_estimator, +) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--author-root", type=Path, required=True) + 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, seed=41) + innovation = DillavouUpdateCorrector( + mode="innovation", bias_ratio=0.2, predictor_rate=1.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 + + # 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, + "autodiff_used_for_learning": False, + }) + + +if __name__ == "__main__": + main() |
