diff options
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/rain_ep_bias_train.py | 36 | ||||
| -rw-r--r-- | experiments/rain_ep_layer_adapter_smoke.py | 139 |
2 files changed, 168 insertions, 7 deletions
diff --git a/experiments/rain_ep_bias_train.py b/experiments/rain_ep_bias_train.py index defd973..fa8a243 100644 --- a/experiments/rain_ep_bias_train.py +++ b/experiments/rain_ep_bias_train.py @@ -19,6 +19,8 @@ sys.path.insert(0, str(ROOT)) from sdil.rain_ep_adapter import ( # noqa: E402 RainGradientCorrector, + RainLayerStateCorrector, + attach_layer_to_rain_estimator, attach_to_rain_estimator, observe_rain_neutral, ) @@ -32,6 +34,8 @@ 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") + parser.add_argument( "--mode", choices=sorted(RainGradientCorrector.MODES), required=True) parser.add_argument("--bias-ratio", type=float, default=0.5) parser.add_argument("--predictor-rate", type=float, default=0.1) @@ -39,6 +43,7 @@ def parse_args() -> argparse.Namespace: "--neutral-cadence", type=int, default=1, help="training steps per neutral update; zero freezes after calibration") parser.add_argument("--calibration-batches", type=int, default=0) + parser.add_argument("--layer-calibration-steps", type=int, default=1) parser.add_argument("--epochs", type=int, default=2) parser.add_argument("--train-limit", type=int, default=2048) parser.add_argument("--test-limit", type=int, default=1024) @@ -139,13 +144,26 @@ def main() -> None: training_minimizer) estimator.variant = "positive" estimator.nudging = 0.25 - corrector = RainGradientCorrector( - mode=args.mode, - bias_ratio=args.bias_ratio, - predictor_rate=args.predictor_rate, - neutral_cadence=args.neutral_cadence, - seed=args.seed + 1729) - attach_to_rain_estimator(estimator, corrector) + if args.adapter == "parameter": + corrector = RainGradientCorrector( + mode=args.mode, + bias_ratio=args.bias_ratio, + predictor_rate=args.predictor_rate, + neutral_cadence=args.neutral_cadence, + seed=args.seed + 1729) + attach_to_rain_estimator(estimator, corrector) + else: + if args.calibration_batches: + raise ValueError( + "layer adapter calibrates inside existing free phases; " + "external calibration batches must be zero") + corrector = RainLayerStateCorrector( + mode=args.mode, + bias_ratio=args.bias_ratio, + predictor_rate=args.predictor_rate, + calibration_steps=args.layer_calibration_steps, + seed=args.seed + 1729) + attach_layer_to_rain_estimator(estimator, corrector) inference_minimizer = FixedPointMinimizer( energy, network.free_layers()) @@ -164,6 +182,8 @@ def main() -> None: calibration_start = time.time() calibration_observations = 0 if args.calibration_batches: + if args.adapter != "parameter": + raise AssertionError("layer calibration was not rejected above") if args.mode not in {"constant", "innovation"}: raise ValueError( "precalibration is defined only for constant or innovation mode") @@ -227,10 +247,12 @@ def main() -> None: "dataset": "FashionMNIST", "network": "author ConvHopfieldEnergy28 32-64-10", "algorithm": "positive equilibrium propagation", + "adapter": args.adapter, "mode": args.mode, "bias_ratio": args.bias_ratio, "predictor_rate": args.predictor_rate, "neutral_cadence": args.neutral_cadence, + "layer_calibration_steps": args.layer_calibration_steps, "calibration_batches": args.calibration_batches, "calibration_observations": calibration_observations, "calibration_seconds": calibration_seconds, diff --git a/experiments/rain_ep_layer_adapter_smoke.py b/experiments/rain_ep_layer_adapter_smoke.py new file mode 100644 index 0000000..8959217 --- /dev/null +++ b/experiments/rain_ep_layer_adapter_smoke.py @@ -0,0 +1,139 @@ +#!/usr/bin/env python3 +"""Strict-locality smoke test for the Rain neuron-state adapter.""" + +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 + RainLayerStateCorrector, + attach_layer_to_rain_estimator, +) + + +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 network, cost, augmented, minimizer, estimator + + +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) + network, cost, augmented, minimizer, estimator = build_estimator( + args.author_root) + x = torch.randn(64, 4) + labels = torch.arange(64) % 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() + oracle_corrector = RainLayerStateCorrector( + mode="oracle", bias_ratio=4.0, seed=19) + attach_layer_to_rain_estimator(estimator, oracle_corrector) + oracle = estimator.compute_gradient() + oracle_relative_error = max( + float((actual - target).norm() / target.norm().clamp_min(1e-30)) + for actual, target in zip(oracle, clean) + ) + assert oracle_relative_error < 2e-5, oracle_relative_error + + first = { + "input": torch.randn(64, 4), + "hidden": torch.randn(64, 7), + "output": torch.randn(64, 3), + } + clean_difference = { + name: 0.01 * torch.randn_like(value) + for name, value in first.items() + } + second = { + name: value + clean_difference[name] + for name, value in first.items() + } + innovation = RainLayerStateCorrector( + mode="innovation", bias_ratio=4.0, predictor_rate=0.2, + calibration_steps=1, seed=31) + constant = RainLayerStateCorrector( + mode="constant", bias_ratio=4.0, predictor_rate=0.2, + calibration_steps=1, seed=31) + layer_names = ["hidden", "output"] + innovation.apply(first, second, layer_names) + constant.apply(first, second, layer_names) + held_first = { + name: 1.25 * value + 0.1 for name, value in first.items() + } + held_clean = { + name: 0.01 * torch.randn_like(value) + for name, value in first.items() + } + held_second = { + name: value + held_clean[name] + for name, value in held_first.items() + } + innovation_used = innovation.apply( + held_first, held_second, layer_names) + constant_used = constant.apply(held_first, held_second, layer_names) + + def residual_rms(used): + errors = [ + used[name] - held_second[name] for name in layer_names + ] + return ( + sum(float(error.square().sum()) for error in errors) + / sum(error.numel() for error in errors) + ) ** 0.5 + + innovation_error = residual_rms(innovation_used) + constant_error = residual_rms(constant_used) + assert innovation_error < 0.1 * constant_error, ( + innovation_error, constant_error) + assert innovation.debiaser.neutral_observations == 64 + assert constant.debiaser.neutral_observations == 64 + assert all(not value.requires_grad for value in innovation_used.values()) + print({ + "oracle_parameter_gradient_relative_error": oracle_relative_error, + "innovation_heldout_state_residual_rms": innovation_error, + "constant_heldout_state_residual_rms": constant_error, + "matched_neutral_observations": 64, + "extra_equilibrium_phases": 0, + "requires_grad": False, + }) + + +if __name__ == "__main__": + main() |
