diff --git a/configs/2b_480_640_agibot.yaml b/configs/2b_480_640_agibot.yaml index b8989c2..6528d78 100644 --- a/configs/2b_480_640_agibot.yaml +++ b/configs/2b_480_640_agibot.yaml @@ -84,11 +84,6 @@ dataloader_train: - /mnt/amlfs-03/shared/datasets/agibot-custom-converted-0829-fullres/agibot.0902_3239 - /mnt/amlfs-03/shared/datasets/agibot-custom-converted-0829-fullres/agibot.0903_3237 - /mnt/amlfs-03/shared/datasets/agibot-custom-converted-0829-fullres/agibot.0903_3248 - - /mnt/amlfs-03/shared/datasets/agibot-custom-converted-0829-fullres/agibot.0911_3228 - - /mnt/amlfs-03/shared/datasets/agibot-custom-converted-0829-fullres/agibot.0911_3228_lang - - /mnt/amlfs-03/shared/datasets/agibot-custom-converted-0829-fullres/agibot.0911_3228_long - - /mnt/amlfs-03/shared/datasets/agibot-custom-converted-0829-fullres/agibot.0911_3228_long_6k - - /mnt/amlfs-03/shared/datasets/agibot-custom-converted-0829-fullres/agibot.0911_3228_test - /mnt/amlfs-03/shared/datasets/agibot-custom-converted-0829-fullres/agibot.0911_3228_train - /mnt/amlfs-03/shared/datasets/agibot-custom-converted-0829-fullres/agibot.0911_3228_train_lang_sep - /mnt/amlfs-03/shared/datasets/agibot-custom-converted-0829-fullres/agibot.0911_3228_train_lang_sep_max600 diff --git a/tests/test_config_no_test_leakage.py b/tests/test_config_no_test_leakage.py new file mode 100644 index 0000000..f43b4a0 --- /dev/null +++ b/tests/test_config_no_test_leakage.py @@ -0,0 +1,45 @@ +from pathlib import Path + + +def _load_dataset_paths(config_path): + """Extract dataset_path entries from a config YAML using only stdlib.""" + paths = [] + in_dataset_path = False + with open(config_path) as f: + for line in f: + stripped = line.strip() + if not stripped: + continue + indent = len(line) - len(line.lstrip()) + if stripped.startswith("dataset_path:"): + in_dataset_path = True + continue + if in_dataset_path: + if stripped.startswith("- "): + if indent > 4: + paths.append(stripped[2:].strip()) + else: + in_dataset_path = False + elif indent <= 4: + in_dataset_path = False + return paths + + +def test_no_test_splits_or_parent_overlap_in_configs(): + config_dir = Path(__file__).parent.parent / "configs" + for config_path in config_dir.glob("*.yaml"): + basenames = [p.rsplit("/", 1)[-1] for p in _load_dataset_paths(config_path)] + base_set = set(basenames) + + test_splits = [b for b in basenames if b.endswith("_test")] + assert not test_splits, ( + f"{config_path}: test splits cannot be used for training: {test_splits}" + ) + + for b in basenames: + if b.endswith("_train") or b.endswith("_test"): + base = b.rsplit("_", 1)[0] + assert base not in base_set, ( + f"{config_path}: base dataset {base} overlaps with its split {b}" + ) +