summaryrefslogtreecommitdiff
path: root/experiments
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 /experiments
parent8cbca53997292045f1bee1e2bbda2c051eb6c578 (diff)
fix: freeze constant update calibration after one probe
Diffstat (limited to 'experiments')
-rw-r--r--experiments/rain_ep_bias_train.py3
-rw-r--r--experiments/rain_ep_dillavou_smoke.py10
2 files changed, 11 insertions, 2 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.