diff options
| -rw-r--r-- | RAIN_EP_DILLAVOU_MATRIX.md | 64 | ||||
| -rw-r--r-- | experiments/rain_ep_bias_train.py | 60 | ||||
| -rw-r--r-- | experiments/rain_ep_dillavou_smoke.py | 124 | ||||
| -rw-r--r-- | sdil/rain_ep_adapter.py | 215 |
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.""" |
