Skip to main content

Classification with MLPs

Examples

From regression to classification

The difference between regression and classification with a network is the same as between linear and logistic regression: only the negative log-likelihood changes. For regression it comes from the normal distribution and gives the mean squared error on an output layer without activation. For classification it comes from the categorical distribution and gives the cross-entropy on a softmax output layer, as for the linear classification of MNIST.

Fitting MNIST

A network with one hidden layer of 100 relu neurons and ten outputs, fitted by AdamW on batches of 32 for 20 epochs.

Code
import numpy as np
import torch
import matplotlib.pyplot as plt
from sklearn.datasets import fetch_openml

mnist = fetch_openml(data_id=554, as_frame=False, parser="auto")
X_all = (mnist.data / 255.0).astype(np.float32)          # pixels scaled to [0, 1]
y_all = mnist.target.astype(np.int64)
X_train, y_train = torch.tensor(X_all[:60000]), torch.tensor(y_all[:60000])
X_test, y_test = torch.tensor(X_all[60000:]), torch.tensor(y_all[60000:])
torch.manual_seed(1)
net = torch.nn.Sequential(
    torch.nn.Linear(784, 100), torch.nn.ReLU(),
    torch.nn.Linear(100, 10),
)
opt = torch.optim.AdamW(net.parameters())
loader = torch.utils.data.DataLoader(
    torch.utils.data.TensorDataset(X_train, y_train), batch_size=32, shuffle=True)

for epoch in range(20):
    for xb, yb in loader:
        opt.zero_grad()
        torch.nn.functional.cross_entropy(net(xb), yb).backward()
        opt.step()

with torch.no_grad():
    train_acc = (net(X_train).argmax(dim=1) == y_train).float().mean()
    test_acc = (net(X_test).argmax(dim=1) == y_test).float().mean()
print(f"training accuracy {float(train_acc):.3f}, test accuracy {float(test_acc):.3f}")
1
Ten outputs, one per digit, without a softmax.
2
cross_entropy applies the softmax to the ten outputs and returns the mean negative log-likelihood of the labels.
training accuracy 0.998, test accuracy 0.977

The test error is between 2% and 3%, against about 7.5% for the linear classifier. The network has \(100 \cdot 785 + 10 \cdot 101 = 79510\) parameters, ten times as many as the linear classifier.

drawing some misclassified test images
with torch.no_grad():
    pred = net(X_test).argmax(dim=1).numpy()
wrong = np.where(pred != y_test.numpy())[0][:10]
fig, axes = plt.subplots(1, 10, figsize=(8.8, 1.3))
for ax, i in zip(axes, wrong):
    ax.imshow(X_test[i].reshape(28, 28), cmap="gray_r")
    ax.set(xticks=[], yticks=[], title=f"{int(y_test[i])} → {pred[i]}")
plt.show()
Figure 34.1: Test images the network gets wrong, with the true label and the prediction.