diff options
Diffstat (limited to 'sdil')
| -rw-r--r-- | sdil/rain_ep_adapter.py | 12 |
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) ) |
