Skip to content
AI-grafen
DAI developerDeep learning· about 45 min· fundamentals that rarely change· verified 2026-09-20· EN

MNIST from scratch

Be able to train a digit classifier and analyse which digits it confuses.

Prerequisites

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.

ChangeExpected outcome
Remove Normalizeslower convergence, slightly worse
Remove both conv layers (Linear only)~92 % instead of ~99 %
Remove Dropoutbetter on training, marginally worse on test
lr=1e-1 instead of 1e-3does 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

All the sources and licences