From 296c95045822c6d13de35515728528e411a7f8cd Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 22 Jul 2026 01:47:20 -0500 Subject: experiments: add synthetic validation isolation --- experiments/synthetic_smoke.py | 25 +++++++++++++++++++++++++ 1 file changed, 25 insertions(+) (limited to 'experiments/synthetic_smoke.py') diff --git a/experiments/synthetic_smoke.py b/experiments/synthetic_smoke.py index d2325a8..2b66892 100644 --- a/experiments/synthetic_smoke.py +++ b/experiments/synthetic_smoke.py @@ -35,6 +35,30 @@ def check_task(dataset, expected_in, expected_out, extra): print(f"{dataset}: input={n_in}, classes={n_out}, fixed task seed={args.task_seed}") +def check_synthetic_validation(): + argv = [ + "run.py", "--dataset", "teacher", "--device", "cpu", + "--task_train_examples", "80", "--task_test_examples", "32", + "--val_examples", "20", "--eval_split", "validation", + "--batch_size", "16", "--depth", "2", "--width", "12", + "--task_n_in", "20", "--task_classes", "4", + "--teacher_depth", "3", "--teacher_width", "10", + ] + old = sys.argv + try: + sys.argv = argv + args = get_args() + finally: + sys.argv = old + train, validation, _, _, split = load_task(args, "cpu") + assert sum(y.numel() for _, y in train) == 60 + assert sum(y.numel() for _, y in validation) == 20 + assert split["evaluation_split"] == "validation" + assert split["validation_index_sha256"] + assert split["split_from_training_only"] is True + print("synthetic validation: 60/20; test generator untouched") + + if __name__ == "__main__": torch.manual_seed(0) check_task("teacher", 20, 4, @@ -44,4 +68,5 @@ if __name__ == "__main__": ["--task_levels", "4", "--task_classes", "4"]) check_task("tentmap", 5, 2, ["--task_levels", "4", "--task_n_in", "5"]) + check_synthetic_validation() print("ALL SYNTHETIC TASK CHECKS PASSED") -- cgit v1.2.3