summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rw-r--r--experiments/cifar_image_smoke.py89
1 files changed, 89 insertions, 0 deletions
diff --git a/experiments/cifar_image_smoke.py b/experiments/cifar_image_smoke.py
new file mode 100644
index 0000000..c23da45
--- /dev/null
+++ b/experiments/cifar_image_smoke.py
@@ -0,0 +1,89 @@
+#!/usr/bin/env python3
+"""Deterministic checks for the convolutional CIFAR data path."""
+import argparse
+import os
+import sys
+
+import torch
+
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+from sdil.data import _CIFARImageLoader, get_cifar_image_splits
+
+
+EXPECTED_SPLIT_100_SHA256 = (
+ "a88b2ca885bba36dc7407171e3afbce8dca2dd83c7c1d4222e4418b0ca3de6f0")
+
+
+def synthetic_checks():
+ x = (torch.arange(12 * 3 * 32 * 32, dtype=torch.float32)
+ .reshape(12, 3, 32, 32) / 1000.0)
+ y = torch.arange(12)
+ first = _CIFARImageLoader(x, y, 5, True, augment=True, seed=17)
+ second = _CIFARImageLoader(x.clone(), y.clone(), 5, True, augment=True, seed=17)
+ epoch_sums = []
+ for _ in range(2):
+ stream_a, stream_b = list(first), list(second)
+ assert len(stream_a) == len(stream_b) == 3
+ assert all(torch.equal(xa, xb) and torch.equal(ya, yb)
+ for (xa, ya), (xb, yb) in zip(stream_a, stream_b))
+ assert sum(len(labels) for _, labels in stream_a) == 12
+ assert all(tuple(images.shape[1:]) == (3, 32, 32)
+ for images, _ in stream_a)
+ epoch_sums.append(sum(float(images.sum()) for images, _ in stream_a))
+ assert epoch_sums[0] != epoch_sums[1], "augmentation RNG did not advance"
+
+ plain_a = _CIFARImageLoader(x, y, 5, False, augment=False, seed=1)
+ plain_b = _CIFARImageLoader(x, y, 5, False, augment=False, seed=999)
+ assert all(torch.equal(xa, xb) and torch.equal(ya, yb)
+ for (xa, ya), (xb, yb) in zip(plain_a, plain_b))
+ try:
+ _CIFARImageLoader(x, y, 5, False, augment=True)
+ except ValueError:
+ pass
+ else:
+ raise AssertionError("evaluation augmentation guard is missing")
+
+
+def real_data_checks(device):
+ args = dict(
+ batch_size=7, device=device, train_limit=13, val_examples=100,
+ split_seed=2027, loader_seed=11)
+ train, validation, test, shape, n_out, metadata = get_cifar_image_splits(**args)
+ assert shape == (3, 32, 32) and n_out == 10
+ assert tuple(train.x.shape) == (13, 3, 32, 32)
+ assert tuple(validation.x.shape) == (100, 3, 32, 32)
+ assert tuple(test.x.shape) == (10000, 3, 32, 32)
+ assert metadata["validation_class_counts"] == {str(i): 10 for i in range(10)}
+ assert metadata["validation_index_sha256"] == EXPECTED_SPLIT_100_SHA256
+
+ validation_before = validation.x.clone()
+ test_before = test.x[:100].clone()
+ list(validation)
+ list(test)
+ assert torch.equal(validation.x, validation_before)
+ assert torch.equal(test.x[:100], test_before)
+
+ paired, _, _, _, _, paired_metadata = get_cifar_image_splits(**args)
+ for (xa, ya), (xb, yb) in zip(train, paired):
+ assert torch.equal(xa, xb) and torch.equal(ya, yb)
+ assert metadata == paired_metadata
+ return metadata
+
+
+def main():
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--device", default="cpu")
+ parser.add_argument("--synthetic-only", action="store_true")
+ args = parser.parse_args()
+ synthetic_checks()
+ metadata = None if args.synthetic_only else real_data_checks(args.device)
+ print({
+ "synthetic": "passed",
+ "real_data": "skipped" if metadata is None else "passed",
+ "validation_index_sha256": (
+ None if metadata is None else metadata["validation_index_sha256"]),
+ })
+
+
+if __name__ == "__main__":
+ main()