diff options
Diffstat (limited to 'experiments/protocol_smoke.py')
| -rw-r--r-- | experiments/protocol_smoke.py | 39 |
1 files changed, 39 insertions, 0 deletions
diff --git a/experiments/protocol_smoke.py b/experiments/protocol_smoke.py new file mode 100644 index 0000000..36da996 --- /dev/null +++ b/experiments/protocol_smoke.py @@ -0,0 +1,39 @@ +"""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() |
