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