| from __future__ import annotations |
|
|
| import json |
| import random |
| from pathlib import Path |
|
|
| import numpy as np |
| import pandas as pd |
| import torch |
| import trackio |
| from model import ContrastiveEncoder, nt_xent, parameter_count |
| from safetensors.torch import save_file |
| from sklearn.linear_model import LogisticRegression |
| from sklearn.metrics import accuracy_score |
| from torch.nn import functional as F |
| from torch.utils.data import DataLoader, TensorDataset |
|
|
| PROJECT_DIR = Path(__file__).resolve().parent |
| ROOT_DIR = PROJECT_DIR.parents[1] |
| DATA_DIR = ROOT_DIR / "projects" / "tiny-vision-foundry" / "data" |
| ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "contrastive-pocket" |
|
|
|
|
| def seed_everything(seed: int) -> None: |
| random.seed(seed) |
| np.random.seed(seed) |
| torch.manual_seed(seed) |
|
|
|
|
| def load_split(name: str) -> tuple[torch.Tensor, np.ndarray]: |
| frame = pd.read_parquet(DATA_DIR / f"{name}.parquet") |
| pixels = np.stack(frame["image"].to_numpy()).astype(np.float32) / 16.0 |
| labels = frame["label"].to_numpy(dtype=np.int64, copy=True) |
| return torch.from_numpy(pixels.reshape(-1, 1, 8, 8)), labels |
|
|
|
|
| def shift_image(image: torch.Tensor, vertical: int, horizontal: int) -> torch.Tensor: |
| shifted = torch.roll(image, shifts=(vertical, horizontal), dims=(-2, -1)) |
| if vertical > 0: |
| shifted[..., :vertical, :] = 0 |
| elif vertical < 0: |
| shifted[..., vertical:, :] = 0 |
| if horizontal > 0: |
| shifted[..., :, :horizontal] = 0 |
| elif horizontal < 0: |
| shifted[..., :, horizontal:] = 0 |
| return shifted |
|
|
|
|
| def augment_batch(pixels: torch.Tensor) -> torch.Tensor: |
| augmented = pixels.clone() |
| for index in range(len(augmented)): |
| vertical = int(torch.randint(-1, 2, (1,)).item()) |
| horizontal = int(torch.randint(-1, 2, (1,)).item()) |
| augmented[index] = shift_image( |
| augmented[index], |
| vertical, |
| horizontal, |
| ) |
| augmented += torch.randn_like(augmented) * 0.045 |
| dropout_mask = torch.rand_like(augmented) < 0.035 |
| augmented[dropout_mask] = 0 |
| contrast = 0.85 + torch.rand(len(augmented), 1, 1, 1) * 0.3 |
| return torch.clamp(augmented * contrast, 0, 1) |
|
|
|
|
| @torch.inference_mode() |
| def embed(model: ContrastiveEncoder, pixels: torch.Tensor) -> np.ndarray: |
| model.eval() |
| embeddings = [] |
| for start in range(0, len(pixels), 256): |
| embeddings.append(model.encode(pixels[start : start + 256]).numpy()) |
| return np.concatenate(embeddings) |
|
|
|
|
| def label_subset(labels: np.ndarray, per_class: int) -> np.ndarray: |
| selected = [] |
| for label in range(10): |
| selected.extend(np.flatnonzero(labels == label)[:per_class].tolist()) |
| return np.asarray(selected, dtype=np.int64) |
|
|
|
|
| def probe( |
| train_features: np.ndarray, |
| train_labels: np.ndarray, |
| test_features: np.ndarray, |
| test_labels: np.ndarray, |
| per_class: int, |
| ) -> float: |
| indices = label_subset(train_labels, per_class) |
| classifier = LogisticRegression( |
| C=3.0, |
| max_iter=3000, |
| random_state=2035, |
| ) |
| classifier.fit(train_features[indices], train_labels[indices]) |
| return float(accuracy_score(test_labels, classifier.predict(test_features))) |
|
|
|
|
| def probe_suite( |
| train_features: np.ndarray, |
| train_labels: np.ndarray, |
| test_features: np.ndarray, |
| test_labels: np.ndarray, |
| ) -> dict[str, float]: |
| return { |
| str(per_class): probe( |
| train_features, |
| train_labels, |
| test_features, |
| test_labels, |
| per_class, |
| ) |
| for per_class in [10, 25, 100] |
| } |
|
|
|
|
| def invariance_score(model: ContrastiveEncoder, pixels: torch.Tensor) -> dict[str, float]: |
| model.eval() |
| with torch.no_grad(): |
| first = F.normalize(model.encode(augment_batch(pixels[:256])), dim=1) |
| second = F.normalize(model.encode(augment_batch(pixels[:256])), dim=1) |
| positive = float((first * second).sum(dim=1).mean()) |
| mismatched = float((first * second.roll(1, dims=0)).sum(dim=1).mean()) |
| return { |
| "positive_pair_cosine": positive, |
| "mismatched_pair_cosine": mismatched, |
| "invariance_gap": positive - mismatched, |
| } |
|
|
|
|
| def main() -> None: |
| seed_everything(2035) |
| train_pixels, train_labels = load_split("train") |
| test_pixels, test_labels = load_split("test") |
| loader = DataLoader( |
| TensorDataset(train_pixels), |
| batch_size=128, |
| shuffle=True, |
| drop_last=True, |
| generator=torch.Generator().manual_seed(2035), |
| ) |
| random_encoder = ContrastiveEncoder() |
| 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() |
|
|
| model = ContrastiveEncoder() |
| optimizer = torch.optim.AdamW(model.parameters(), lr=0.0015, weight_decay=0.002) |
| epochs = 220 |
| scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs) |
| trackio.init( |
| project="contrastive-pocket", |
| name="simclr-8x8-v1", |
| config={ |
| "parameters": parameter_count(model), |
| "epochs": epochs, |
| "temperature": 0.18, |
| "labels_used_in_pretraining": False, |
| }, |
| ) |
| for epoch in range(1, epochs + 1): |
| model.train() |
| losses = [] |
| for (pixels,) in loader: |
| first = model(augment_batch(pixels)) |
| second = model(augment_batch(pixels)) |
| loss = nt_xent(first, second, temperature=0.18) |
| optimizer.zero_grad(set_to_none=True) |
| loss.backward() |
| optimizer.step() |
| losses.append(loss.item()) |
| scheduler.step() |
| if epoch == 1 or epoch % 10 == 0: |
| trackio.log( |
| { |
| "epoch": epoch, |
| "contrastive_loss": float(np.mean(losses)), |
| "learning_rate": scheduler.get_last_lr()[0], |
| } |
| ) |
| trackio.finish() |
|
|
| learned_train = embed(model, train_pixels) |
| learned_test = embed(model, test_pixels) |
| results = { |
| "model": "Contrastive Pocket", |
| "parameters": parameter_count(model), |
| "unlabeled_pretraining_examples": len(train_pixels), |
| "epochs": epochs, |
| "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_cosine": { |
| "contrastive_encoder": invariance_score(model, train_pixels), |
| "random_encoder": invariance_score(random_encoder, train_pixels), |
| }, |
| "embedding_standard_deviation": float(learned_train.std(axis=0).mean()), |
| } |
| ARTIFACT_DIR.mkdir(parents=True, exist_ok=True) |
| save_file(model.state_dict(), ARTIFACT_DIR / "model.safetensors") |
| (ARTIFACT_DIR / "evaluation.json").write_text( |
| json.dumps(results, indent=2), |
| encoding="utf-8", |
| ) |
| print(json.dumps(results, indent=2)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|