diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 01:47:20 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 01:47:20 -0500 |
| commit | 296c95045822c6d13de35515728528e411a7f8cd (patch) | |
| tree | be02a1c0053b1658781393669098ca68ff090a63 /experiments/run.py | |
| parent | 4f0c65d252d21438da48974ac77deb4a862abcd3 (diff) | |
experiments: add synthetic validation isolation
Diffstat (limited to 'experiments/run.py')
| -rw-r--r-- | experiments/run.py | 23 |
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): |
