#!/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()