summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--RAIN_EP_DILLAVOU_MATRIX.md64
-rw-r--r--experiments/rain_ep_bias_train.py60
-rw-r--r--experiments/rain_ep_dillavou_smoke.py124
-rw-r--r--sdil/rain_ep_adapter.py215
4 files changed, 452 insertions, 11 deletions
diff --git a/RAIN_EP_DILLAVOU_MATRIX.md b/RAIN_EP_DILLAVOU_MATRIX.md
new file mode 100644
index 0000000..3452ac2
--- /dev/null
+++ b/RAIN_EP_DILLAVOU_MATRIX.md
@@ -0,0 +1,64 @@
+# Rain EP Dillavou-Imperfection Matrix
+
+## Question
+
+Can a local neutral predictor remove a fixed hardware update offset without
+using a larger EP nudging voltage?
+
+The primary corruption is the model in Dillavou et al., Eq. (7). Every local
+parameter update receives a fixed, unknown, parameter-specific offset after
+the two-state EP estimate has been formed:
+
+\[
+g^{\rm measured}_i = g^{\rm EP}_i + B_i.
+\]
+
+`B_i` is fixed across examples, epochs, and beta signs. Its RMS is specified
+relative to the first clean local update only to set a reproducible simulation
+scale. This normalization is not visible to the learner. The optional
+state-drift ratio is zero in all exact-model experiments.
+
+The correction is local and uses no backpropagation. With the teaching input
+disabled, the update circuit exposes `B_i`. The predictor stores this neutral
+measurement and subtracts it from subsequent updates. Under the strictly
+constant model, constant calibration and SDIL innovation are expected to
+coincide. A difference between them is neither predicted nor claimed.
+
+## Frozen author protocol
+
+- author repository: `rain-neuromorphics/energy-based-learning`;
+- revision: `6b253fd8a5d267535f58ab79992256ef10031ceb`;
+- endpoint: `experiments/rain_ep_bias_train.py`;
+- network protocol: `comparative32`;
+- FashionMNIST with the author's 32x32 augmentation;
+- author ConvHopfieldEnergy32 architecture, gains and per-layer rates;
+- beta 0.25, 15 training relaxation iterations, 60 inference iterations;
+- batch size 100 and 100-epoch cosine schedule.
+
+## S0 development screen
+
+S0 uses only a fixed training/holdout split with 10,000 training and 2,000
+holdout examples. One epoch selects a non-catastrophic offset magnitude from
+the frozen grid `0.003, 0.01, 0.03, 0.1`. The same wave includes clean PEP,
+random-sign beta, centered EP, and one SDIL arm. S0 is development evidence
+and is never reported as a final result.
+
+## Confirmation matrix
+
+After S0 freezes one offset magnitude, the full author horizon compares:
+
+| EP estimator | no offset | raw offset | constant calibration | SDIL |
+|---|---:|---:|---:|---:|
+| positive beta | yes | yes | yes | yes |
+| random beta sign | yes | yes | no | no |
+| centered EP | yes | yes | no | no |
+
+An oracle subtraction arm checks implementation correctness. A large-beta
+EP sweep is reported as a strong-clamp proxy but is not called overclamping:
+Dillavou overclamping also changes the output force and update duration, so a
+beta sweep alone is not the published method.
+
+The exact constant model establishes the hardware failure and the limit of
+beta centering. It cannot establish an advantage over ordinary local offset
+calibration. Any claimed SDIL advantage requires a separately labeled
+state-dependent extension or real measured device drift.
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()
diff --git a/sdil/rain_ep_adapter.py b/sdil/rain_ep_adapter.py
index ecdb5e5..9be1128 100644
--- a/sdil/rain_ep_adapter.py
+++ b/sdil/rain_ep_adapter.py
@@ -215,6 +215,221 @@ def attach_to_rain_estimator(estimator, corrector: RainGradientCorrector):
return estimator
+class DillavouUpdateCorrector:
+ """Fixed per-parameter update bias from Dillavou et al. Eq. (7).
+
+ The corruption is added *after* the two-state EP estimator has formed its
+ local parameter update. It is therefore independent of the sign or
+ centering of beta. ``constant`` and ``innovation`` intentionally coincide
+ when ``drift_ratio`` is zero: under the paper's strictly constant bias
+ model, the Harnett-style predictor reduces to a local intercept estimate.
+
+ ``drift_ratio`` is an optional state-dependent extension. It is kept
+ separate from the exact Dillavou model so reports cannot conflate them.
+ """
+
+ MODES = RainGradientCorrector.MODES
+
+ def __init__(
+ self,
+ *,
+ mode: str,
+ bias_ratio: float,
+ predictor_rate: float = 1.0,
+ neutral_cadence: int = 1,
+ drift_ratio: float = 0.0,
+ seed: int = 1729,
+ ) -> None:
+ if mode not in self.MODES:
+ raise ValueError(f"unrecognized correction mode {mode}")
+ if bias_ratio < 0.0 or drift_ratio < 0.0:
+ raise ValueError("Dillavou bias ratios must be nonnegative")
+ if not 0.0 < predictor_rate <= 1.0:
+ raise ValueError("predictor rate must lie in (0, 1]")
+ if neutral_cadence < 0:
+ raise ValueError("neutral cadence must be nonnegative")
+ self.mode = mode
+ self.bias_ratio = bias_ratio
+ self.predictor_rate = predictor_rate
+ self.neutral_cadence = neutral_cadence
+ self.drift_ratio = drift_ratio
+ self.seed = seed
+ self.steps = 0
+ self.debiaser: LocalAffineDebiaser | None = None
+ self._offsets: list[Tensor] | None = None
+ self._slopes: list[Tensor] | None = None
+ self._centers: list[Tensor] | None = None
+ self._scales: list[Tensor] | None = None
+ self._noise_generators: list[torch.Generator] | None = None
+ self.last_diagnostics: dict[str, float | int | str] = {}
+
+ @staticmethod
+ @torch.no_grad()
+ def _unit_pattern(template: Tensor, phase: float) -> Tensor:
+ index = torch.arange(
+ template.numel(), dtype=template.dtype, device=template.device
+ ).reshape(template.shape)
+ pattern = torch.sin((index + phase) * 1.61803398875)
+ pattern.add_(0.5 * torch.cos((index + phase) * 0.754877666))
+ return pattern / pattern.square().mean().sqrt().clamp_min(1e-30)
+
+ @torch.no_grad()
+ def _initialize(
+ self, clean: list[Tensor], parameter_states: list[Tensor]
+ ) -> None:
+ if len(clean) != len(parameter_states):
+ raise ValueError("gradient and parameter collections disagree")
+ self._offsets = []
+ self._slopes = []
+ self._centers = []
+ self._scales = []
+ feature_scales = []
+ for index, (gradient, parameter) in enumerate(
+ zip(clean, parameter_states)
+ ):
+ gradient_scale = gradient.square().mean().sqrt().clamp_min(1e-12)
+ parameter_scale = parameter.square().mean().sqrt().clamp_min(1e-6)
+ offset_pattern = self._unit_pattern(
+ gradient, float(self.seed + 97 * index))
+ slope_pattern = self._unit_pattern(
+ gradient, float(self.seed + 193 * index + 41))
+ self._offsets.append(
+ self.bias_ratio * gradient_scale * offset_pattern)
+ self._slopes.append(
+ self.drift_ratio * gradient_scale * slope_pattern)
+ self._centers.append(parameter.clone())
+ self._scales.append(parameter_scale.clone())
+ feature_scales.append(torch.ones_like(parameter_scale))
+ self.debiaser = LocalAffineDebiaser(
+ clean,
+ feature_centers=[0.0] * len(clean),
+ feature_scales=feature_scales,
+ affine=self.mode == "innovation",
+ )
+
+ @torch.no_grad()
+ def _measure(
+ self, parameter_states: list[Tensor]
+ ) -> tuple[list[Tensor], list[Tensor]]:
+ features = []
+ biases = []
+ for parameter, center, scale, offset, slope in zip(
+ parameter_states,
+ self._centers,
+ self._scales,
+ self._offsets,
+ self._slopes,
+ ):
+ feature = torch.tanh((parameter - center) / scale)
+ features.append(feature)
+ biases.append(offset + slope * feature)
+ return features, biases
+
+ @torch.no_grad()
+ def apply(
+ self, clean: Iterable[Tensor], parameter_states: Iterable[Tensor]
+ ) -> list[Tensor]:
+ clean = list(clean)
+ parameter_states = list(parameter_states)
+ if any(value.requires_grad for value in clean + parameter_states):
+ raise ValueError("Dillavou adapter received a requires-grad tensor")
+ if self.mode == "clean":
+ return [value.clone() for value in clean]
+ if self._offsets is None:
+ self._initialize(clean, parameter_states)
+ features, bias = self._measure(parameter_states)
+ 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.neutral_cadence > 0
+ and self.steps % self.neutral_cadence == 0
+ ):
+ self.debiaser.update_neutral(
+ features, bias, self.predictor_rate)
+ corrected = self.debiaser.residual(features, 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,
+ "bias_model": (
+ "dillavou_constant_update"
+ if self.drift_ratio == 0.0
+ else "dillavou_plus_local_state_drift"
+ ),
+ "clean_update_rms": clean_rms,
+ "bias_update_rms": bias_rms,
+ "residual_update_rms": residual_rms,
+ "bias_to_clean_update_rms": (
+ bias_rms / max(clean_rms, 1e-30)),
+ "residual_to_clean_update_rms": (
+ residual_rms / max(clean_rms, 1e-30)),
+ "neutral_observations": self.debiaser.neutral_observations,
+ }
+ self.steps += 1
+ return corrected
+
+
+def attach_dillavou_to_rain_estimator(
+ estimator, corrector: DillavouUpdateCorrector
+):
+ """Add the hardware update offset after Rain's local EP estimate."""
+
+ @torch.no_grad()
+ def corrected_standard_param_grads(self, layers_first, layers_second):
+ 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 = layers_second[layer.name]
+ grads_second = [updater.grad() for updater in self._param_updaters]
+ denominator = self._second_nudging - self._first_nudging
+ clean = [
+ (second - first) / denominator
+ for first, second in zip(grads_first, grads_second)
+ ]
+ parameter_states = [parameter.state for parameter in self._params]
+ return corrector.apply(clean, parameter_states)
+
+ estimator._standard_param_grads = MethodType(
+ corrected_standard_param_grads, estimator)
+ estimator.sdil_corrector = corrector
+ return estimator
+
+
@torch.no_grad()
def observe_rain_neutral(estimator, corrector: RainGradientCorrector) -> None:
"""Expose one label-free Rain equilibrium to the local predictor."""