summaryrefslogtreecommitdiff
path: root/experiments/synthetic_smoke.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 01:47:20 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 01:47:20 -0500
commit296c95045822c6d13de35515728528e411a7f8cd (patch)
treebe02a1c0053b1658781393669098ca68ff090a63 /experiments/synthetic_smoke.py
parent4f0c65d252d21438da48974ac77deb4a862abcd3 (diff)
experiments: add synthetic validation isolation
Diffstat (limited to 'experiments/synthetic_smoke.py')
-rw-r--r--experiments/synthetic_smoke.py25
1 files changed, 25 insertions, 0 deletions
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")