diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-06 16:44:51 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-06 16:44:51 -0500 |
| commit | fa19cd0020a9c356a9f849f93d857cce040c1bb7 (patch) | |
| tree | 7fa10922891f1689f767f669e23f5dfb68d33640 /experiments/rain_ep_adapter_smoke.py | |
| parent | 0be0a72cc0a82afc1b6811f95e87a0dfc1927a49 (diff) | |
feat: add counted Rain neutral calibration
Diffstat (limited to 'experiments/rain_ep_adapter_smoke.py')
| -rw-r--r-- | experiments/rain_ep_adapter_smoke.py | 10 |
1 files changed, 8 insertions, 2 deletions
diff --git a/experiments/rain_ep_adapter_smoke.py b/experiments/rain_ep_adapter_smoke.py index 8ad73b7..425cc26 100644 --- a/experiments/rain_ep_adapter_smoke.py +++ b/experiments/rain_ep_adapter_smoke.py @@ -96,6 +96,8 @@ def main() -> None: for scale in torch.linspace(-1.2, 1.2, 50) ] generator = torch.Generator().manual_seed(1988) + for local in local_sequence: + innovation.observe_neutral(local) for _ in range(20): for index in torch.randperm(len(local_sequence), generator=generator): local = local_sequence[int(index)] @@ -108,14 +110,18 @@ def main() -> None: constant_error = constant.last_diagnostics["residual_bias_rms"] assert innovation_error < 0.25 * constant_error, ( innovation_error, constant_error) - assert innovation.debiaser.neutral_observations == constant.debiaser.neutral_observations + assert ( + innovation.debiaser.neutral_observations + == constant.debiaser.neutral_observations + len(local_sequence) + ) print({ "rain_revision_expected": RAIN_REVISION, "parameter_tensors": len(clean), "oracle_matches_clean_bitwise": True, "innovation_residual_bias_rms": innovation_error, "constant_residual_bias_rms": constant_error, - "matched_neutral_observations": innovation.debiaser.neutral_observations, + "innovation_neutral_observations": innovation.debiaser.neutral_observations, + "constant_neutral_observations": constant.debiaser.neutral_observations, "requires_grad": False, }) |
