From 00cd80855fb00bf8e06a76e5ba88a697acf92776 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Thu, 6 Aug 2026 16:37:12 -0500 Subject: feat: integrate local debiaser with Rain EP --- experiments/rain_ep_adapter_smoke.py | 124 +++++++++++++++++++++++ sdil/rain_ep_adapter.py | 189 +++++++++++++++++++++++++++++++++++ 2 files changed, 313 insertions(+) create mode 100644 experiments/rain_ep_adapter_smoke.py create mode 100644 sdil/rain_ep_adapter.py diff --git a/experiments/rain_ep_adapter_smoke.py b/experiments/rain_ep_adapter_smoke.py new file mode 100644 index 0000000..8ad73b7 --- /dev/null +++ b/experiments/rain_ep_adapter_smoke.py @@ -0,0 +1,124 @@ +#!/usr/bin/env python3 +"""Integration smoke test against the pinned Rain EP implementation.""" + +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 + RainGradientCorrector, + attach_to_rain_estimator, +) + + +RAIN_REVISION = "6b253fd8a5d267535f58ab79992256ef10031ceb" + + +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) + output = energy.layers()[-1] + cost = SquaredError(output) + 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 free_state(network, cost, minimizer, augmented, x, labels): + network.set_input(x, reset=True) + cost.set_target(labels) + augmented.nudging = 0.0 + minimizer.compute_equilibrium() + return [layer.state.clone() for layer in minimizer._layers] + + +def restore(minimizer, states): + for layer, state in zip(minimizer._layers, states): + layer.state = state.clone() + + +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) + energy, network, cost, augmented, minimizer, estimator = build_estimator( + args.author_root) + x = torch.randn(8, 4) + labels = torch.arange(8) % 3 + states = free_state(network, cost, minimizer, augmented, x, labels) + clean = [value.clone() for value in estimator.compute_gradient()] + assert all(not value.requires_grad for value in clean) + + restore(minimizer, states) + oracle_corrector = RainGradientCorrector( + mode="oracle", bias_ratio=0.5, seed=19) + attach_to_rain_estimator(estimator, oracle_corrector) + oracle = estimator.compute_gradient() + assert all(torch.equal(a, b) for a, b in zip(clean, oracle)) + + # Exercise the shared corrector on a sequence of local states. The + # structured field is exactly affine in its fixed local basis; innovation + # should learn it while a constant filter retains state-dependent error. + template_states = [torch.randn_like(value) for value in clean] + innovation = RainGradientCorrector( + mode="innovation", bias_ratio=0.5, predictor_rate=0.2, seed=31) + constant = RainGradientCorrector( + mode="constant", bias_ratio=0.5, predictor_rate=0.2, seed=31) + zero = [torch.zeros_like(value) for value in clean] + local_sequence = [ + [scale * value for value in template_states] + for scale in torch.linspace(-1.2, 1.2, 50) + ] + generator = torch.Generator().manual_seed(1988) + for _ in range(20): + for index in torch.randperm(len(local_sequence), generator=generator): + local = local_sequence[int(index)] + innovation.apply(zero, local) + constant.apply(zero, local) + held = [1.45 * value for value in template_states] + innovation.apply(zero, held) + constant.apply(zero, held) + innovation_error = innovation.last_diagnostics["residual_bias_rms"] + constant_error = constant.last_diagnostics["residual_bias_rms"] + assert innovation_error < 0.25 * constant_error, ( + innovation_error, constant_error) + assert innovation.debiaser.neutral_observations == constant.debiaser.neutral_observations + print({ + "rain_revision_expected": RAIN_REVISION, + "parameter_tensors": len(clean), + "oracle_matches_clean_bitwise": True, + "innovation_residual_bias_rms": innovation_error, + "constant_residual_bias_rms": constant_error, + "matched_neutral_observations": innovation.debiaser.neutral_observations, + "requires_grad": False, + }) + + +if __name__ == "__main__": + main() diff --git a/sdil/rain_ep_adapter.py b/sdil/rain_ep_adapter.py new file mode 100644 index 0000000..81442aa --- /dev/null +++ b/sdil/rain_ep_adapter.py @@ -0,0 +1,189 @@ +"""No-autograd structured-bias adapter for the Rain EP implementation. + +The adapter monkey-patches only the estimator's two-state parameter-gradient +measurement. Rain's interaction objects already compute dense, convolutional +and bias energy derivatives by explicit local tensor operations, so no reverse +mode is introduced here. +""" + +from __future__ import annotations + +from types import MethodType +from typing import Iterable + +import torch + +from sdil.two_state_debias import LocalAffineDebiaser + + +Tensor = torch.Tensor + + +class LocalStructuredBias: + """Fixed per-element bias affine in a local first-state measurement.""" + + def __init__(self, ratio: float, seed: int = 1729) -> None: + if ratio < 0.0: + raise ValueError("bias ratio must be nonnegative") + self.ratio = ratio + self.seed = seed + self._metadata: list[tuple[Tensor, Tensor, Tensor]] | None = None + + @torch.no_grad() + def _initialize(self, local_states: list[Tensor]) -> None: + self._metadata = [] + for tensor_index, state in enumerate(local_states): + scale = state.square().mean().sqrt().clamp_min(1e-6) + flat_index = torch.arange( + state.numel(), dtype=state.dtype, device=state.device + ).reshape(state.shape) + phase = flat_index + float(self.seed + 97 * tensor_index) + offset = torch.where( + torch.remainder(phase, 2.0) < 1.0, + torch.full_like(state, -0.5), + torch.full_like(state, 0.5), + ) + slope = 0.5 + torch.remainder(phase * 0.61803398875, 1.0) + amplitude = self.ratio * scale + self._metadata.append((scale, offset, amplitude * slope)) + + @torch.no_grad() + def measure(self, local_states: Iterable[Tensor]) -> tuple[list[Tensor], list[Tensor]]: + local_states = list(local_states) + if any(state.requires_grad for state in local_states): + raise ValueError("local bias state must be detached") + if self._metadata is None: + self._initialize(local_states) + if len(local_states) != len(self._metadata): + raise ValueError("parameter collection changed after bias initialization") + bases = [] + biases = [] + for state, (scale, offset, scaled_slope) in zip( + local_states, self._metadata + ): + basis = torch.tanh(state / scale) + amplitude = self.ratio * scale + bias = amplitude * offset + scaled_slope * basis + bases.append(basis) + biases.append(bias) + return bases, biases + + +class RainGradientCorrector: + """Apply clean/raw/constant/SDIL/oracle/noise measurement policies.""" + + MODES = { + "clean", "raw", "constant", "innovation", "oracle", "same_rms_noise" + } + + def __init__( + self, + *, + mode: str, + bias_ratio: float, + predictor_rate: float = 0.05, + neutral_cadence: int = 1, + seed: int = 1729, + ) -> None: + if mode not in self.MODES: + raise ValueError(f"unrecognized correction mode {mode}") + if neutral_cadence < 1: + raise ValueError("neutral cadence must be positive") + self.mode = mode + self.predictor_rate = predictor_rate + self.neutral_cadence = neutral_cadence + self.bias = LocalStructuredBias(bias_ratio, seed) + self.debiaser: LocalAffineDebiaser | None = None + self.steps = 0 + self.last_diagnostics: dict[str, float | int] = {} + self._noise_generators: list[torch.Generator] | None = None + self.seed = seed + + @torch.no_grad() + def _initialize_debiaser(self, templates: list[Tensor]) -> None: + self.debiaser = LocalAffineDebiaser( + templates, + feature_centers=[0.0] * len(templates), + feature_scales=[1.0] * len(templates), + affine=self.mode == "innovation", + ) + + @torch.no_grad() + def apply(self, clean: Iterable[Tensor], local_states: Iterable[Tensor]) -> list[Tensor]: + clean = list(clean) + local_states = list(local_states) + if any(value.requires_grad for value in clean + local_states): + raise ValueError("Rain adapter received a requires-grad tensor") + if self.mode == "clean": + return [value.clone() for value in clean] + bases, bias = self.bias.measure(local_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.debiaser is None: + self._initialize_debiaser(clean) + if self.steps % self.neutral_cadence == 0: + 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 bias) + bias_square = sum(float(value.square().sum()) for value in bias) + residual_square = sum( + float(value.square().sum()) for value in residual_bias) + self.last_diagnostics = { + "step": self.steps, + "bias_rms": (bias_square / total_elements) ** 0.5, + "residual_bias_rms": (residual_square / total_elements) ** 0.5, + "neutral_observations": ( + 0 if self.debiaser is None else self.debiaser.neutral_observations + ), + } + self.steps += 1 + return corrected + + +def attach_to_rain_estimator(estimator, corrector: RainGradientCorrector): + """Replace Rain's standard two-state measurement with a corrected one.""" + + @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) + ] + return corrector.apply(clean, grads_first) + + estimator._standard_param_grads = MethodType( + corrected_standard_param_grads, estimator) + estimator.sdil_corrector = corrector + return estimator + -- cgit v1.2.3