summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rw-r--r--experiments/rain_ep_adapter_smoke.py10
-rw-r--r--experiments/rain_ep_bias_train.py27
2 files changed, 35 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,