CNN Backbones and Pretraining

Masked image modelling

Hide most of an image and train a network to fill in what is missing, which teaches it about the world without a single label.

Read these first

On this page 7
  1. Why this was tried
  2. How it was made hard enough
  3. The clever part about speed
  4. Where you have already seen it
  5. The honest part
  6. Remember this
  7. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

Masked image modelling hides most of a picture and trains a network to paint back what was covered.

You have walked past a torn cinema poster on a wall. Half the letters are gone and the actor's face is partly ripped away.

You still read the film's name and recognise the face. Your eyes get a fraction of the poster, and your head supplies the rest.

You can only do that because you know how posters, letters and faces normally look. Anyone who could fill in the gaps has, without noticing, learned a great deal about the world.

That is the whole idea. Cover most of the picture. Ask the network to fill it in. Whatever it needs to learn in order to succeed is exactly what we wanted it to learn.

Why this was tried

Language models had already won with this trick. Hide a word in a sentence, predict it, and the model ends up understanding grammar and facts.

Images resisted for years. The reason is a real difference between the two, and it is worth understanding.

Words are already chunks with meaning. Pixels are not. A single pixel carries almost nothing, and its neighbours are nearly identical to it.

So hiding one pixel teaches nothing. The network copies a neighbour and scores well. The task has to be made much harder.

How it was made hard enough

Two changes did it.

Hide whole blocks, not single pixels. The picture is cut into square patches, and entire patches are removed. There is no neighbouring pixel to copy from inside a missing patch.

Hide most of them. Language models hide about fifteen percent of the words. Image models hide seventy-five percent, and sometimes more. That much removal forces real understanding rather than local smoothing.

   original            given to the model         asked to produce
   ###  ###           ###   ??                     ###  ###
   ###  ###     ->     ??   ??            ->       ###  ###
   ###  ###            ??  ###                     ###  ###

The clever part about speed

If seventy-five percent of the patches are removed, why carry them through the network at all?

One design throws them away entirely. The heavy part of the network sees only the quarter that remains. A small, light part is then given the leftovers to fill in.

This makes training several times faster, and cheaper training means you can afford a bigger model. The measurement further down shows the attention work dropping by a factor of sixteen.

Where you have already seen it

  • Photo editing tools that remove an object and fill the hole convincingly.
  • Old photo restoration that repairs scratches and tears.
  • Medical imaging models pretrained on scans that were never labelled by a radiologist.

The honest part

The result of this training is not a network that paints pretty pictures. The painting part is thrown away.

What you keep is the part that learned to understand. It gets attached to a classifier, or a detector, and trained on the real task. That final step needs far fewer labels than usual.

Also worth knowing: the filled-in patches usually look blurry. That is expected. When the model is unsure, averaging over the possibilities is the safest answer, and averaging looks like blur.

Remember this

  • Hide most of the image in whole patches, then train the network to fill it in.
  • Very high hiding rates are necessary, because neighbouring pixels are too similar.
  • The filling-in part is discarded afterwards; the understanding part is what you keep.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install "torch==2.5.1"

Run against PyTorch 2.5.1 on CPU. It takes about seventeen seconds.

A complete masked autoencoder, small enough to read

The pictures are toy textures, so the model can learn in three hundred steps on a laptop. Every mechanism in the real thing is present: patching, random masking, an encoder that sees only what is visible, a mask token, and a loss on hidden patches only.

tiny_mae.py
import torch
import torch.nn as nn
import torch.nn.functional as F

P, G, D = 4, 4, 64            # 4x4 pixel patches, a 4x4 grid of them, 64-wide tokens
NT = G * G                    # 16 patches per picture
MASK_RATIO = 0.75
KEEP = int(NT * (1 - MASK_RATIO))


