From 25f9b58f2fe478b7a4c728f404364ef8b8a92155 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Thu, 6 Aug 2026 17:19:02 -0500 Subject: feat: make Rain hardware bias beta-independent --- sdil/rain_ep_adapter.py | 12 ++++++++++-- 1 file changed, 10 insertions(+), 2 deletions(-) (limited to 'sdil/rain_ep_adapter.py') 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) ) -- cgit v1.2.3