diff options
| -rw-r--r-- | experiments/rain_ep_bias_train.py | 14 |
1 files changed, 14 insertions, 0 deletions
diff --git a/experiments/rain_ep_bias_train.py b/experiments/rain_ep_bias_train.py index fcca67e..8b6e9ad 100644 --- a/experiments/rain_ep_bias_train.py +++ b/experiments/rain_ep_bias_train.py @@ -43,6 +43,7 @@ def parse_args() -> argparse.Namespace: parser.add_argument("--batch-size", type=int, default=128) parser.add_argument("--training-iterations", type=int, default=12) parser.add_argument("--inference-iterations", type=int, default=30) + parser.add_argument("--deterministic", action="store_true") parser.add_argument("--seed", type=int, default=1988) parser.add_argument("--output", type=Path, required=True) return parser.parse_args() @@ -95,6 +96,10 @@ def main() -> None: torch.manual_seed(args.seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(args.seed) + if args.deterministic: + torch.backends.cudnn.benchmark = False + torch.backends.cudnn.deterministic = True + torch.use_deterministic_algorithms(True) device = torch.device(args.device) training_data, test_data = load_fashion_mnist( @@ -191,17 +196,24 @@ def main() -> None: parameter.clamp_() test_accuracy, test_cost = evaluate( network, cost, inference_minimizer, test_loader) + finite = all( + bool(torch.isfinite(parameter.state).all()) + for parameter in energy.params() + ) record = { "epoch": epoch, "train_accuracy": total_correct / total, "train_cost": total_cost / total, "test_accuracy": test_accuracy, "test_cost": test_cost, + "finite": finite, "corrector": dict(corrector.last_diagnostics), "wall_seconds": time.time() - start, } metrics.append(record) print(json.dumps(record), flush=True) + if not finite: + break report = { "schema": "rain_ep_structured_bias_screen_v1", @@ -228,9 +240,11 @@ def main() -> None: "inference_iterations": args.inference_iterations, "seed": args.seed, "device": str(device), + "deterministic": args.deterministic, "autodiff_used_for_learning": False, }, "metrics": metrics, + "epochs_completed": len(metrics), "final": metrics[-1], } args.output.parent.mkdir(parents=True, exist_ok=True) |
