summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rw-r--r--experiments/rain_ep_bias_train.py60
-rw-r--r--experiments/rain_ep_dillavou_smoke.py124
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()