From fa19cd0020a9c356a9f849f93d857cce040c1bb7 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Thu, 6 Aug 2026 16:44:51 -0500 Subject: feat: add counted Rain neutral calibration --- experiments/rain_ep_bias_train.py | 27 +++++++++++++++++++++++++++ 1 file changed, 27 insertions(+) (limited to 'experiments/rain_ep_bias_train.py') diff --git a/experiments/rain_ep_bias_train.py b/experiments/rain_ep_bias_train.py index 871d094..fcca67e 100644 --- a/experiments/rain_ep_bias_train.py +++ b/experiments/rain_ep_bias_train.py @@ -20,6 +20,7 @@ sys.path.insert(0, str(ROOT)) from sdil.rain_ep_adapter import ( # noqa: E402 RainGradientCorrector, attach_to_rain_estimator, + observe_rain_neutral, ) @@ -35,6 +36,7 @@ def parse_args() -> argparse.Namespace: parser.add_argument("--bias-ratio", type=float, default=0.5) parser.add_argument("--predictor-rate", type=float, default=0.1) parser.add_argument("--neutral-cadence", type=int, default=1) + parser.add_argument("--calibration-batches", type=int, default=0) parser.add_argument("--epochs", type=int, default=2) parser.add_argument("--train-limit", type=int, default=2048) parser.add_argument("--test-limit", type=int, default=1024) @@ -73,6 +75,8 @@ def evaluate(network, cost, minimizer, loader) -> tuple[float, float]: def main() -> None: args = parse_args() + if args.calibration_batches < 0: + raise ValueError("calibration batches must be nonnegative") author_root = args.author_root.resolve() author_revision = revision(author_root) if author_revision != PINNED_REVISION: @@ -103,6 +107,11 @@ def main() -> None: torch.utils.data.Subset(training_data, train_indices.tolist()), batch_size=args.batch_size, shuffle=True, generator=train_generator, num_workers=0) + calibration_generator = torch.Generator().manual_seed(args.seed + 104729) + calibration_loader = torch.utils.data.DataLoader( + torch.utils.data.Subset(training_data, train_indices.tolist()), + batch_size=args.batch_size, shuffle=True, + generator=calibration_generator, num_workers=0) test_loader = torch.utils.data.DataLoader( torch.utils.data.Subset(test_data, test_indices.tolist()), batch_size=args.batch_size, shuffle=False, num_workers=0) @@ -145,6 +154,21 @@ def main() -> None: metrics = [] start = time.time() + calibration_start = time.time() + calibration_observations = 0 + if args.calibration_batches: + if args.mode not in {"constant", "innovation"}: + raise ValueError( + "precalibration is defined only for constant or innovation mode") + while calibration_observations < args.calibration_batches: + for x, _ in calibration_loader: + network.set_input(x, reset=False) + inference_minimizer.compute_equilibrium() + observe_rain_neutral(estimator, corrector) + calibration_observations += 1 + if calibration_observations == args.calibration_batches: + break + calibration_seconds = time.time() - calibration_start for epoch in range(1, args.epochs + 1): total_cost = 0.0 total_correct = 0.0 @@ -193,6 +217,9 @@ def main() -> None: "bias_ratio": args.bias_ratio, "predictor_rate": args.predictor_rate, "neutral_cadence": args.neutral_cadence, + "calibration_batches": args.calibration_batches, + "calibration_observations": calibration_observations, + "calibration_seconds": calibration_seconds, "epochs": args.epochs, "train_limit": args.train_limit, "test_limit": args.test_limit, -- cgit v1.2.3