summaryrefslogtreecommitdiff
path: root/sdil/rain_ep_adapter.py
diff options
context:
space:
mode:
Diffstat (limited to 'sdil/rain_ep_adapter.py')
-rw-r--r--sdil/rain_ep_adapter.py26
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(