def make_batch(n, g):
    """Three 16x16 textures: horizontal stripes, vertical stripes, checkerboard."""
    y = torch.randint(0, 3, (n,), generator=g)
    phase = torch.randint(0, 4, (n, 2), generator=g)
    rows = torch.arange(16).view(16, 1).expand(16, 16)
    cols = torch.arange(16).view(1, 16).expand(16, 16)
    x = torch.zeros(n, 16, 16)
    for i, k in enumerate(y.tolist()):
        a, b = phase[i].tolist()
        if k == 0:
            x[i] = (((rows + a) // 2) % 2).float()
        elif k == 1:
            x[i] = (((cols + b) // 2) % 2).float()
        else:
            x[i] = ((((rows + a) // 2) + ((cols + b) // 2)) % 2).float()
    return x, y


def patchify(x):
    n = x.shape[0]
    return x.view(n, G, P, G, P).permute(0, 1, 3, 2, 4).reshape(n, NT, P * P)


def unpatchify(t):
    n = t.shape[0]
    return t.view(n, G, G, P, P).permute(0, 1, 3, 2, 4).reshape(n, 16, 16)


class TinyMAE(nn.Module):
    def __init__(self):
        super().__init__()
        self.embed = nn.Linear(P * P, D)
        self.pos = nn.Parameter(torch.zeros(NT, D))
        self.encoder = nn.TransformerEncoder(
            nn.TransformerEncoderLayer(D, 4, 128, batch_first=True, dropout=0.0), 2)
        self.decoder = nn.TransformerEncoder(
            nn.TransformerEncoderLayer(D, 4, 128, batch_first=True, dropout=0.0), 1)
        self.mask_token = nn.Parameter(torch.zeros(D))
        self.out = nn.Linear(D, P * P)

    def forward(self, x, keep_idx, mask_idx):
        n = x.shape[0]
        patches = patchify(x)
        tokens = self.embed(patches) + self.pos
        visible = torch.gather(tokens, 1, keep_idx[..., None].expand(-1, -1, D))
        latent = self.encoder(visible)                       # 4 tokens, not 16

        full = self.mask_token.expand(n, NT, D) + self.pos   # placeholders everywhere
        full = full.scatter(1, keep_idx[..., None].expand(-1, -1, D), latent)
        pred = self.out(self.decoder(full))

        want = torch.gather(patches, 1, mask_idx[..., None].expand(-1, -1, P * P))
        got = torch.gather(pred, 1, mask_idx[..., None].expand(-1, -1, P * P))
        return F.mse_loss(got, want), pred                   # loss on hidden patches only


torch.manual_seed(0)
g = torch.Generator().manual_seed(0)
model = TinyMAE()
opt = torch.optim.Adam(model.parameters(), lr=2e-3)

print(f"patches per picture {NT}, hidden {MASK_RATIO:.0%}, encoder sees {KEEP}")
print(f"attention pairs if the encoder saw everything : {NT * NT}")
print(f"attention pairs the way MAE does it           : {KEEP * KEEP}"
      f"   ({NT * NT / (KEEP * KEEP):.0f}x cheaper)")
print(f"MSE of guessing flat grey for every pixel     : 0.25\n")

for step in range(1, 301):
    x, _ = make_batch(64, g)
    perm = torch.argsort(torch.rand(64, NT, generator=g), dim=1)
    loss, _ = model(x, perm[:, :KEEP], perm[:, KEEP:])
    opt.zero_grad()
    loss.backward()
    opt.step()
    if step in (1, 10, 50, 100, 300):
        print(f"step {step:4d}   MSE on the hidden patches {loss.item():.4f}")

# one held-out picture
gt = torch.Generator().manual_seed(7)
x, _ = make_batch(1, gt)
perm = torch.argsort(torch.rand(1, NT, generator=gt), dim=1)
keep_idx, mask_idx = perm[:, :KEEP], perm[:, KEEP:]
with torch.no_grad():
    _, pred = model(x, keep_idx, mask_idx)

patches = patchify(x)
seen = torch.zeros_like(patches)
seen.scatter_(1, keep_idx[..., None].expand(-1, -1, P * P),
              torch.gather(patches, 1, keep_idx[..., None].expand(-1, -1, P * P)))
filled = pred.clone()
filled.scatter_(1, keep_idx[..., None].expand(-1, -1, P * P),
                torch.gather(patches, 1, keep_idx[..., None].expand(-1, -1, P * P)))

RAMP = " .:-=+*#%@"


def draw(img, title):
    print("\n" + title)
    for row in img[0]:
        print("  " + "".join(RAMP[min(9, max(0, int(v.item() * 9.999)))] for v in row))


draw(x, "the original")
draw(unpatchify(seen), f"what the encoder was given ({KEEP} of {NT} patches)")
draw(unpatchify(pred), "the decoder's raw output for all 16 patches")
draw(unpatchify(filled), "hidden patches filled in, visible patches pasted back")
Output
patches per picture 16, hidden 75%, encoder sees 4
attention pairs if the encoder saw everything : 256
attention pairs the way MAE does it           : 16   (16x cheaper)
MSE of guessing flat grey for every pixel     : 0.25

step    1   MSE on the hidden patches 0.7514
step   10   MSE on the hidden patches 0.0796
step   50   MSE on the hidden patches 0.0007
step  100   MSE on the hidden patches 0.0000
step  300   MSE on the hidden patches 0.0000

the original
                  
                  
  @@@@@@@@@@@@@@@@
  @@@@@@@@@@@@@@@@
                  
                  
  @@@@@@@@@@@@@@@@
  @@@@@@@@@@@@@@@@
                  
                  
  @@@@@@@@@@@@@@@@
  @@@@@@@@@@@@@@@@
                  
                  
  @@@@@@@@@@@@@@@@
  @@@@@@@@@@@@@@@@

what the encoder was given (4 of 16 patches)
                  
                  
      @@@@    @@@@
      @@@@    @@@@
                  
                  
                  
                  
                  
                  
          @@@@    
          @@@@    
                  
                  
          @@@@    
          @@@@    

the decoder's raw output for all 16 patches
      . -.    . -.
        *       * 
  @@@@+: @@@@@+: @
  @@@@=@  @@@@=@  
                  
                  
  @@@@@@@@@@@@@@@@
  @@@@@@@@@@@@@@@@
          . -.    
            *     
  @@@@@@@@+: @@@@@
  @@@@@@@@=@  @@@@
          . -.    
            *     
  @@@@@@@@*: @@@@@
  @@@@@@@@=@  @@@@

hidden patches filled in, visible patches pasted back
                  
                  
  @@@@@@@@@@@@@@@@
  @@@@@@@@@@@@@@@@
                  
                  
  @@@@@@@@@@@@@@@@
  @@@@@@@@@@@@@@@@
                  
                  
  @@@@@@@@@@@@@@@@
  @@@@@@@@@@@@@@@@
                  
                  
  @@@@@@@@@@@@@@@@
  @@@@@@@@@@@@@@@@

Reading the output

Four visible patches were enough to rebuild the whole picture. The final drawing matches the original exactly. From a quarter of the image, the model inferred both the texture class and its phase, which is the shift of the stripes. There is no way to do that without having learned what these textures are.

Look at the raw decoder output, and where the mess is. The garbage sits exactly at the four positions the encoder was given. Those patches are never in the loss, so nothing constrains their predictions. This is why MAE figures always paste the visible patches back before showing you the result. It is easy to mistake for a bug.

The trivial baseline is 0.25 and training reaches 0.0000. Guessing a flat mid-grey for every pixel scores 0.25 on this data. Quoting a reconstruction loss without its trivial baseline is meaningless, in the same way that a contrastive loss is meaningless without its chance level.

The 16x saving comes from the encoder never seeing the hidden patches. Attention cost grows with the square of the token count. Keeping a quarter of the tokens costs a sixteenth. On a real ViT-Large at 75% masking, this is what turns a large pretraining run into an affordable one.

The mask token is a single learned vector, reused at every hidden position. Position is restored by adding the positional embedding back. So the decoder is told "something is here, and here is where here is", and nothing more.

Common mistakes

Computing the loss over every patch. The MAE paper measures this and reports it as worse. Supervising the visible patches lets the network spend capacity on copying rather than inferring.

Masking too little. At 15%, the rate that works for text, an image model interpolates from neighbours and learns little. MAE's ablation puts the optimum near 75%, far higher than anyone expected before the experiment.

Feeding mask tokens through the encoder. SimMIM does this and works; MAE does not and is much faster. Both are valid, but if you keep the mask tokens in the encoder you have given up the speed argument, so do it on purpose.

Forgetting to normalise the target patches. MAE reports better representations when each patch's target is normalised by its own mean and standard deviation. It removes low-frequency brightness that the model would otherwise waste capacity on.

Judging the method by how the reconstructions look. Sharper pixels do not mean better features. Evaluate with a linear probe or a fine-tune on a real task.

Try it yourself

Set MASK_RATIO = 0.25 and watch the loss fall faster while the task becomes easier and less informative. Then set it to 0.9375, leaving one visible patch, and see whether phase can still be recovered. The point where the task stops being solvable tells you how much redundancy this data has.

What to learn next

Researcher — Mathematics and papers.

The three founding designs

Bao, Dong, Piao and Wei (2021), BEiT: BERT Pre-Training of Image Transformers, arxiv.org/abs/2106.08254, predict discrete visual tokens produced by a separately trained discrete VAE, rather than pixels. Blockwise masking is used, and the objective is classification over a vocabulary. Reported: 83.2% ImageNet-1K top-1 for the base model against DeiT's 81.8%, and 86.3% for the large model, exceeding a supervised ViT-L trained on ImageNet-22K at 85.2%.

He, Chen, Xie, Li, Dollár and Girshick (2021), Masked Autoencoders Are Scalable Vision Learners, arxiv.org/abs/2111.06377, make two design decisions that define the method. The encoder processes only the visible patches; the decoder is small and receives mask tokens plus the encoded visible tokens. Masking is random at a high ratio, and the target is raw pixels. Reported: training accelerated by 3x or more, and ViT-Huge at 87.8% on ImageNet-1K using ImageNet-1K data alone.

Xie, Zhang, Cao, Lin, Bao, Yao, Dai and Hu (2021), SimMIM, arxiv.org/abs/2111.09886, show the elaborate parts are unnecessary. Random masking with a moderately large patch size of 32 pixels, direct regression on raw pixels, and a single linear prediction head match or beat the heavier alternatives. ViT-B reaches 83.8% fine-tuned; SwinV2-H reaches 87.1%; a 3-billion-parameter SwinV2-G reaches state-of-the-art results with 40x less data than previous practice.

Their disagreement is instructive. BEiT argues the target must be semantic; SimMIM shows raw pixels are enough; MAE shows the encoder should not see mask tokens at all. All three produce strong models, which suggests the essential ingredient is the difficulty of the task rather than the specific target.

Why the ratio must be so high

Language is discrete and information-dense, so masking 15% of tokens leaves a hard problem. Images are continuous and spatially redundant, so a low masking ratio is solvable by interpolation.

MAE's ablation sweeps the ratio and finds fine-tuning and linear-probe accuracy both peaking around 75%, with the linear-probe curve sharper. The framing in the paper is that a high ratio removes the redundancy that would otherwise permit a local solution, forcing a global one.

The asymmetric encoder-decoder then makes the high ratio a compute advantage rather than a cost. With 25% of tokens retained, encoder attention costs $(0.25)^2 = 1/16$ of full attention, which the developer block measures directly at toy scale.

Contrastive versus masked pretraining

They fail and succeed in different places, and the differences are consistent across studies:

  • Linear probing. Contrastive and self-distillation methods produce features that are linearly separable straight away. MAE's linear probe is notably weaker than its fine-tuned accuracy, because the representation is not organised around semantic classes.
  • Fine-tuning. MAE-style pretraining tends to win, particularly for large models and dense prediction.
  • Augmentation dependence. Contrastive learning depends heavily on the augmentation set. MAE works with almost none, only cropping, which makes it easier to move to a new domain.
  • Data scaling. Masked modelling shows less saturation at large model sizes, which is the argument in MAE's title.

Park et al. (2023), What Do Self-Supervised Vision Transformers Learn?, analyse the difference in terms of attention behaviour: contrastive methods produce more global, shape-biased attention while masked modelling produces more local, texture-sensitive attention, and the two are complementary. Several later methods combine them, including iBOT and DINOv2.

Convolutional networks and masking

Masking is awkward for convolutions. A convolution has no way to skip a hidden region, and a masked input is not a valid dense tensor, so the mask leaks into every downstream activation.

Woo et al. (2023), ConvNeXt V2, arxiv.org/abs/2301.00808, solve this by treating the masked input as a sparse tensor and using sparse convolutions during pretraining, then converting to dense convolutions for fine-tuning. Combining a ConvNeXt with a masked autoencoder naively gave, in their words, subpar performance. They diagnosed feature collapse, with many channels going inactive, and added Global Response Normalization to restore inter-channel competition. Results: 76.7% top-1 for a 3.7M Atto model, and 88.9% for a 650M Huge model on public data only.

The general lesson is that pretraining objective and architecture are coupled. Transplanting a recipe across architectures needs measurement, not assumption.

Newer directions

  • Predicting in feature space rather than pixel space. I-JEPA (Assran et al., 2023) predicts the representations of masked regions rather than their pixels, removing the pressure to model unpredictable low-level detail. It reaches strong linear-probe accuracy without hand-designed augmentations.
  • Predicting the outputs of a momentum teacher, as in data2vec (Baevski et al., 2022), unifies the objective across images, speech and text.
  • Combining masked and self-distillation objectives, as in iBOT and then DINOv2, currently gives the strongest general-purpose visual features.

The trajectory is away from pixel targets. Pixels contain a large amount of information that is genuinely unpredictable, and spending capacity on it is wasted.

What to learn next