Few-shot image classification
With a handful of examples per class you stop training the network and start comparing descriptions, because averaging a few frozen feature vectors beats fine-tuning a model that would rather memorise.
- 12 min read
- 3 reading levels
- Updated
Read these first
On this page 7
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
Few-shot classification means learning a new class from a handful of photos, sometimes one.
You met a new colleague once, on Monday, for two minutes. On Thursday you spot her in a corridor from behind, at an angle, in different clothes. One meeting was enough.
That works because you already knew everything general about faces before you met her. Shape, lighting, angles, hair. The only new thing you stored was how her particular face differs from all the others.
A pretrained vision model is in the same position. The general knowledge is already inside it. Learning a new class is adding one short note, not re-reading the textbook.
Why the usual approach fails here
Fine-tuning has millions of adjustable numbers. Give it forty photos and it will find a way to score perfectly on those forty. It memorises the tea stain in the corner of one photo, and calls that "Asha".
You have all heard the equivalent story about the student who memorised last year's question paper. It works on that paper and nowhere else.
So few-shot work uses a different move. Do not train. Compare.
How comparing works
5 photos of Asha -> [ frozen model ] -> 5 descriptions -> average -> "Asha point"
5 photos of Ravi -> [ frozen model ] -> 5 descriptions -> average -> "Ravi point"
a new photo -> [ frozen model ] -> description
|
which point is it closest to?The average of a class's descriptions is called its prototype — one point standing for the whole class. Classifying is measuring distance to each prototype and taking the nearest.
Nothing was trained. No gradients, no epochs, no learning rate. Adding a new person is one more average, computed in a second.
The vocabulary you will meet
People write few-shot setups as N-way K-shot. N is how many classes you must tell apart. K is how many examples you get of each. "5-way 1-shot" means five classes, one photo each.
Guessing at random in a 40-way problem gets you 2.5 percent. Keep that number in mind, because few-shot scores look low until you compare them with chance.
Where you have already seen this
- Adding a new face to your phone's photo albums with two or three tagged pictures.
- A shop's checkout camera learning a new product from a few snaps.
- A security system that recognises a visitor from one enrolment photo.
None of these retrain a neural network while you wait. They all store a point and measure distance.
The honest part
Few-shot results are noisy. Which five photos you happen to get changes the score a lot. A paper reporting "62.3 percent" without a spread across many random draws is telling you very little.
And prototypes assume a class forms one tight cluster. A class like "damaged packaging" covers dents, tears and water stains, which do not sit together. One average for three different things is a poor summary.
Remember this
- Few-shot means compare against stored examples, not train on them.
- A prototype is the average description of a class, and nearest-prototype is a complete classifier.
- Always compare a few-shot score against random guessing, and always report the spread.
What to learn next
- Domain adaptation for vision — when your few examples come from a different camera than the ones ahead.
- Embeddings — the general idea of turning things into comparable points.
- CLIP — prototypes built from text, which drops K to zero.
Developer — Code and libraries.
Setup
pip install torch torchvision scikit-learnWritten and run against torch 2.13.0 (CPU), torchvision 0.28.0, scikit-learn 1.7.2. Uses fetch_olivetti_faces (400 photos, 40 people, about 4.5 MB) and the 44.7 MB ResNet-18 checkpoint. Runs in a few seconds after the downloads.
Nearest prototype, and a probe, at 1 / 2 / 5 shots
import torch, torch.nn as nn, torch.nn.functional as F, numpy as np
from sklearn.datasets import fetch_olivetti_faces
from sklearn.linear_model import LogisticRegression
from torchvision.models import resnet18, ResNet18_Weights
torch.manual_seed(0); np.random.seed(0)
faces = fetch_olivetti_faces()
x = torch.tensor(faces.images).unsqueeze(1).repeat(1, 3, 1, 1)
x = F.interpolate(x, size=112, mode="bilinear", align_corners=False)
mean = torch.tensor([0.485,0.456,0.406]).view(1,3,1,1)
std = torch.tensor([0.229,0.224,0.225]).view(1,3,1,1)
x = (x - mean) / std
y = torch.tensor(faces.target)
bb = resnet18(weights=ResNet18_Weights.DEFAULT); bb.fc = nn.Identity(); bb.eval()
with torch.no_grad():
f = torch.cat([bb(x[i:i+40]) for i in range(0, 400, 40)]) # 400 x 512, computed once
f = F.normalize(f, dim=1) # compare by direction
te = torch.cat([torch.arange(i*10+5, i*10+10) for i in range(40)]) # last 5 per person: never seen
print(f"{'shots':>6} {'nearest centroid':>18} {'logistic probe':>16}")
for k in (1, 2, 5):
tr = torch.cat([torch.arange(i*10, i*10+k) for i in range(40)])
proto = torch.stack([f[tr][y[tr] == c].mean(0) for c in range(40)]) # one prototype per person
proto = F.normalize(proto, dim=1)
nc = ((f[te] @ proto.T).argmax(1) == y[te]).float().mean().item()
lr = LogisticRegression(max_iter=3000).fit(f[tr], y[tr]).score(f[te], y[te])
print(f"{k:>6} {nc:>18.3f} {lr:>16.3f}") shots nearest centroid logistic probe
1 0.650 0.625
2 0.665 0.680
5 0.780 0.705Deterministic across runs with the seeds set.
What that table is worth
One photo per person, forty people, 65 percent correct. Random guessing is 2.5 percent. No training happened at all: one forward pass per image, one average per class, one matrix multiply to classify.
At five shots the plain average beats the fitted classifier. 0.780 against 0.705. Logistic regression has 40 x 512 weights to fit from 200 rows and starts to overfit them. The average has no parameters to overfit. This ordering flips once you have enough examples, which is exactly the point.
These are single draws, not measurements. The first K photos of each person is one arbitrary choice. A proper report repeats over many random draws and gives a confidence interval, as the researcher section describes.
Why not fine-tune instead?
Because it memorises. Same data, one shot per person, full fine-tuning:
import torch, torch.nn as nn, torch.nn.functional as F, numpy as np
from sklearn.datasets import fetch_olivetti_faces
from torchvision.models import resnet18, ResNet18_Weights
torch.manual_seed(0); np.random.seed(0)
faces = fetch_olivetti_faces()
x = torch.tensor(faces.images).unsqueeze(1).repeat(1, 3, 1, 1)
x = F.interpolate(x, size=112, mode="bilinear", align_corners=False)
mean = torch.tensor([0.485,0.456,0.406]).view(1,3,1,1)
std = torch.tensor([0.229,0.224,0.225]).view(1,3,1,1)
x = (x - mean) / std
y = torch.tensor(faces.target)
tr = torch.cat([torch.arange(i*10, i*10+1) for i in range(40)]) # ONE photo per person
te = torch.cat([torch.arange(i*10+5, i*10+10) for i in range(40)])
m = resnet18(weights=ResNet18_Weights.DEFAULT); m.fc = nn.Linear(512, 40)
opt = torch.optim.AdamW(m.parameters(), lr=1e-4)
for ep in range(8):
m.train()
perm = tr[torch.randperm(40)]
for i in range(0, 40, 10):
b = perm[i:i+10]
opt.zero_grad(); F.cross_entropy(m(x[b]), y[b]).backward(); opt.step()
m.eval()
with torch.no_grad():
tr_acc = (m(x[tr]).argmax(1) == y[tr]).float().mean().item()
te_acc = (torch.cat([m(x[te][i:i+40]) for i in range(0, 200, 40)]).argmax(1) == y[te]).float().mean().item()
print(f"epoch {ep+1}: train {tr_acc:.3f} held-out {te_acc:.3f}")epoch 1: train 0.025 held-out 0.035 epoch 2: train 0.125 held-out 0.035 epoch 3: train 0.200 held-out 0.090 epoch 4: train 0.350 held-out 0.100 epoch 5: train 0.525 held-out 0.095 epoch 6: train 0.625 held-out 0.140 epoch 7: train 0.725 held-out 0.170 epoch 8: train 0.775 held-out 0.180
Train accuracy climbing to 0.775 while held-out reaches 0.180. The model is learning the forty specific photographs. Nearest prototype, on the identical data and the identical backbone, reached 0.650 without a single gradient step.
Keep training and the gap widens. This is overfitting in its purest available form.
Line by line
F.normalize(f, dim=1) rescales every feature vector to unit length, so f[te] @ proto.T is a cosine similarity. Without it, image brightness and contrast leak into the distance and dominate it.
f[tr][y[tr] == c].mean(0) averages the K feature vectors of class c. At K=1 the prototype is that single vector, so nearest-prototype degenerates to nearest-neighbour.
bb.fc = nn.Identity() and bb.eval() are the same frozen-feature setup as the probing lesson. The features are computed once for all 400 images and reused for every value of K.
Common mistakes
Reporting one episode. A single choice of support photos gives a number with a standard deviation of several points. Fix: sample 200 to 600 random episodes and report mean with a 95 percent interval.
Letting a test photo into the prototypes. With ten photos per person it is easy to average all ten and then test on some of them. Fix: build support and query index lists first, assert they do not intersect.
Skipping normalisation. Unnormalised Euclidean distance on ResNet features ranks by vector magnitude as much as by direction. Fix: L2-normalise, or use cosine similarity.
Treating a multi-mode class as one prototype. "Defective" covering three unrelated defect types averages to a point representing none of them. Fix: split it into three classes, or store several prototypes per class and take the nearest.
Try it yourself
Wrap the prototype code in a loop over 300 random episodes. Each episode: sample 5 of the 40 people. Take K support photos and 5 query photos each, then record accuracy. Print the mean and the 2.5th and 97.5th percentiles. That interval is the honest version of the numbers above.
What to learn next
- Domain adaptation for vision — when your few examples come from a different camera than the ones ahead.
- Embeddings — the general idea of turning things into comparable points.
- CLIP — prototypes built from text, which drops K to zero.
Researcher — Mathematics and papers.
Prototypical networks, stated
Snell, Swersky and Zemel (2017), Prototypical Networks for Few-shot Learning (NeurIPS), define the prototype of class $c$ as
$$ p_c = \frac{1}{|S_c|} \sum_{(x_i, y_i) \in S_c} f_\phi(x_i) $$
- $S_c$ — the support set for class $c$, of size $K$.
- $f_\phi$ — the embedding network with parameters $\phi$.
- $p_c \in \mathbb{R}^{d}$ — the prototype.
and classify a query $x$ by
$$ P(y = c \mid x) = \frac{\exp(-\lVert f_\phi(x) - p_c \rVert_2^2)}{\sum_{c'} \exp(-\lVert f_\phi(x) - p_{c'} \rVert_2^2)} $$
- $\lVert \cdot \rVert_2^2$ — squared Euclidean distance in embedding space.
Their key observation: with squared Euclidean distance this is a linear classifier in embedding space. Expanding the square leaves terms linear in $f_\phi(x)$ plus a per-class constant. Prototypical networks are therefore a linear model whose weights are set by averaging rather than fitting. That explains the table above. They beat a fitted linear model when data is scarce, and lose when it is plentiful.
Euclidean distance mattered empirically. Snell et al. report a large gap over cosine distance. They attribute it to Euclidean being a Bregman divergence, whose optimal representative point is the cluster mean.
The correction that reshaped the field
Three papers argued that episodic meta-learning was largely unnecessary:
- Chen et al. (2019), A Closer Look at Few-shot Classification (ICLR), showed pretraining plus a cosine-distance classifier competitive with meta-learning. The gap narrows further with deeper backbones.
- Tian et al. (2020) (ECCV) showed a plain linear classifier on a well-trained embedding beating state-of-the-art meta-learners, with a further gain from self-distillation. Their title asks whether a good embedding is all you need.
- Dhillon et al. (2020), A Baseline for Few-Shot Image Classification (ICLR), added transductive fine-tuning at test time as a strong, under-reported baseline.
The practical consequence for a working engineer: spend your effort on the representation, not on the few-shot algorithm.
Evaluation protocol, and why published numbers disagree
The standard protocol samples episodes: draw $N$ classes, $K$ support and $Q$ query images per class, classify, repeat. Reported accuracy is the mean over 600 or more episodes, with a 95% confidence interval. On miniImageNet that interval is typically $\pm 0.4$ to $\pm 0.8$ points.
Three recurring comparability problems:
- Backbone. Conv-4 and ResNet-12 results are not comparable; the backbone often explains more variance than the method.
- Transduction. Some methods see all query images at once, through query-set batch normalisation or transductive fine-tuning. They have strictly more information than inductive ones.
- Base-class overlap. Any pretraining data overlapping the novel classes invalidates the result. This is a live problem for large web-scale backbones evaluated on standard benchmarks.
Zero-shot changes the framing
Radford et al. (2021), CLIP, produce class prototypes from text rather than images. Encode "a photo of a {class}" and use the resulting text embedding as $p_c$. K becomes zero. Zhang et al. (2022), Tip-Adapter, then combine text prototypes with a cache of few-shot image features. That gains few-shot accuracy without gradient training. That hybrid is text prototypes refined by a handful of images. It is the pragmatic default when a vision-language backbone is available.
Papers
- Snell et al., Prototypical Networks for Few-shot Learning, 2017 — arxiv.org/abs/1703.05175
- Vinyals et al., Matching Networks for One Shot Learning, 2016 — arxiv.org/abs/1606.04080
- Finn et al., Model-Agnostic Meta-Learning (MAML), 2017 — arxiv.org/abs/1703.03400
- Chen et al., A Closer Look at Few-shot Classification, 2019 — arxiv.org/abs/1904.04232
- Tian et al., Rethinking Few-Shot Image Classification, 2020 — arxiv.org/abs/2003.11539
- Zhang et al., Tip-Adapter, 2022 — arxiv.org/abs/2207.09519
What to learn next
- Domain adaptation for vision — when your few examples come from a different camera than the ones ahead.
- Embeddings — the general idea of turning things into comparable points.
- CLIP — prototypes built from text, which drops K to zero.