summaryrefslogtreecommitdiff
path: root/sdil
diff options
context:
space:
mode:
Diffstat (limited to 'sdil')
-rw-r--r--sdil/rain_ep_adapter.py12
1 files changed, 10 insertions, 2 deletions
diff --git a/sdil/rain_ep_adapter.py b/sdil/rain_ep_adapter.py
index 73cd70d..ecdb5e5 100644
--- a/sdil/rain_ep_adapter.py
+++ b/sdil/rain_ep_adapter.py
@@ -234,6 +234,7 @@ class RainLayerStateCorrector:
bias_ratio: float,
predictor_rate: float = 0.2,
calibration_steps: int = 1,
+ bias_normalization: str = "clean_difference",
seed: int = 1729,
) -> None:
if mode not in self.MODES:
@@ -242,10 +243,13 @@ class RainLayerStateCorrector:
raise ValueError("bias ratio must be nonnegative")
if calibration_steps < 0:
raise ValueError("calibration steps must be nonnegative")
+ if bias_normalization not in {"clean_difference", "first_state"}:
+ raise ValueError("unrecognized layer-bias normalization")
self.mode = mode
self.bias_ratio = bias_ratio
self.predictor_rate = predictor_rate
self.calibration_steps = calibration_steps
+ self.bias_normalization = bias_normalization
self.seed = seed
self.steps = 0
self.last_diagnostics: dict[str, float | int] = {}
@@ -275,8 +279,12 @@ class RainLayerStateCorrector:
basis = torch.tanh(first / scale)
source = offset.unsqueeze(0) + slope.unsqueeze(0) * basis
source_rms = source.square().mean().sqrt().clamp_min(1e-30)
- clean_rms = clean.square().mean().sqrt()
- gain = self.bias_ratio * clean_rms / source_rms
+ reference_rms = (
+ clean.square().mean().sqrt()
+ if self.bias_normalization == "clean_difference"
+ else first.square().mean().sqrt()
+ )
+ gain = self.bias_ratio * reference_rms / source_rms
self._metadata.append(
(scale, gain * offset, gain * slope)
)