pocket-jepa / source /train.py
ARotting's picture
Publish 14.8K parameter masked joint-embedding predictive encoder
dcaea4a verified
Raw
History Blame Contribute Delete
9.38 kB
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()