summaryrefslogtreecommitdiff
path: root/experiments/rain_ep_bias_train.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 17:01:42 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 17:01:42 -0500
commitc77affedb1d04f20c01fd88a0047efb07e023915 (patch)
tree1cf09638e68cc2511aae27b84445e41f483d2dd5 /experiments/rain_ep_bias_train.py
parent038e97b684b7b460dedec3018677729f956ac691 (diff)
feat: add fixed Rain validation holdout
Diffstat (limited to 'experiments/rain_ep_bias_train.py')
-rw-r--r--experiments/rain_ep_bias_train.py27
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,