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