diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-07 12:57:26 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-07 12:57:26 -0500 |
| commit | e824fc9ec147fc7b94b5493bc96e7a7127a33c01 (patch) | |
| tree | 6b4b07e526df37e1f80901f0988b8dddc5bf556b | |
| parent | 8cbca53997292045f1bee1e2bbda2c051eb6c578 (diff) | |
fix: freeze constant update calibration after one probe
| -rw-r--r-- | experiments/rain_ep_bias_train.py | 3 | ||||
| -rw-r--r-- | experiments/rain_ep_dillavou_smoke.py | 10 | ||||
| -rw-r--r-- | sdil/rain_ep_adapter.py | 15 |
3 files changed, 23 insertions, 5 deletions
diff --git a/experiments/rain_ep_bias_train.py b/experiments/rain_ep_bias_train.py index 79e9d6b..fa62137 100644 --- a/experiments/rain_ep_bias_train.py +++ b/experiments/rain_ep_bias_train.py @@ -54,6 +54,7 @@ def parse_args() -> argparse.Namespace: parser.add_argument( "--dillavou-drift-ratio", type=float, default=0.0, help="zero is the exact fixed update-offset model from Dillavou et al.") + parser.add_argument("--dillavou-calibration-steps", type=int, default=1) parser.add_argument("--predictor-rate", type=float, default=0.1) parser.add_argument( "--neutral-cadence", type=int, default=1, @@ -207,6 +208,7 @@ def main() -> None: mode=args.mode, bias_ratio=args.bias_ratio, predictor_rate=args.predictor_rate, + calibration_steps=args.dillavou_calibration_steps, neutral_cadence=args.neutral_cadence, seed=args.seed + 1729) attach_to_rain_estimator(estimator, corrector) @@ -351,6 +353,7 @@ def main() -> None: "mode": args.mode, "bias_ratio": args.bias_ratio, "dillavou_drift_ratio": args.dillavou_drift_ratio, + "dillavou_calibration_steps": args.dillavou_calibration_steps, "predictor_rate": args.predictor_rate, "neutral_cadence": args.neutral_cadence, "layer_calibration_steps": args.layer_calibration_steps, diff --git a/experiments/rain_ep_dillavou_smoke.py b/experiments/rain_ep_dillavou_smoke.py index d4a634a..9519210 100644 --- a/experiments/rain_ep_dillavou_smoke.py +++ b/experiments/rain_ep_dillavou_smoke.py @@ -70,9 +70,11 @@ def main() -> None: assert fixed_relative_error < 2e-6, fixed_relative_error constant = DillavouUpdateCorrector( - mode="constant", bias_ratio=0.2, predictor_rate=1.0, seed=41) + mode="constant", bias_ratio=0.2, predictor_rate=1.0, + calibration_steps=1, neutral_cadence=0, seed=41) innovation = DillavouUpdateCorrector( - mode="innovation", bias_ratio=0.2, predictor_rate=1.0, seed=41) + mode="innovation", bias_ratio=0.2, predictor_rate=1.0, + calibration_steps=1, neutral_cadence=0, seed=41) corrected_constant = constant.apply(clean_a, parameters_a) corrected_innovation = innovation.apply(clean_a, parameters_a) constant_error = max( @@ -85,6 +87,10 @@ def main() -> None: ) assert constant_error < 2e-7, constant_error assert innovation_error < 2e-7, innovation_error + constant.apply(clean_b, parameters_b) + innovation.apply(clean_b, parameters_b) + assert constant.debiaser.neutral_observations == 1 + assert innovation.debiaser.neutral_observations == 1 # Integration check: the corruption is attached after Rain's hand-written # local EP estimator and introduces no autograd graph. diff --git a/sdil/rain_ep_adapter.py b/sdil/rain_ep_adapter.py index 9be1128..3711a6a 100644 --- a/sdil/rain_ep_adapter.py +++ b/sdil/rain_ep_adapter.py @@ -236,6 +236,7 @@ class DillavouUpdateCorrector: mode: str, bias_ratio: float, predictor_rate: float = 1.0, + calibration_steps: int = 1, neutral_cadence: int = 1, drift_ratio: float = 0.0, seed: int = 1729, @@ -248,9 +249,12 @@ class DillavouUpdateCorrector: raise ValueError("predictor rate must lie in (0, 1]") if neutral_cadence < 0: raise ValueError("neutral cadence must be nonnegative") + if calibration_steps < 0: + raise ValueError("calibration steps must be nonnegative") self.mode = mode self.bias_ratio = bias_ratio self.predictor_rate = predictor_rate + self.calibration_steps = calibration_steps self.neutral_cadence = neutral_cadence self.drift_ratio = drift_ratio self.seed = seed @@ -365,10 +369,15 @@ class DillavouUpdateCorrector: noise.mul_(corruption.square().mean().sqrt()) corrected.append(value + noise) else: - if ( + calibrating = self.steps < self.calibration_steps + tracking = ( self.neutral_cadence > 0 - and self.steps % self.neutral_cadence == 0 - ): + and self.steps >= self.calibration_steps + and ( + self.steps - self.calibration_steps + ) % self.neutral_cadence == 0 + ) + if calibrating or tracking: self.debiaser.update_neutral( features, bias, self.predictor_rate) corrected = self.debiaser.residual(features, measured) |
