summaryrefslogtreecommitdiff
path: root/sdil/rain_ep_adapter.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 16:37:12 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 16:37:12 -0500
commit00cd80855fb00bf8e06a76e5ba88a697acf92776 (patch)
tree56898cd08d2ddcc6efc3081c7ce9550863fcd00c /sdil/rain_ep_adapter.py
parent5d7adb2604f9f1e083523e71880533ac750c1cc4 (diff)
feat: integrate local debiaser with Rain EP
Diffstat (limited to 'sdil/rain_ep_adapter.py')
-rw-r--r--sdil/rain_ep_adapter.py189
1 files changed, 189 insertions, 0 deletions
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
+