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()
|