from __future__ import annotations import copy import json from pathlib import Path import numpy as np import pandas as pd import torch import trackio from model import JEncoder, JPredictor, parameter_count from safetensors.torch import save_file from sklearn.datasets import load_digits from sklearn.linear_model import LogisticRegression from sklearn.metrics import accuracy_score from sklearn.model_selection import train_test_split from torch import nn from torch.nn import functional as F from torch.utils.data import DataLoader, TensorDataset PROJECT_DIR = Path(__file__).resolve().parent ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "pocket-jepa" DATA_DIR = PROJECT_DIR / "data" def seed_everything(seed: int) -> None: np.random.seed(seed) torch.manual_seed(seed) torch.set_num_threads(1) def mask_blocks(images: torch.Tensor, rng: np.random.Generator) -> torch.Tensor: masked = images.clone() for index in range(len(masked)): size = int(rng.integers(2, 5)) row = int(rng.integers(0, 9 - size)) column = int(rng.integers(0, 9 - size)) masked[index, row : row + size, column : column + size] = 0 masked += 0.035 * torch.randn_like(masked) return masked.clamp(0, 1) def deterministic_masks(images: np.ndarray, seed: int) -> np.ndarray: rng = np.random.default_rng(seed) output = images.copy() for index in range(len(output)): size = 3 row = int(rng.integers(0, 6)) column = int(rng.integers(0, 6)) output[index, row : row + size, column : column + size] = 0 return output def redundancy_loss(prediction: torch.Tensor, target: torch.Tensor) -> torch.Tensor: prediction = (prediction - prediction.mean(0)) / ( prediction.std(0, unbiased=False) + 1e-4 ) target = (target - target.mean(0)) / ( target.std(0, unbiased=False) + 1e-4 ) correlation = prediction.T @ target / len(prediction) diagonal = torch.diagonal(correlation) identity = (diagonal - 1).pow(2).mean() off_diagonal = correlation - torch.diag(diagonal) return identity + 0.01 * off_diagonal.pow(2).sum() / prediction.shape[1] @torch.inference_mode() def encode(model: nn.Module, images: np.ndarray) -> np.ndarray: model.eval() tensor = torch.from_numpy(images.astype(np.float32)) embeddings = model(tensor).numpy() norms = np.linalg.norm(embeddings, axis=1, keepdims=True) return embeddings / np.maximum(norms, 1e-8) def label_budget_indices( labels: np.ndarray, per_class: int, seed: int ) -> np.ndarray: rng = np.random.default_rng(seed) selected = [] for label in np.unique(labels): candidates = np.flatnonzero(labels == label) selected.extend(rng.choice(candidates, per_class, replace=False)) return np.asarray(selected) def probe( train_embeddings: np.ndarray, train_labels: np.ndarray, test_embeddings: np.ndarray, test_labels: np.ndarray, selected: np.ndarray, ) -> float: classifier = LogisticRegression(C=3.0, max_iter=2_000) classifier.fit(train_embeddings[selected], train_labels[selected]) return float( accuracy_score(test_labels, classifier.predict(test_embeddings)) ) def main() -> None: seed_everything(2043) digits = load_digits() images = (digits.images / 16.0).astype(np.float32) labels = digits.target.astype(np.int64) indices = np.arange(len(images)) train_indices, test_indices = train_test_split( indices, test_size=0.25, random_state=2043, stratify=labels ) train_images = images[train_indices] test_images = images[test_indices] train_labels = labels[train_indices] test_labels = labels[test_indices] online = JEncoder() random_control = copy.deepcopy(online) target = copy.deepcopy(online) predictor = JPredictor() for parameter in target.parameters(): parameter.requires_grad = False optimizer = torch.optim.AdamW( [*online.parameters(), *predictor.parameters()], lr=2e-3, weight_decay=2e-4, ) loader = DataLoader( TensorDataset(torch.from_numpy(train_images)), batch_size=256, shuffle=True, drop_last=True, generator=torch.Generator().manual_seed(2043), ) epochs = 260 trackio.init( project="pocket-jepa", name="masked-j-space-v1", config={ "encoder_parameters": parameter_count(online), "predictor_parameters": parameter_count(predictor), "epochs": epochs, "labels_per_class_for_probe": 10, "target_ema": 0.99, }, ) rng = np.random.default_rng(2043) history = [] online.train() predictor.train() for epoch in range(1, epochs + 1): losses = [] for (complete,) in loader: context = mask_blocks(complete, rng) prediction = predictor(online(context)) with torch.no_grad(): target_embedding = target(complete) cosine = 1 - F.cosine_similarity( prediction, target_embedding, dim=1 ).mean() decorrelation = redundancy_loss(prediction, target_embedding) loss = cosine + 0.35 * decorrelation optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_( [*online.parameters(), *predictor.parameters()], 1.0 ) optimizer.step() with torch.no_grad(): for target_parameter, online_parameter in zip( target.parameters(), online.parameters(), strict=True ): target_parameter.mul_(0.99).add_( online_parameter, alpha=0.01 ) losses.append(float(loss.detach())) record = {"epoch": epoch, "pretraining_loss": float(np.mean(losses))} history.append(record) if epoch % 20 == 0: trackio.log(record) selected = label_budget_indices(train_labels, per_class=10, seed=3043) learned_train = encode(target, train_images) learned_test = encode(target, test_images) learned_masked = encode(target, deterministic_masks(test_images, 4043)) random_train = encode(random_control, train_images) random_test = encode(random_control, test_images) random_masked = encode( random_control, deterministic_masks(test_images, 4043) ) raw_train = train_images.reshape(len(train_images), -1) raw_test = test_images.reshape(len(test_images), -1) raw_masked = deterministic_masks(test_images, 4043).reshape( len(test_images), -1 ) results = { "model": "Pocket JEPA", "encoder_parameters": parameter_count(target), "predictor_training_parameters": parameter_count(predictor), "unlabeled_pretraining_examples": len(train_images), "pretraining_epochs": epochs, "linear_probe_labels": int(len(selected)), "labels_per_class": 10, "clean_accuracy": { "pocket_jepa": probe( learned_train, train_labels, learned_test, test_labels, selected, ), "random_encoder": probe( random_train, train_labels, random_test, test_labels, selected, ), "raw_pixels": probe( raw_train, train_labels, raw_test, test_labels, selected ), }, "masked_accuracy": { "pocket_jepa": probe( learned_train, train_labels, learned_masked, test_labels, selected, ), "random_encoder": probe( random_train, train_labels, random_masked, test_labels, selected, ), "raw_pixels": probe( raw_train, train_labels, raw_masked, test_labels, selected ), }, "final_pretraining_loss": history[-1]["pretraining_loss"], } ARTIFACT_DIR.mkdir(parents=True, exist_ok=True) DATA_DIR.mkdir(parents=True, exist_ok=True) save_file(target.state_dict(), ARTIFACT_DIR / "model.safetensors") (ARTIFACT_DIR / "evaluation.json").write_text( json.dumps(results, indent=2), encoding="utf-8" ) np.savez_compressed( ARTIFACT_DIR / "j_space.npz", embeddings=learned_test, images=test_images, labels=test_labels, ) pd.DataFrame( { "source_index": indices, "label": labels, "split": np.where( np.isin(indices, test_indices), "test", "unlabeled_train" ), } ).to_parquet(DATA_DIR / "split_manifest.parquet", index=False) trackio.log( { "clean_probe_accuracy": results["clean_accuracy"]["pocket_jepa"], "masked_probe_accuracy": results["masked_accuracy"]["pocket_jepa"], "random_clean_accuracy": results["clean_accuracy"][ "random_encoder" ], } ) trackio.finish() print(json.dumps(results, indent=2)) if __name__ == "__main__": main()