diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-06 17:19:02 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-06 17:19:02 -0500 |
| commit | 25f9b58f2fe478b7a4c728f404364ef8b8a92155 (patch) | |
| tree | fca17dc57dac2513cbbcd5ddb4cdacb35255f26e /sdil/rain_ep_adapter.py | |
| parent | 7ea36581a0469532630e891a032622c1f21a914b (diff) | |
feat: make Rain hardware bias beta-independent
Diffstat (limited to 'sdil/rain_ep_adapter.py')
| -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) ) |
