summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 16:56:11 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 16:56:11 -0500
commitddb034680ddc46d848278ca89d9bb97f59cbe3d3 (patch)
tree9e53bbac894717f172948607a63afd26327b2262
parentd686f83aa35c2f1d87adfb5705ff02abe1901e9d (diff)
feat: add BP-free Rain layer-state adapter
-rw-r--r--experiments/rain_ep_bias_train.py36
-rw-r--r--experiments/rain_ep_layer_adapter_smoke.py139
-rw-r--r--sdil/rain_ep_adapter.py186
3 files changed, 353 insertions, 8 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()
diff --git a/sdil/rain_ep_adapter.py b/sdil/rain_ep_adapter.py
index 27314bd..73cd70d 100644
--- a/sdil/rain_ep_adapter.py
+++ b/sdil/rain_ep_adapter.py
@@ -13,7 +13,10 @@ from typing import Iterable
import torch
-from sdil.two_state_debias import LocalAffineDebiaser
+from sdil.two_state_debias import (
+ BatchedLocalAffineDebiaser,
+ LocalAffineDebiaser,
+)
Tensor = torch.Tensor
@@ -217,3 +220,184 @@ def observe_rain_neutral(estimator, corrector: RainGradientCorrector) -> None:
"""Expose one label-free Rain equilibrium to the local predictor."""
local_states = [updater.grad() for updater in estimator._param_updaters]
corrector.observe_neutral(local_states)
+
+
+class RainLayerStateCorrector:
+ """Correct a structured per-neuron bias before Rain's local EP update."""
+
+ MODES = RainGradientCorrector.MODES
+
+ def __init__(
+ self,
+ *,
+ mode: str,
+ bias_ratio: float,
+ predictor_rate: float = 0.2,
+ calibration_steps: int = 1,
+ seed: int = 1729,
+ ) -> None:
+ if mode not in self.MODES:
+ raise ValueError(f"unrecognized correction mode {mode}")
+ if bias_ratio < 0.0:
+ raise ValueError("bias ratio must be nonnegative")
+ if calibration_steps < 0:
+ raise ValueError("calibration steps must be nonnegative")
+ self.mode = mode
+ self.bias_ratio = bias_ratio
+ self.predictor_rate = predictor_rate
+ self.calibration_steps = calibration_steps
+ self.seed = seed
+ self.steps = 0
+ self.last_diagnostics: dict[str, float | int] = {}
+ self.debiaser: BatchedLocalAffineDebiaser | None = None
+ self._metadata: list[tuple[Tensor, Tensor, Tensor]] | None = None
+ self._noise_generators: list[torch.Generator] | None = None
+
+ @torch.no_grad()
+ def _initialize(
+ self, first_states: list[Tensor], clean_differences: list[Tensor]
+ ) -> None:
+ self._metadata = []
+ for index, (first, clean) in enumerate(zip(
+ first_states, clean_differences
+ )):
+ scale = first.square().mean(dim=0).sqrt().clamp_min(1e-6)
+ flat_index = torch.arange(
+ scale.numel(), dtype=first.dtype, device=first.device
+ ).reshape(scale.shape)
+ phase = flat_index + float(self.seed + 97 * index)
+ offset = torch.where(
+ torch.remainder(phase, 2.0) < 1.0,
+ torch.full_like(scale, -0.5),
+ torch.full_like(scale, 0.5),
+ )
+ slope = 0.5 + torch.remainder(phase * 0.61803398875, 1.0)
+ basis = torch.tanh(first / scale)
+ source = offset.unsqueeze(0) + slope.unsqueeze(0) * basis
+ source_rms = source.square().mean().sqrt().clamp_min(1e-30)
+ clean_rms = clean.square().mean().sqrt()
+ gain = self.bias_ratio * clean_rms / source_rms
+ self._metadata.append(
+ (scale, gain * offset, gain * slope)
+ )
+ self.debiaser = BatchedLocalAffineDebiaser(
+ first_states,
+ feature_centers=[0.0] * len(first_states),
+ feature_scales=[1.0] * len(first_states),
+ affine=self.mode == "innovation",
+ )
+
+ @torch.no_grad()
+ def _measure(
+ self, first_states: list[Tensor]
+ ) -> tuple[list[Tensor], list[Tensor]]:
+ bases = []
+ biases = []
+ for first, (scale, offset, slope) in zip(
+ first_states, self._metadata
+ ):
+ basis = torch.tanh(first / scale)
+ bases.append(basis)
+ biases.append(
+ offset.unsqueeze(0) + slope.unsqueeze(0) * basis)
+ return bases, biases
+
+ @torch.no_grad()
+ def apply(
+ self,
+ layers_first: dict[str, Tensor],
+ layers_second: dict[str, Tensor],
+ layer_names: list[str],
+ ) -> dict[str, Tensor]:
+ if self.mode == "clean":
+ return dict(layers_second)
+ first = [layers_first[name] for name in layer_names]
+ second = [layers_second[name] for name in layer_names]
+ if any(value.requires_grad for value in first + second):
+ raise ValueError("Rain layer adapter received a requires-grad tensor")
+ clean = [after - before for before, after in zip(first, second)]
+ if self._metadata is None:
+ self._initialize(first, clean)
+ bases, bias = self._measure(first)
+ measured = [value + corruption for value, corruption in zip(clean, bias)]
+ if self.mode == "raw":
+ corrected = measured
+ elif self.mode == "oracle":
+ corrected = [value.clone() for value in clean]
+ elif self.mode == "same_rms_noise":
+ if self._noise_generators is None:
+ self._noise_generators = []
+ for index, template in enumerate(clean):
+ generator = torch.Generator(device=template.device)
+ generator.manual_seed(self.seed + 1009 * index)
+ self._noise_generators.append(generator)
+ corrected = []
+ for value, corruption, generator in zip(
+ clean, bias, self._noise_generators
+ ):
+ noise = torch.randn(
+ value.shape, dtype=value.dtype, device=value.device,
+ generator=generator)
+ noise.mul_(corruption.square().mean().sqrt())
+ corrected.append(value + noise)
+ else:
+ if self.steps < self.calibration_steps:
+ self.debiaser.update_neutral(
+ bases, bias, self.predictor_rate)
+ corrected = self.debiaser.residual(bases, measured)
+ residual_bias = [
+ value - target for value, target in zip(corrected, clean)
+ ]
+ total_elements = sum(value.numel() for value in clean)
+ clean_square = sum(float(value.square().sum()) for value in clean)
+ bias_square = sum(float(value.square().sum()) for value in bias)
+ residual_square = sum(
+ float(value.square().sum()) for value in residual_bias)
+ clean_rms = (clean_square / total_elements) ** 0.5
+ bias_rms = (bias_square / total_elements) ** 0.5
+ residual_rms = (residual_square / total_elements) ** 0.5
+ self.last_diagnostics = {
+ "step": self.steps,
+ "clean_state_difference_rms": clean_rms,
+ "bias_state_difference_rms": bias_rms,
+ "residual_state_difference_rms": residual_rms,
+ "bias_to_clean_state_difference_rms": (
+ bias_rms / max(clean_rms, 1e-30)),
+ "residual_to_clean_state_difference_rms": (
+ residual_rms / max(clean_rms, 1e-30)),
+ "neutral_observations": (
+ 0 if self.debiaser is None else self.debiaser.neutral_observations),
+ }
+ used_second = dict(layers_second)
+ for name, before, difference in zip(layer_names, first, corrected):
+ used_second[name] = before + difference
+ self.steps += 1
+ return used_second
+
+
+def attach_layer_to_rain_estimator(
+ estimator, corrector: RainLayerStateCorrector
+):
+ """Patch Rain immediately before its hand-written local parameter rule."""
+ layer_names = [layer.name for layer in estimator._layers[1:]]
+
+ @torch.no_grad()
+ def corrected_standard_param_grads(self, layers_first, layers_second):
+ used_second = corrector.apply(
+ layers_first, layers_second, layer_names)
+ for layer in self._layers:
+ layer.state = layers_first[layer.name]
+ grads_first = [updater.grad() for updater in self._param_updaters]
+ for layer in self._layers:
+ layer.state = used_second[layer.name]
+ grads_second = [updater.grad() for updater in self._param_updaters]
+ denominator = self._second_nudging - self._first_nudging
+ return [
+ (second - first) / denominator
+ for first, second in zip(grads_first, grads_second)
+ ]
+
+ estimator._standard_param_grads = MethodType(
+ corrected_standard_param_grads, estimator)
+ estimator.sdil_corrector = corrector
+ return estimator