From fa19cd0020a9c356a9f849f93d857cce040c1bb7 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Thu, 6 Aug 2026 16:44:51 -0500 Subject: feat: add counted Rain neutral calibration --- sdil/rain_ep_adapter.py | 20 ++++++++++++++++++++ 1 file changed, 20 insertions(+) (limited to 'sdil') diff --git a/sdil/rain_ep_adapter.py b/sdil/rain_ep_adapter.py index 81442aa..9c7ab8f 100644 --- a/sdil/rain_ep_adapter.py +++ b/sdil/rain_ep_adapter.py @@ -108,6 +108,20 @@ class RainGradientCorrector: affine=self.mode == "innovation", ) + @torch.no_grad() + def observe_neutral(self, local_states: Iterable[Tensor]) -> None: + """Fit the local bias field from one instruction-off observation.""" + if self.mode not in {"constant", "innovation"}: + raise ValueError( + "neutral predictor observations require constant or innovation mode") + local_states = list(local_states) + if any(value.requires_grad for value in local_states): + raise ValueError("Rain adapter received a requires-grad tensor") + bases, bias = self.bias.measure(local_states) + if self.debiaser is None: + self._initialize_debiaser(local_states) + self.debiaser.update_neutral(bases, bias, self.predictor_rate) + @torch.no_grad() def apply(self, clean: Iterable[Tensor], local_states: Iterable[Tensor]) -> list[Tensor]: clean = list(clean) @@ -187,3 +201,9 @@ def attach_to_rain_estimator(estimator, corrector: RainGradientCorrector): 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.""" + local_states = [updater.grad() for updater in estimator._param_updaters] + corrector.observe_neutral(local_states) -- cgit v1.2.3