diff options
| -rw-r--r-- | experiments/rain_ep_bias_train.py | 4 | ||||
| -rw-r--r-- | sdil/rain_ep_adapter.py | 9 |
2 files changed, 9 insertions, 4 deletions
diff --git a/experiments/rain_ep_bias_train.py b/experiments/rain_ep_bias_train.py index 8b6e9ad..7bdff26 100644 --- a/experiments/rain_ep_bias_train.py +++ b/experiments/rain_ep_bias_train.py @@ -35,7 +35,9 @@ def parse_args() -> argparse.Namespace: "--mode", choices=sorted(RainGradientCorrector.MODES), required=True) parser.add_argument("--bias-ratio", type=float, default=0.5) parser.add_argument("--predictor-rate", type=float, default=0.1) - parser.add_argument("--neutral-cadence", type=int, default=1) + parser.add_argument( + "--neutral-cadence", type=int, default=1, + help="training steps per neutral update; zero freezes after calibration") parser.add_argument("--calibration-batches", type=int, default=0) parser.add_argument("--epochs", type=int, default=2) parser.add_argument("--train-limit", type=int, default=2048) diff --git a/sdil/rain_ep_adapter.py b/sdil/rain_ep_adapter.py index bcdc410..27314bd 100644 --- a/sdil/rain_ep_adapter.py +++ b/sdil/rain_ep_adapter.py @@ -87,8 +87,8 @@ class RainGradientCorrector: ) -> None: if mode not in self.MODES: raise ValueError(f"unrecognized correction mode {mode}") - if neutral_cadence < 1: - raise ValueError("neutral cadence must be positive") + if neutral_cadence < 0: + raise ValueError("neutral cadence must be nonnegative") self.mode = mode self.predictor_rate = predictor_rate self.neutral_cadence = neutral_cadence @@ -155,7 +155,10 @@ class RainGradientCorrector: else: if self.debiaser is None: self._initialize_debiaser(clean) - if self.steps % self.neutral_cadence == 0: + if ( + self.neutral_cadence > 0 + and self.steps % self.neutral_cadence == 0 + ): self.debiaser.update_neutral( bases, bias, self.predictor_rate) corrected = self.debiaser.residual(bases, measured) |
