summaryrefslogtreecommitdiff
path: root/sdil/rain_ep_adapter.py
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 /sdil/rain_ep_adapter.py
parentd686f83aa35c2f1d87adfb5705ff02abe1901e9d (diff)
feat: add BP-free Rain layer-state adapter
Diffstat (limited to 'sdil/rain_ep_adapter.py')
-rw-r--r--sdil/rain_ep_adapter.py186
1 files changed, 185 insertions, 1 deletions
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