summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--experiments/rain_ep_bias_train.py4
-rw-r--r--sdil/rain_ep_adapter.py9
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)