MNIST from scratch
Be able to train a digit classifier and analyse which digits it confuses.
Prerequisites
- DDataloaders, batches and epochsrequired
- DTrain a neural network in PyTorchrequired
Intuition
MNIST is 70 000 handwritten digits in 28×28 greyscale — the field's «hello world». A simple model reaches 97 %, a small CNN over 99 %.
But the accuracy is not the interesting part. The interesting part is which 1 % come out wrong.
The confusion matrix shows that. Row = the truth, column = the guess:
truth\guess 0 1 2 3 4 5 6 7 8 9
4 970 0 1 0 0 0 3 2 1 5
9 2 1 0 2 9 1 0 5 3 986
The row for 4 says: 970 correct, and most of the mistakes went to 9. That is not random — 4 and 9 look alike when written carelessly, and the network makes the same mistake a human would.
The common confusions in MNIST are exactly 4↔9, 3↔5, 7↔1 and 8↔3. A model making understandable mistakes is a good sign.
Code
import torch, torch.nn as nn
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
from sklearn.metrics import confusion_matrix, classification_report
import numpy as np
tf = transforms.Compose([transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,))])
train = datasets.MNIST("data", train=True, download=True, transform=tf)
test = datasets.MNIST("data", train=False, transform=tf)
tl = DataLoader(train, batch_size=128, shuffle=True)
vl = DataLoader(test, batch_size=512)
model = nn.Sequential(
nn.Conv2d(1, 16, 3, padding=1), nn.ReLU(), nn.MaxPool2d(2),
nn.Conv2d(16, 32, 3, padding=1), nn.ReLU(), nn.MaxPool2d(2),
nn.Flatten(), nn.Dropout(0.25), nn.Linear(32 * 7 * 7, 10),
)
opt = torch.optim.Adam(model.parameters(), lr=1e-3)
for epoch in range(3):
model.train()
for x, y in tl:
loss = nn.functional.cross_entropy(model(x), y)
opt.zero_grad(); loss.backward(); opt.step()
model.eval()
pred, truth = [], []
with torch.no_grad():
for x, y in vl:
pred += model(x).argmax(1).tolist(); truth += y.tolist()
print(f"accuracy {np.mean(np.array(pred) == np.array(truth)):.4f}") # ~0.988
# What is actually interesting
M = confusion_matrix(truth, pred)
np.fill_diagonal(M, 0)
for _ in range(4):
i, j = np.unravel_index(M.argmax(), M.shape)
print(f" {i} classified as {j}: {M[i, j]} times")
M[i, j] = 0
# 4 classified as 9: 9 times
# 3 classified as 5: 7 times
# 7 classified as 1: 6 times
# 8 classified as 3: 5 times
print(classification_report(truth, pred, digits=3))
The next step after the matrix: pull out the images that were actually misclassified and look at them. In MNIST a striking share of the mistakes are ones a human would hesitate over too — and a few are outright mislabelled in the original data. That sets a ceiling on how good any model can get, and it is worth knowing before you chase the last tenths.
Interactive
Five experiments on the same model. Run the base version, note the accuracy, and change one thing at a time.
| Change | Expected outcome |
|---|---|
Remove Normalize | slower convergence, slightly worse |
| Remove both conv layers (Linear only) | ~92 % instead of ~99 % |
Remove Dropout | better on training, marginally worse on test |
lr=1e-1 instead of 1e-3 | does not train at all, or very unstably |
| Train on 1 000 images instead of 60 000 | ~95 % — surprisingly good |
The last row is the most instructive: MNIST is too easy to separate good methods from bad. Almost everything works, which makes the dataset excellent for learning the tools and unsuitable for comparing models.
If you want a more honest test: run the same model on Fashion-MNIST (the same format, clothes instead of digits). The accuracy falls to around 91 %, and the differences between methods suddenly become visible.
And the real final exam: write a digit yourself on paper, photograph it, scale it to 28×28 greyscale and invert it. The model that got 99 % will probably get it wrong. That is domain shift in its purest form — the MNIST digits are centred, size-normalised and written with a particular kind of pen, and your image is none of those things.
Mastery means
- Trains a classifier on MNIST
- Reads a confusion matrix
- Analyses the mistakes rather than just the accuracy
Sign in to do the exercises and build your mastery up.
Sources
- PyTorch — tutorials (BSD-3) — BSD-3-Clause
- Dive into Deep Learning (CC BY-SA 4.0) — CC BY-SA 4.0
- Fashion-MNIST (MIT) — MIT