{ "cells": [ { "cell_type": "markdown", "metadata": {}, "source": [ "# Lesson 51: Masked Autoencoders\n", "\n", "Lesson 48's autoencoder compressed a whole image through a narrow bottleneck vector. **Masked autoencoders** (MAE, He et al., 2021) create the information bottleneck a completely different way: chop the image into patches (Lesson 47's ViT patchify), hide most of them entirely — not noised, not blurred, just never shown to the encoder at all — and train the network to reconstruct the missing patches from whatever's left. This is the third self-supervised pretraining paradigm in this part, after contrastive learning (Lesson 49) and self-distillation (Lesson 50): no labels, no negative pairs, no teacher network — just a plausible-sounding fill-in-the-blank task invented directly from unlabeled images." ] }, { "cell_type": "code", "execution_count": null, "id": "ffcc071c", "metadata": {}, "outputs": [], "source": [ "import numpy as np\n", "import torch\n", "import torch.nn as nn\n", "import torch.nn.functional as F\n", "import matplotlib.pyplot as plt" ] }, { "cell_type": "markdown", "id": "31168735", "metadata": {}, "source": [ "## Patchify and randomly mask\n", "\n", "Same `16x16` image, `4x4` patches as Lesson 47 — 16 patches total. For each image, pick a random subset to keep **visible** and discard the rest; the discarded patches are not part of the encoder's input in any form, not even as zeros or noise." ] }, { "cell_type": "code", "execution_count": null, "id": "ec6180a0", "metadata": {}, "outputs": [], "source": [ "IMG = 16\n", "PATCH = 4\n", "N_PATCHES = (IMG // PATCH) ** 2 # 16\n", "PATCH_DIM = PATCH * PATCH # 16 (grayscale)\n", "\n", "def patchify(img, patch_size=PATCH):\n", " B, C, H, W = img.shape\n", " P = patch_size\n", " patches = img.unfold(2, P, P).unfold(3, P, P)\n", " patches = patches.contiguous().view(B, C, -1, P, P).permute(0, 2, 1, 3, 4)\n", " return patches.reshape(B, -1, C * P * P)\n", "\n", "def unpatchify(patches, patch_size=PATCH, img_size=IMG):\n", " B, N, D = patches.shape\n", " P = patch_size\n", " n_side = img_size // P\n", " x = patches.view(B, n_side, n_side, P, P).permute(0, 1, 3, 2, 4).contiguous()\n", " return x.view(B, 1, img_size, img_size)\n", "\n", "def positional_encoding(T, D):\n", " pos = torch.arange(T).unsqueeze(1).float()\n", " i = torch.arange(D).unsqueeze(0).float()\n", " angle_rates = 1.0 / (10000 ** (2 * (i // 2) / D))\n", " angles = pos * angle_rates\n", " pe = torch.zeros(T, D)\n", " pe[:, 0::2] = torch.sin(angles[:, 0::2])\n", " pe[:, 1::2] = torch.cos(angles[:, 1::2])\n", " return pe\n", "\n", "def random_masking(x, mask_ratio, generator=None):\n", " B, N, D = x.shape\n", " n_visible = max(1, int(N * (1 - mask_ratio)))\n", " noise = torch.rand(B, N, generator=generator)\n", " ids_shuffle = torch.argsort(noise, dim=1) # random permutation of patch indices per image\n", " ids_keep = ids_shuffle[:, :n_visible] # kept as VISIBLE, in original patch-index units\n", " ids_mask = ids_shuffle[:, n_visible:] # HIDDEN entirely from the encoder\n", " x_visible = torch.gather(x, 1, ids_keep.unsqueeze(-1).expand(-1, -1, D))\n", " return x_visible, ids_keep, ids_mask\n", "\n", "def make_image(shape_type, cx, cy, size=IMG):\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", "\n", "def 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", "\n", "rng = np.random.default_rng(6)\n", "X_train, y_train = make_dataset(rng, 400)\n", "X_test, y_test = make_dataset(rng, 150)\n", "Xt = torch.tensor(X_train).unsqueeze(1)\n", "Xte = torch.tensor(X_test).unsqueeze(1)\n", "\n", "print(f'{N_PATCHES} patches per image, {PATCH_DIM} pixels each')" ] }, { "cell_type": "markdown", "id": "b50c4de0", "metadata": {}, "source": [ "## An asymmetric encoder-decoder\n", "\n", "The **encoder** is a small Transformer (Lesson 46) that only ever processes the visible patches — masked patches never enter it, so the encoder is fast and never wastes computation on tokens that get thrown away. The **decoder** is where the masked patches reappear: a shared, learned **mask token** vector stands in for every hidden patch, combined with that patch's original positional encoding (so the decoder knows *where* each blank belongs, even though it never saw *what* was there), and a lightweight Transformer predicts each masked patch's raw pixel values. Loss is computed **only on the masked patches** — the network gets no credit for copying pixels it was already shown." ] }, { "cell_type": "code", "execution_count": null, "id": "adb05664", "metadata": {}, "outputs": [], "source": [ "class MAEEncoder(nn.Module):\n", " def __init__(self, patch_dim=PATCH_DIM, embed_dim=32, n_heads=4, n_layers=2):\n", " super().__init__()\n", " self.embed = nn.Linear(patch_dim, embed_dim)\n", " self.register_buffer('pe', positional_encoding(N_PATCHES, embed_dim))\n", " layer = nn.TransformerEncoderLayer(embed_dim, n_heads, dim_feedforward=embed_dim * 2,\n", " batch_first=True, dropout=0.0)\n", " self.encoder = nn.TransformerEncoder(layer, num_layers=n_layers)\n", "\n", " def forward(self, x_visible, ids_keep):\n", " tok = self.embed(x_visible) + self.pe[ids_keep] # positional encoding at each patch's TRUE index\n", " return self.encoder(tok)\n", "\n", "class MAEDecoder(nn.Module):\n", " def __init__(self, embed_dim=32, decoder_dim=16, patch_dim=PATCH_DIM, n_heads=4, n_layers=1):\n", " super().__init__()\n", " self.embed = nn.Linear(embed_dim, decoder_dim)\n", " self.mask_token = nn.Parameter(torch.zeros(1, 1, decoder_dim))\n", " self.register_buffer('pe', positional_encoding(N_PATCHES, decoder_dim))\n", " layer = nn.TransformerEncoderLayer(decoder_dim, n_heads, dim_feedforward=decoder_dim * 2,\n", " batch_first=True, dropout=0.0)\n", " self.decoder = nn.TransformerEncoder(layer, num_layers=n_layers)\n", " self.pred = nn.Linear(decoder_dim, patch_dim)\n", " self.decoder_dim = decoder_dim\n", "\n", " def forward(self, encoded_visible, ids_keep, ids_mask):\n", " B, n_mask = ids_mask.shape\n", " vis_tok = self.embed(encoded_visible) + self.pe[ids_keep]\n", " mask_tok = self.mask_token.expand(B, n_mask, -1) + self.pe[ids_mask]\n", " full = torch.zeros(B, N_PATCHES, self.decoder_dim)\n", " full = full.scatter(1, ids_keep.unsqueeze(-1).expand(-1, -1, self.decoder_dim), vis_tok)\n", " full = full.scatter(1, ids_mask.unsqueeze(-1).expand(-1, -1, self.decoder_dim), mask_tok)\n", " return self.pred(self.decoder(full)) # (B, N_PATCHES, patch_dim), full sequence, original order\n", "\n", "class MAE(nn.Module):\n", " def __init__(self, mask_ratio=0.5, embed_dim=32, decoder_dim=16):\n", " super().__init__()\n", " self.mask_ratio = mask_ratio\n", " self.encoder = MAEEncoder(embed_dim=embed_dim)\n", " self.decoder = MAEDecoder(embed_dim=embed_dim, decoder_dim=decoder_dim)\n", "\n", " def forward(self, img, generator=None):\n", " patches = patchify(img)\n", " x_visible, ids_keep, ids_mask = random_masking(patches, self.mask_ratio, generator=generator)\n", " encoded = self.encoder(x_visible, ids_keep)\n", " pred = self.decoder(encoded, ids_keep, ids_mask)\n", " D = patches.shape[-1]\n", " pred_masked = torch.gather(pred, 1, ids_mask.unsqueeze(-1).expand(-1, -1, D))\n", " target_masked = torch.gather(patches, 1, ids_mask.unsqueeze(-1).expand(-1, -1, D))\n", " loss = F.mse_loss(pred_masked, target_masked)\n", " return loss, pred, ids_mask, ids_keep, patches\n", "\n", "def train_mae(mask_ratio, epochs=600, lr=0.001, seed=0):\n", " torch.manual_seed(seed)\n", " model = MAE(mask_ratio=mask_ratio)\n", " opt = torch.optim.Adam(model.parameters(), lr=lr)\n", " for _ in range(epochs):\n", " opt.zero_grad()\n", " loss, *_ = model(Xt)\n", " loss.backward()\n", " opt.step()\n", " return model" ] }, { "cell_type": "markdown", "id": "abb0a964", "metadata": {}, "source": [ "## How much can be reconstructed from how little?\n", "\n", "Sweep the fraction of patches hidden from 25% to 75%, and compare each model's masked-patch reconstruction error against a trivial baseline: predicting the training set's mean pixel value for every masked patch, regardless of image." ] }, { "cell_type": "code", "execution_count": null, "id": "b484ca19", "metadata": {}, "outputs": [], "source": [ "train_patches = patchify(Xt)\n", "mean_pixel = train_patches.mean()\n", "\n", "print(f'{\"mask ratio\":>10} {\"MAE test MSE\":>14} {\"mean-pixel baseline\":>20}')\n", "mae_models = {}\n", "for mr in [0.25, 0.5, 0.75]:\n", " model = train_mae(mask_ratio=mr)\n", " mae_models[mr] = model\n", " with torch.no_grad():\n", " test_loss, pred, ids_mask, ids_keep, patches = model(Xte, generator=torch.Generator().manual_seed(123))\n", " D = patches.shape[-1]\n", " target_masked = torch.gather(patches, 1, ids_mask.unsqueeze(-1).expand(-1, -1, D))\n", " baseline_mse = F.mse_loss(torch.full_like(target_masked, mean_pixel), target_masked).item()\n", " print(f'{mr:>10} {test_loss.item():>14.4f} {baseline_mse:>20.4f}')" ] }, { "cell_type": "markdown", "id": "432d63e0", "metadata": {}, "source": [ "Reconstruction gets harder as more of the image is hidden — unsurprising, there's simply less evidence to work with — but the model beats the naive baseline by a wide margin even at 75% masking, where only 4 of the 16 patches are ever shown to the encoder. That's only possible because these shapes have enormous internal redundancy: a circle's silhouette is entirely determined by its center and radius, so a handful of visible boundary patches, combined with attention across all of them, is enough to infer the rest. Real MAE finds the same thing at a much larger scale — natural photographs are also highly spatially redundant (a patch of sky, a patch of skin, a patch of brick predicts its neighbors) — which is *why* He et al. found the best pretraining mask ratio for images is around 75%, dramatically higher than the ~15% word-masking ratio BERT uses for language, where far less of any given sentence is predictable from the rest." ] }, { "cell_type": "code", "execution_count": null, "id": "34ad594b", "metadata": {}, "outputs": [], "source": [ "idx = 7\n", "model75 = mae_models[0.75]\n", "with torch.no_grad():\n", " _, pred, ids_mask, ids_keep, patches = model75(Xte, generator=torch.Generator().manual_seed(42))\n", "\n", "masked_view = patches.clone()\n", "masked_view[torch.arange(len(patches)).unsqueeze(1), ids_mask] = 0.5 # gray out hidden patches\n", "recon_view = patches.clone()\n", "recon_view[torch.arange(len(patches)).unsqueeze(1), ids_mask] = pred[torch.arange(len(patches)).unsqueeze(1), ids_mask]\n", "\n", "orig_img = unpatchify(patches)[idx, 0].numpy()\n", "masked_img = unpatchify(masked_view)[idx, 0].numpy()\n", "recon_img = unpatchify(recon_view)[idx, 0].numpy()\n", "\n", "fig, axes = plt.subplots(1, 3, figsize=(7, 2.8))\n", "for ax, im, title in zip(axes, [orig_img, masked_img, recon_img],\n", " ['original', 'what the encoder saw\\n(75% hidden)', 'reconstructed']):\n", " ax.imshow(im, cmap='gray', vmin=0, vmax=1); ax.set_title(title, fontsize=9); ax.axis('off')\n", "plt.show()" ] }, { "cell_type": "markdown", "id": "e88cea66", "metadata": {}, "source": [ "## From reconstruction to representation: a linear probe\n", "\n", "The reconstruction task itself is thrown away after pretraining — what real MAE keeps is the **encoder**, exactly as Lessons 48 and 49 kept only the trained encoder/teacher and discarded the contrastive/distillation machinery around it. Freeze the mask-ratio-0.75 encoder, feed it *every* patch at once (no masking, `mask_ratio=0` at evaluation time — the whole point of pretraining under a harsh bottleneck was to prepare the encoder for the easy case), average-pool its output into one feature vector per image, and train a linear probe on a handful of labeled examples." ] }, { "cell_type": "code", "execution_count": null, "id": "1eef0af3", "metadata": {}, "outputs": [], "source": [ "def encode_full(encoder, img):\n", " patches = patchify(img)\n", " B, N, _ = patches.shape\n", " ids_all = torch.arange(N).unsqueeze(0).expand(B, -1)\n", " with torch.no_grad():\n", " encoded = encoder(patches, ids_all)\n", " return encoded.mean(dim=1) # average-pool over all 16 patches -> one vector per image\n", "\n", "def make_noisy_dataset(rng_local, n, noise=0.2, 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", "\n", "probe_rng = np.random.default_rng(11)\n", "X_probe_train, y_probe_train = make_noisy_dataset(probe_rng, 8)\n", "X_probe_test, y_probe_test = make_noisy_dataset(probe_rng, 150)\n", "Xpt = torch.tensor(X_probe_train).unsqueeze(1)\n", "Xpte = torch.tensor(X_probe_test).unsqueeze(1)\n", "\n", "def linear_probe_acc(encoder, seed):\n", " torch.manual_seed(seed)\n", " feat_train = encode_full(encoder, Xpt)\n", " feat_test = encode_full(encoder, Xpte)\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", "\n", "mae_accs, random_accs = [], []\n", "for seed in range(5):\n", " mae_accs.append(linear_probe_acc(mae_models[0.75].encoder, seed=seed + 50))\n", " torch.manual_seed(seed)\n", " random_encoder = MAEEncoder()\n", " random_accs.append(linear_probe_acc(random_encoder, seed=seed + 50))\n", "\n", "print(f'linear probe on RANDOM (untrained) MAE encoder: {np.mean(random_accs):.1%} (+/- {np.std(random_accs):.1%})')\n", "print(f'linear probe on MAE-pretrained encoder: {np.mean(mae_accs):.1%} (+/- {np.std(mae_accs):.1%})')" ] }, { "cell_type": "markdown", "id": "6d67f201", "metadata": {}, "source": [ "That's a genuinely disappointing number — the MAE-pretrained encoder's frozen, average-pooled features probe *worse* than a random, untrained encoder's. This isn't a bug; it's a real, documented property of masked-reconstruction pretraining, and it's worth understanding why. Contrastive learning (Lesson 49) and self-distillation (Lesson 50) both directly optimize a pooled, whole-image embedding to be useful on its own — that's the entire training signal. MAE's training signal never touches a pooled representation at all: the encoder only has to produce per-patch tokens rich enough that attention *within the decoder* can fill in missing patches. Nothing forces the *average* of those tokens to be a good summary of \"what shape is this\" — a perfectly good reconstruction-supporting representation can average out to something a linear probe finds useless. The original MAE paper reports exactly this pattern at real scale: substantially lower linear-probe accuracy than contrastive methods, with MAE's real payoff only showing up after **full fine-tuning** (updating every encoder weight on the downstream task, not just a linear head on top of frozen features) — a genuinely different deployment story from Lessons 48 and 49, not a worse one." ] }, { "cell_type": "markdown", "id": "ca9d4625", "metadata": {}, "source": [ "## Three self-supervised paradigms, one shared pattern\n", "\n", "Contrastive learning (Lesson 49), self-distillation (Lesson 50), and masked reconstruction (this lesson) invent three completely different training signals from the same raw material — unlabeled images — yet all three exist to answer the same question: what's the cheapest information-destroying task whose solution forces a network to learn something genuinely useful about images? Contrastive learning destroys identity across augmented views and asks the network to recover it; self-distillation destroys the teacher's certainty and asks the student to match it anyway; masked reconstruction destroys most of the image outright and asks the network to fill in the rest. But this lesson's probe result is a real warning against assuming all three produce interchangeable representations: MAE's training signal is architecturally cheap (the encoder never even processes masked tokens, unlike the other two, which always see the whole image) and excellent for downstream *fine-tuning*, but it earns that efficiency by never once forcing a pooled, linearly-useful summary vector to exist — which is exactly what contrastive learning and self-distillation are built around. Choosing among these three in practice is a genuine engineering tradeoff, not a strict ranking.\n", "\n", "### Exercises\n", "\n", "1. Increase `mask_ratio` to `0.9` (only 1-2 patches visible). Does reconstruction MSE still beat the mean-pixel baseline, or does it collapse to roughly the baseline's error — and does that reveal a point past which this dataset's redundancy runs out?\n", "2. The decoder here is deliberately tiny (`n_layers=1`, `decoder_dim=16`, smaller than the encoder). Try making the decoder *larger* than the encoder (e.g. `decoder_dim=64`, `n_layers=3`) while keeping the encoder fixed. Does reconstruction quality improve much — and does that match real MAE's design choice of a small decoder, on the reasoning that the decoder is discarded after pretraining anyway?\n", "3. Instead of freezing the encoder for the linear probe, fine-tune it end to end: let gradients from the 8-example classification loss update the encoder's own weights too, not just a linear head on top of frozen features. Does end-to-end fine-tuning recover the gap between MAE and the random baseline — matching the real MAE paper's finding that fine-tuned MAE features are excellent, even when frozen ones probe poorly?" ] } ], "metadata": { "kernelspec": { "display_name": "Python 3", "language": "python", "name": "python3" }, "language_info": { "name": "python", "version": "3.10.0" } }, "nbformat": 4, "nbformat_minor": 5 }