summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rw-r--r--experiments/rain_ep_bias_train.py14
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)