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 | |
| parent | 0be0a72cc0a82afc1b6811f95e87a0dfc1927a49 (diff) | |
feat: add counted Rain neutral calibration
| -rw-r--r-- | experiments/rain_ep_adapter_smoke.py | 10 | ||||
| -rw-r--r-- | experiments/rain_ep_bias_train.py | 27 | ||||
| -rw-r--r-- | sdil/rain_ep_adapter.py | 20 |
3 files changed, 55 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, }) 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, diff --git a/sdil/rain_ep_adapter.py b/sdil/rain_ep_adapter.py index 81442aa..9c7ab8f 100644 --- a/sdil/rain_ep_adapter.py +++ b/sdil/rain_ep_adapter.py @@ -109,6 +109,20 @@ class RainGradientCorrector: ) @torch.no_grad() + def observe_neutral(self, local_states: Iterable[Tensor]) -> None: + """Fit the local bias field from one instruction-off observation.""" + if self.mode not in {"constant", "innovation"}: + raise ValueError( + "neutral predictor observations require constant or innovation mode") + local_states = list(local_states) + if any(value.requires_grad for value in local_states): + raise ValueError("Rain adapter received a requires-grad tensor") + bases, bias = self.bias.measure(local_states) + if self.debiaser is None: + self._initialize_debiaser(local_states) + self.debiaser.update_neutral(bases, bias, self.predictor_rate) + + @torch.no_grad() def apply(self, clean: Iterable[Tensor], local_states: Iterable[Tensor]) -> list[Tensor]: clean = list(clean) local_states = list(local_states) @@ -187,3 +201,9 @@ def attach_to_rain_estimator(estimator, corrector: RainGradientCorrector): estimator.sdil_corrector = corrector return estimator + +@torch.no_grad() +def observe_rain_neutral(estimator, corrector: RainGradientCorrector) -> None: + """Expose one label-free Rain equilibrium to the local predictor.""" + local_states = [updater.grad() for updater in estimator._param_updaters] + corrector.observe_neutral(local_states) |
