"""Fast checks for frozen train/validation/test protocol plumbing.""" import os import sys import torch sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from sdil.data import get_dataset_splits def loader_labels(loader): return torch.cat([labels.cpu() for _, labels in loader]) def main(): first = get_dataset_splits( "mnist", batch_size=256, device="cpu", val_examples=1000, split_seed=2027) second = get_dataset_splits( "mnist", batch_size=256, device="cpu", val_examples=1000, split_seed=2027) train, validation, test, n_in, n_out, metadata = first _, validation_2, _, _, _, metadata_2 = second assert (n_in, n_out) == (784, 10) assert metadata["train_examples"] == 59000 assert metadata["validation_examples"] == 1000 assert metadata["test_examples"] == 10000 assert metadata["validation_index_sha256"] == metadata_2["validation_index_sha256"] assert metadata["validation_class_counts"] == {str(i): 100 for i in range(10)} assert torch.equal(loader_labels(validation), loader_labels(validation_2)) assert sum(y.numel() for _, y in train) == 59000 assert sum(y.numel() for _, y in validation) == 1000 assert sum(y.numel() for _, y in test) == 10000 print("validation split hash:", metadata["validation_index_sha256"]) print("train/validation/test: 59000/1000/10000; stratification exact") print("ALL PROTOCOL SMOKE CHECKS PASSED") if __name__ == "__main__": main()