File size: 2,238 Bytes
02323ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
from __future__ import annotations

import json

from model import ContrastiveEncoder, parameter_count
from safetensors.torch import load_file
from train import (
    ARTIFACT_DIR,
    embed,
    invariance_score,
    load_split,
    probe_suite,
    seed_everything,
)


def main() -> None:
    seed_everything(2035)
    train_pixels, train_labels = load_split("train")
    test_pixels, test_labels = load_split("test")
    learned = ContrastiveEncoder()
    learned.load_state_dict(load_file(ARTIFACT_DIR / "model.safetensors"))
    random_encoder = ContrastiveEncoder()
    learned_train = embed(learned, train_pixels)
    learned_test = embed(learned, test_pixels)
    random_train = embed(random_encoder, train_pixels)
    random_test = embed(random_encoder, test_pixels)
    raw_train = train_pixels.reshape(len(train_pixels), -1).numpy()
    raw_test = test_pixels.reshape(len(test_pixels), -1).numpy()
    results = {
        "model": "Contrastive Pocket",
        "parameters": parameter_count(learned),
        "unlabeled_pretraining_examples": len(train_pixels),
        "epochs": 220,
        "embedding_dimensions": learned_train.shape[1],
        "linear_probe_accuracy_by_examples_per_class": {
            "contrastive_encoder": probe_suite(
                learned_train,
                train_labels,
                learned_test,
                test_labels,
            ),
            "random_encoder": probe_suite(
                random_train,
                train_labels,
                random_test,
                test_labels,
            ),
            "raw_pixels": probe_suite(
                raw_train,
                train_labels,
                raw_test,
                test_labels,
            ),
        },
        "augmentation_invariance": {
            "contrastive_encoder": invariance_score(learned, train_pixels),
            "random_encoder": invariance_score(random_encoder, train_pixels),
        },
        "embedding_standard_deviation": float(learned_train.std(axis=0).mean()),
    }
    (ARTIFACT_DIR / "evaluation.json").write_text(
        json.dumps(results, indent=2),
        encoding="utf-8",
    )
    print(json.dumps(results, indent=2))


if __name__ == "__main__":
    main()