summaryrefslogtreecommitdiff
path: root/experiments/protocol_smoke.py
blob: 551987cab7807d41c578ad4bc53753479aecc540 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
"""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
from sdil.core import SDILConfig, SDILNet
from experiments.run import calibration_work_per_event


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

    net = SDILNet([10, 8, 8, 3], device="cpu")
    simultaneous = calibration_work_per_event(
        net, SDILConfig(pert_ndirs=2, pert_mode="simultaneous"))
    layerwise = calibration_work_per_event(
        net, SDILConfig(pert_ndirs=2, pert_mode="layerwise"))
    assert simultaneous == {
        "batch_loss_evaluations": 4,
        "forward_equivalent_batches": 5.0,
        "perturbation_batch_expansion": 4,
    }
    assert layerwise["batch_loss_evaluations"] == 9
    assert abs(layerwise["forward_equivalent_batches"] - 11.0 / 3.0) < 1e-12
    print("validation split hash:", metadata["validation_index_sha256"])
    print("train/validation/test: 59000/1000/10000; stratification exact")
    print("simultaneous/layerwise calibration cost accounting: exact")
    print("ALL PROTOCOL SMOKE CHECKS PASSED")


if __name__ == "__main__":
    main()