summaryrefslogtreecommitdiff
path: root/sdil/rain_ep_adapter.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-07 12:51:59 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-07 12:51:59 -0500
commitaea79bc3251f5e8a1ae798be97d773cf2639f46f (patch)
tree3a1e6484a1a0ab5c934e15f34ef6d1467c94af08 /sdil/rain_ep_adapter.py
parentffaa5c82b643b4313d288cdb9a4453254207f075 (diff)
feat: add Dillavou EP update-bias protocol
Diffstat (limited to 'sdil/rain_ep_adapter.py')
-rw-r--r--sdil/rain_ep_adapter.py215
1 files changed, 215 insertions, 0 deletions
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."""