diff options
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/rain_ep_bias_train.py | 27 |
1 files changed, 22 insertions, 5 deletions
diff --git a/experiments/rain_ep_bias_train.py b/experiments/rain_ep_bias_train.py index 15838d6..abc4138 100644 --- a/experiments/rain_ep_bias_train.py +++ b/experiments/rain_ep_bias_train.py @@ -47,6 +47,10 @@ def parse_args() -> argparse.Namespace: 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) + parser.add_argument( + "--evaluation-split", choices=("test", "train_holdout"), + default="test") + parser.add_argument("--data-seed", type=int, default=1988) 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) @@ -111,10 +115,21 @@ def main() -> None: training_data, test_data = load_fashion_mnist( normalize=True, augment_32x32=False) - train_generator = torch.Generator().manual_seed(args.seed) - train_indices = torch.randperm( - len(training_data), generator=train_generator)[:args.train_limit] - test_indices = torch.arange(min(args.test_limit, len(test_data))) + split_generator = torch.Generator().manual_seed(args.data_seed) + split_indices = torch.randperm( + len(training_data), generator=split_generator) + if args.evaluation_split == "train_holdout": + if args.train_limit + args.test_limit > len(training_data): + raise ValueError("training and holdout subsets overlap") + train_indices = split_indices[:args.train_limit] + evaluation_data = training_data + evaluation_indices = split_indices[ + args.train_limit:args.train_limit + args.test_limit] + else: + train_indices = split_indices[:args.train_limit] + evaluation_data = test_data + evaluation_indices = torch.arange(min(args.test_limit, len(test_data))) + train_generator = torch.Generator().manual_seed(args.seed + 65537) training_loader = torch.utils.data.DataLoader( torch.utils.data.Subset(training_data, train_indices.tolist()), batch_size=args.batch_size, shuffle=True, generator=train_generator, @@ -125,7 +140,7 @@ def main() -> None: 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()), + torch.utils.data.Subset(evaluation_data, evaluation_indices.tolist()), batch_size=args.batch_size, shuffle=False, num_workers=0) energy = ConvHopfieldEnergy28( @@ -268,6 +283,8 @@ def main() -> None: "epochs": args.epochs, "train_limit": args.train_limit, "test_limit": args.test_limit, + "evaluation_split": args.evaluation_split, + "data_seed": args.data_seed, "batch_size": args.batch_size, "training_iterations": args.training_iterations, "inference_iterations": args.inference_iterations, |
