summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-07 12:57:26 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-07 12:57:26 -0500
commite824fc9ec147fc7b94b5493bc96e7a7127a33c01 (patch)
tree6b4b07e526df37e1f80901f0988b8dddc5bf556b
parent8cbca53997292045f1bee1e2bbda2c051eb6c578 (diff)
fix: freeze constant update calibration after one probe
-rw-r--r--experiments/rain_ep_bias_train.py3
-rw-r--r--experiments/rain_ep_dillavou_smoke.py10
-rw-r--r--sdil/rain_ep_adapter.py15
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)