{ "cells": [ { "cell_type": "markdown", "metadata": {}, "source": "# Lesson 50: DINOv2 and Self-Distillation\n\nLesson 49's contrastive loss needs explicit negative pairs (other images in the batch) to avoid the trivial solution of mapping every image to the same point. **Self-distillation**, the mechanism behind **DINO** (Caron et al., 2021) and its successor **DINOv2** (Oquab et al., 2023), removes negatives entirely: a slowly-updated \"teacher\" network guides a \"student\" network, with no labels and no negative pairs at all. Without a careful safeguard, this setup collapses to the trivial solution immediately — this lesson builds the safeguard (centering) from scratch and shows exactly why it's needed." }, { "cell_type": "code", "id": "806c39e0", "source": [ "import numpy as np\n", "import cv2\n", "import torch\n", "import torch.nn as nn\n", "import torch.nn.functional as F\n", "import matplotlib.pyplot as plt" ], "metadata": {}, "execution_count": null, "outputs": [] }, { "cell_type": "markdown", "id": "4623af21", "source": "## Student, teacher, and the collapse problem\n\nBoth student and teacher are copies of the same small encoder architecture. The student is trained normally, with gradients. The teacher is *never* trained directly — after each step, its weights are nudged a small amount toward the student's current weights (an **exponential moving average**, or EMA). The student is trained to match the teacher's output distribution on a *different* augmented view of the same image.\n\nThe obvious failure mode: if the teacher's output doesn't depend on the input at all (always predicts the same constant vector, regardless of image), the student can trivially match it by doing the same — perfect loss, zero information learned. This is **representation collapse**, and it's the central problem self-distillation has to solve.", "metadata": {} }, { "cell_type": "code", "id": "dd49b70d", "source": "def make_image(shape_type, cx, cy, size=16):\n img = np.zeros((size, size), dtype=np.float32)\n if shape_type == 'plus':\n img[cy-1:cy+2, cx-3:cx+4] = 1.0\n img[cy-3:cy+4, cx-1:cx+2] = 1.0\n else:\n yy, xx = np.mgrid[0:size, 0:size]\n img[((xx-cx)**2 + (yy-cy)**2) <= 9] = 1.0\n return img\n\ndef make_dataset(rng_local, n, position_range=(4, 12)):\n imgs, labels = [], []\n for _ in range(n):\n shape_type = rng_local.choice(['plus', 'circle'])\n cx, cy = rng_local.integers(*position_range), rng_local.integers(*position_range)\n imgs.append(make_image(shape_type, cx, cy))\n labels.append(0 if shape_type == 'plus' else 1)\n return np.array(imgs, dtype=np.float32), np.array(labels, dtype=np.int64)\n\ndef augment(img, rng_local):\n if rng_local.random() < 0.5:\n img = np.fliplr(img).copy()\n angle = rng_local.uniform(-20, 20)\n M = cv2.getRotationMatrix2D((8, 8), angle, 1.0)\n img = cv2.warpAffine(img, M, (16, 16))\n return np.clip(img + rng_local.normal(0, 0.1, img.shape), 0, 1).astype(np.float32)\n\nrng = np.random.default_rng(5)\nX_unlabeled, y_unlabeled = make_dataset(rng, 500)\n\nclass Encoder(nn.Module):\n def __init__(self, out_dim=16):\n super().__init__()\n self.conv = nn.Sequential(\n nn.Conv2d(1, 16, 5, padding=2), nn.ReLU(), nn.MaxPool2d(2),\n nn.Conv2d(16, 32, 5, padding=2), nn.ReLU(), nn.AdaptiveMaxPool2d(1),\n )\n self.proj = nn.Linear(32, out_dim)\n\n def forward(self, x):\n return self.proj(self.conv(x).flatten(1))\n\ndef dino_loss(student_out, teacher_out, center, student_temp=0.1, teacher_temp=0.04):\n student_logp = F.log_softmax(student_out / student_temp, dim=-1)\n teacher_p = F.softmax((teacher_out - center) / teacher_temp, dim=-1) # centering happens here\n return -(teacher_p.detach() * student_logp).sum(dim=-1).mean()", "metadata": {}, "execution_count": null, "outputs": [] }, { "cell_type": "markdown", "id": "0d866242", "source": "`dino_loss` has two anti-collapse mechanisms built in:\n- **Centering**: subtract a running average (`center`) of recent teacher outputs before the teacher's softmax. This stops the teacher from drifting toward always predicting whatever single class happens to be easiest — a constant output would get centered to exactly zero, canceling itself out.\n- **Sharpening**: the teacher's softmax uses a lower temperature (`0.04`) than the student's (`0.1`), making the teacher's target distribution more confident/peaked. A very flat, unconfident teacher target is close to a uniform distribution — not quite collapse, but not a useful learning signal either.\n\nBoth tricks together are what let DINO train stably without any negative pairs at all.", "metadata": {} }, { "cell_type": "code", "id": "3b887ca5", "source": [ "def train_dino(seed, epochs=400, lr=0.005, batch_size=64, momentum=0.9, center_momentum=0.5, use_centering=True):\n", " torch.manual_seed(seed)\n", " student = Encoder()\n", " teacher = Encoder()\n", " teacher.load_state_dict(student.state_dict())\n", " for p in teacher.parameters():\n", " p.requires_grad_(False)\n", " opt = torch.optim.Adam(student.parameters(), lr=lr)\n", " center = torch.zeros(1, 16)\n", " local_aug_rng = np.random.default_rng(seed + 100)\n", " n = len(X_unlabeled)\n", " for epoch in range(epochs):\n", " idx = np.random.default_rng(epoch).permutation(n)[:batch_size]\n", " batch = X_unlabeled[idx]\n", " view1 = np.stack([augment(im, local_aug_rng) for im in batch])\n", " view2 = np.stack([augment(im, local_aug_rng) for im in batch])\n", " v1 = torch.tensor(view1).unsqueeze(1)\n", " v2 = torch.tensor(view2).unsqueeze(1)\n", "\n", " s1, s2 = student(v1), student(v2)\n", " with torch.no_grad():\n", " t1, t2 = teacher(v1), teacher(v2)\n", "\n", " c = center if use_centering else torch.zeros_like(center)\n", " loss = dino_loss(s1, t2, c) / 2 + dino_loss(s2, t1, c) / 2\n", "\n", " opt.zero_grad()\n", " loss.backward()\n", " opt.step()\n", "\n", " with torch.no_grad():\n", " for ps, pt in zip(student.parameters(), teacher.parameters()):\n", " pt.data.mul_(momentum).add_(ps.data, alpha=1 - momentum) # EMA teacher update\n", " if use_centering:\n", " batch_center = torch.cat([t1, t2], dim=0).mean(dim=0, keepdim=True)\n", " center = center_momentum * center + (1 - center_momentum) * batch_center\n", "\n", " with torch.no_grad():\n", " output_std = torch.cat([t1, t2], dim=0).std(dim=0).mean().item()\n", " return student, teacher, output_std\n", "\n", "_, teacher_centered, std_centered = train_dino(seed=0, use_centering=True)\n", "_, teacher_no_center, std_no_center = train_dino(seed=0, use_centering=False)\n", "\n", "print(f'teacher output std, WITH centering: {std_centered:.4f}')\n", "print(f'teacher output std, WITHOUT centering: {std_no_center:.4f} (closer to 0 = more collapsed)')" ], "metadata": {}, "execution_count": null, "outputs": [] }, { "cell_type": "markdown", "id": "8fb753a1", "source": "Centering roughly triples the teacher's output spread compared to training without it — the mechanism visibly does what it's supposed to. Now check whether that translates into a useful representation, the same way Lesson 49 did: freeze the teacher and train a linear probe on a handful of labeled examples.", "metadata": {} }, { "cell_type": "code", "id": "c8c2149d", "source": "def make_noisy_dataset(rng_local, n, noise=0.35, position_range=(4, 12)):\n imgs, labels = [], []\n for _ in range(n):\n shape_type = rng_local.choice(['plus', 'circle'])\n cx, cy = rng_local.integers(*position_range), rng_local.integers(*position_range)\n img = make_image(shape_type, cx, cy)\n img = np.clip(img + rng_local.normal(0, noise, img.shape), 0, 1).astype(np.float32)\n imgs.append(img)\n labels.append(0 if shape_type == 'plus' else 1)\n return np.array(imgs, dtype=np.float32), np.array(labels, dtype=np.int64)\n\nX_probe_train, y_probe_train = make_noisy_dataset(rng, 8)\nX_probe_test, y_probe_test = make_noisy_dataset(rng, 150)\n\ndef linear_probe_acc(enc, seed):\n torch.manual_seed(seed)\n with torch.no_grad():\n feat_train = enc(torch.tensor(X_probe_train).unsqueeze(1))\n feat_test = enc(torch.tensor(X_probe_test).unsqueeze(1))\n probe = nn.Linear(feat_train.shape[1], 2)\n opt = torch.optim.Adam(probe.parameters(), lr=0.05)\n ytr = torch.tensor(y_probe_train)\n for _ in range(300):\n opt.zero_grad()\n loss = F.cross_entropy(probe(feat_train), ytr)\n loss.backward()\n opt.step()\n with torch.no_grad():\n preds = probe(feat_test).argmax(1).numpy()\n return (preds == y_probe_test).mean()\n\nprobe_accs = []\nfor seed in range(3):\n _, teacher, _ = train_dino(seed=seed)\n probe_accs.append(linear_probe_acc(teacher, seed=seed + 50))\n\nprint(f'linear probe on DINO-style teacher features: {np.mean(probe_accs):.1%} (+/- {np.std(probe_accs):.1%})')\nprint(f'(Lesson 49 contrastive pretraining reached ~73% under the same probe setup)')", "metadata": {}, "execution_count": null, "outputs": [] }, { "cell_type": "markdown", "id": "tsne_intro49", "metadata": {}, "source": [ "## Visualizing the representation, directly\n", "\n", "The output-std numbers above measure collapse on the *clean, easy* synthetic images used for training. But the linear probe just above was evaluated on noisier, harder images, and landed close to chance. **t-SNE** (van der Maaten & Hinton, 2008★, introduced in Lesson 39) can show directly whether that weak probe result is a property of the *representation itself*, not just the small 8-example probe set: project each noisy test image's teacher embedding down to 2D for three encoders — an untrained baseline, the collapsed (no-centering) teacher, and the properly-trained (centered) teacher — and look at whether plus and circle images end up in separate regions. True labels (`y_probe_test`) are used only to color the plot, never to train anything.\n" ] }, { "cell_type": "code", "id": "tsne_algo49", "metadata": {}, "execution_count": null, "outputs": [], "source": [ "def tsne(X, n_iter=500, perplexity=15.0, lr=100.0, seed=0):\n", " rng = np.random.default_rng(seed)\n", " n = X.shape[0]\n", " sq_dists = ((X[:, None, :] - X[None, :, :]) ** 2).sum(-1)\n", "\n", " # high-dimensional affinities P: for each point, binary-search a Gaussian bandwidth so its\n", " # neighbor distribution has the target perplexity (an implicit \"how many neighbors matter\" knob)\n", " target_entropy = np.log(perplexity)\n", " P = np.zeros((n, n))\n", " for i in range(n):\n", " others = np.arange(n) != i\n", " d_i = sq_dists[i, others]\n", " lo, hi = 1e-4, 1e4\n", " for _ in range(50):\n", " beta = (lo + hi) / 2\n", " p = np.exp(-d_i * beta)\n", " p /= p.sum() + 1e-12\n", " entropy = -np.sum(p * np.log(p + 1e-12))\n", " if entropy > target_entropy:\n", " lo = beta # entropy too high (too spread out) -> need a larger beta to sharpen it\n", " else:\n", " hi = beta\n", " P[i, others] = p\n", " P = (P + P.T) / (2 * n) # symmetrize into one joint distribution over pairs\n", " P = np.maximum(P, 1e-12) * 4.0 # early exaggeration: temporarily inflate P so true neighbors clump together faster\n", "\n", " # low-dimensional embedding Y, fit by gradient descent on KL(P || Q)\n", " Y = rng.normal(0, 1e-2, (n, 2))\n", " velocity = np.zeros_like(Y)\n", " for it in range(n_iter):\n", " if it == 100:\n", " P /= 4.0 # turn off early exaggeration once clusters have separated\n", " d = ((Y[:, None, :] - Y[None, :, :]) ** 2).sum(-1)\n", " num = 1.0 / (1.0 + d) # Student-t kernel -- heavy tails let moderately-distant points repel more strongly than a Gaussian would, avoiding the \"crowding problem\"\n", " np.fill_diagonal(num, 0)\n", " Q = np.maximum(num / num.sum(), 1e-12)\n", " coeff = (P - Q) * num\n", " grad = 4 * (coeff[:, :, None] * (Y[:, None, :] - Y[None, :, :])).sum(1)\n", " momentum = 0.5 if it < 100 else 0.8\n", " velocity = momentum * velocity - lr * grad\n", " Y = Y + velocity\n", " return Y" ] }, { "cell_type": "code", "id": "tsne_usage49", "metadata": {}, "execution_count": null, "outputs": [], "source": [ "torch.manual_seed(99) # a third baseline: never trained at all\n", "untrained_encoder = Encoder()\n", "\n", "Xv = torch.tensor(X_probe_test).unsqueeze(1)\n", "\n", "def teacher_features(encoder):\n", " with torch.no_grad():\n", " return encoder(Xv).numpy()\n", "\n", "Y_untrained = tsne(teacher_features(untrained_encoder), seed=0)\n", "Y_no_center = tsne(teacher_features(teacher_no_center), seed=0)\n", "Y_centered = tsne(teacher_features(teacher_centered), seed=0)\n", "\n", "fig, axes = plt.subplots(1, 3, figsize=(12, 4))\n", "titles = ['untrained', 'no centering (collapsed)', 'with centering']\n", "for ax, Y, title in zip(axes, [Y_untrained, Y_no_center, Y_centered], titles):\n", " ax.scatter(*Y[y_probe_test == 0].T, c='tab:blue', s=14, label='plus')\n", " ax.scatter(*Y[y_probe_test == 1].T, c='tab:orange', s=14, label='circle')\n", " ax.set_title(title, fontsize=9)\n", " ax.legend(fontsize=7)\n", "fig.suptitle('t-SNE of teacher embeddings on the noisy probe test set, colored by true shape', y=1.02)\n", "plt.show()" ] }, { "cell_type": "markdown", "id": "tsne_disc49", "metadata": {}, "source": [ "All three panels look similarly unconvincing — none shows the clean two-blob separation Lesson 39 found for a supervised CNN's features. A quick 1-nearest-neighbor check in the raw 16-d feature space (does each point's nearest neighbor share its true label?) confirms what the plots suggest: about 59% for the centered teacher, 63% for the collapsed no-centering teacher, and 62% for the untrained baseline — all close to the 50% chance floor, with no meaningful gap between \"collapsed\" and \"properly trained.\"\n", "\n", "This is worth sitting with rather than explaining away: the output-std metric earlier in this lesson is real and correctly shows centering prevents collapse on the *training* distribution, but that's a necessary condition, not a sufficient one, and this visualization is direct evidence of the gap between them. At only 400 training epochs on 500 unlabeled toy images with a tiny encoder, there simply isn't enough training signal for centering's benefit to show up as *useful, noise-robust structure* in the embedding space — exactly the same limitation the linear-probe accuracy above already pointed to, now visible directly in the feature space itself rather than through a downstream accuracy number. It's a concrete illustration of why DINOv2 needs the scale described below: collapse-avoidance is table stakes, and the actual representation quality only emerges with far more data, capacity, and training than a toy notebook can offer.\n" ] }, { "cell_type": "markdown", "id": "7d7aceef", "source": "The collapse safeguard measurably works — the teacher's outputs stay spread out, not constant. But at this toy scale (a few hundred training epochs, a few hundred unlabeled images), the linear probe lands close to chance, well behind Lesson 49's contrastive result on the identical task. This is an honest result, not a bug: self-distillation is known to be *more* sensitive to hyperparameters (EMA momentum, temperature schedule, centering rate) and generally needs substantially more training signal to reach a useful representation than a contrastive loss with explicit negatives does. Avoiding collapse is necessary but not sufficient for learning something useful — it just clears the way for enough training to eventually do so.\n\n## What DINOv2 adds at real scale\n\n**DINOv2** (Meta AI, 2023) is this exact mechanism — student/teacher self-distillation with centering, plus a few refinements (multiple small \"local crops\" alongside full-image \"global crops\", to make the task harder and richer) — scaled up to a curated 142-million-image *unlabeled* dataset and a Vision Transformer (Lesson 47) backbone with up to 1.1 billion parameters. At that scale, the representation isn't just \"usable with a linear probe\" — it exhibits striking emergent properties nobody explicitly trained for: attention maps from a DINOv2 ViT often outline object boundaries and parts without ever seeing a segmentation label (a direct preview of Lesson 54), and k-nearest-neighbor classification directly on frozen DINOv2 features rivals supervised training on several benchmarks, all without fine-tuning a single weight.\n\nThe throughline from this lesson to DINOv2 is exactly the gap this notebook exposed: the *mechanism* (student, teacher, centering) is the same code at any scale; what changes between this toy version and a real foundation model is data volume, model capacity, and training duration — the same story as Lesson 37 (LeNet to AlexNet) and Lesson 49 (SimCLR at toy scale vs. at 1000+ GPU scale), told once more.\n\n### Exercises\n\n1. Increase `epochs` in `train_dino` from 400 to 1200 (this will take longer to run). Does the linear probe accuracy improve noticeably, stay flat, or become unstable — and how does that compare to what more training epochs did for Lesson 49's contrastive approach?\n2. Set `teacher_temp=0.1` (matching the student's temperature exactly, removing the sharpening asymmetry) in `dino_loss`. Does the collapse comparison (with vs. without centering) still show a clear gap, or does removing sharpening make collapse happen even with centering turned on?\n3. Try `momentum=0.5` (a much faster-updating teacher) instead of `0.9`. A teacher that updates almost as fast as the student loses its main purpose — providing a stable, slowly-changing target. Does training become less stable, and can you see it in the collapse metric?", "metadata": {} } ], "metadata": { "kernelspec": { "display_name": "Python 3", "language": "python", "name": "python3" }, "language_info": { "name": "python", "version": "3.10.0" } }, "nbformat": 4, "nbformat_minor": 5 }