summaryrefslogtreecommitdiff
path: root/experiments/rain_ep_layer_adapter_smoke.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/rain_ep_layer_adapter_smoke.py')
-rw-r--r--experiments/rain_ep_layer_adapter_smoke.py139
1 files changed, 139 insertions, 0 deletions
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()