diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-10 04:31:50 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-10 04:31:50 -0500 |
| commit | 927241c496d2508344f463c9cf96b280b5cc2830 (patch) | |
| tree | f2a8bb017f305d57581103c71345a47ba09ca55f /sdil/rain_ep_adapter.py | |
| parent | 29808ed4f8d550eb0ddcbc300b33f2dc596721e3 (diff) | |
feat: add local online least-squares residual predictor
Diffstat (limited to 'sdil/rain_ep_adapter.py')
| -rw-r--r-- | sdil/rain_ep_adapter.py | 26 |
1 files changed, 19 insertions, 7 deletions
diff --git a/sdil/rain_ep_adapter.py b/sdil/rain_ep_adapter.py index b16e96d..7dd47eb 100644 --- a/sdil/rain_ep_adapter.py +++ b/sdil/rain_ep_adapter.py @@ -18,6 +18,7 @@ import torch from sdil.two_state_debias import ( BatchedLocalAffineDebiaser, LocalAffineDebiaser, + LocalLeastSquaresDebiaser, ) @@ -306,6 +307,7 @@ class DillavouUpdateCorrector: neutral_cadence: int = 1, drift_ratio: float = 0.0, empirical_profile: DillavouBiasProfile | None = None, + predictor_kind: str = "nlms", seed: int = 1729, ) -> None: if mode not in self.MODES: @@ -321,6 +323,8 @@ class DillavouUpdateCorrector: raise ValueError("neutral cadence must be nonnegative") if calibration_steps < 0: raise ValueError("calibration steps must be nonnegative") + if predictor_kind not in {"nlms", "ols"}: + raise ValueError("predictor kind must be nlms or ols") self.mode = mode self.bias_ratio = bias_ratio self.predictor_rate = predictor_rate @@ -328,9 +332,10 @@ class DillavouUpdateCorrector: self.neutral_cadence = neutral_cadence self.drift_ratio = drift_ratio self.empirical_profile = empirical_profile + self.predictor_kind = predictor_kind self.seed = seed self.steps = 0 - self.debiaser: LocalAffineDebiaser | None = None + self.debiaser: LocalAffineDebiaser | LocalLeastSquaresDebiaser | None = None self._offsets: list[Tensor] | None = None self._slopes: list[Tensor] | None = None self._centers: list[Tensor] | None = None @@ -402,12 +407,19 @@ class DillavouUpdateCorrector: 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", - ) + if self.mode == "innovation" and self.predictor_kind == "ols": + self.debiaser = LocalLeastSquaresDebiaser( + clean, + feature_centers=[0.0] * len(clean), + feature_scales=feature_scales, + ) + else: + self.debiaser = LocalAffineDebiaser( + clean, + feature_centers=[0.0] * len(clean), + feature_scales=feature_scales, + affine=self.mode == "innovation", + ) @torch.no_grad() def _measure( |
