summaryrefslogtreecommitdiff
path: root/experiments/run.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/run.py')
-rw-r--r--experiments/run.py23
1 files changed, 17 insertions, 6 deletions
diff --git a/experiments/run.py b/experiments/run.py
index c66acca..8c59a50 100644
--- a/experiments/run.py
+++ b/experiments/run.py
@@ -25,7 +25,8 @@ from sdil.baselines import BPNet, dfa_config, evaluate
from sdil.local_baselines import FANet
from sdil import probes
from sdil.data import (get_dataset_splits, onehot, make_hierarchical,
- make_teacher_student, make_tentmap)
+ make_teacher_student, make_tentmap,
+ split_training_loader)
REAL_DATASETS = ("mnist", "fmnist", "cifar10")
@@ -128,17 +129,27 @@ def load_task(args, device):
else:
raise ValueError(f"unknown dataset: {args.dataset}")
train, test, n_in, n_out = task
+ validation = None
+ split_details = {}
+ if args.val_examples:
+ train, validation, split_details = split_training_loader(
+ train, args.val_examples, args.split_seed, args.batch_size)
+ if args.eval_split == "validation":
+ if validation is None:
+ raise ValueError("--eval_split validation requires --val_examples > 0")
+ evaluation = validation
+ else:
+ evaluation = test
split = {
"dataset": args.dataset,
"task_seed": args.task_seed,
- "train_examples": args.task_train_examples,
+ "train_examples": len(train.x),
"test_examples": args.task_test_examples,
- "evaluation_split": "test",
+ "evaluation_split": args.eval_split,
"synthetic_generator": True,
}
- if args.eval_split != "test" or args.val_examples:
- raise ValueError("synthetic validation splits are not implemented yet")
- return train, test, n_in, n_out, split
+ split.update(split_details)
+ return train, evaluation, n_in, n_out, split
def train(args):