From ddb034680ddc46d848278ca89d9bb97f59cbe3d3 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Thu, 6 Aug 2026 16:56:11 -0500 Subject: feat: add BP-free Rain layer-state adapter --- experiments/rain_ep_bias_train.py | 36 ++++-- experiments/rain_ep_layer_adapter_smoke.py | 139 +++++++++++++++++++++ sdil/rain_ep_adapter.py | 186 ++++++++++++++++++++++++++++- 3 files changed, 353 insertions(+), 8 deletions(-) create mode 100644 experiments/rain_ep_layer_adapter_smoke.py 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, ) @@ -31,6 +33,8 @@ def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() 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) @@ -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 -- cgit v1.2.3