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