diff options
Diffstat (limited to 'sdil/rain_ep_adapter.py')
| -rw-r--r-- | sdil/rain_ep_adapter.py | 15 |
1 files changed, 12 insertions, 3 deletions
diff --git a/sdil/rain_ep_adapter.py b/sdil/rain_ep_adapter.py index 9be1128..3711a6a 100644 --- a/sdil/rain_ep_adapter.py +++ b/sdil/rain_ep_adapter.py @@ -236,6 +236,7 @@ class DillavouUpdateCorrector: mode: str, bias_ratio: float, predictor_rate: float = 1.0, + calibration_steps: int = 1, neutral_cadence: int = 1, drift_ratio: float = 0.0, seed: int = 1729, @@ -248,9 +249,12 @@ class DillavouUpdateCorrector: raise ValueError("predictor rate must lie in (0, 1]") if neutral_cadence < 0: raise ValueError("neutral cadence must be nonnegative") + if calibration_steps < 0: + raise ValueError("calibration steps must be nonnegative") self.mode = mode self.bias_ratio = bias_ratio self.predictor_rate = predictor_rate + self.calibration_steps = calibration_steps self.neutral_cadence = neutral_cadence self.drift_ratio = drift_ratio self.seed = seed @@ -365,10 +369,15 @@ class DillavouUpdateCorrector: noise.mul_(corruption.square().mean().sqrt()) corrected.append(value + noise) else: - if ( + calibrating = self.steps < self.calibration_steps + tracking = ( self.neutral_cadence > 0 - and self.steps % self.neutral_cadence == 0 - ): + and self.steps >= self.calibration_steps + and ( + self.steps - self.calibration_steps + ) % self.neutral_cadence == 0 + ) + if calibrating or tracking: self.debiaser.update_neutral( features, bias, self.predictor_rate) corrected = self.debiaser.residual(features, measured) |
