Most real classification problems have more than two classes: ten digits, dozens of document categories, thousands of product types, fifty thousand vocabulary words for next-token prediction. Today we extend logistic regression to $K$ classes. The result โ softmax regression (multinomial logistic regression) โ is also exactly the final layer of almost every neural network classifier.
The softmax function#
Give each class $k$ its own weight vector $\mathbf{w}_k$ and compute a score (logit) $z_k = \mathbf{w}_k^\top\mathbf{x}$. Convert scores to probabilities with the softmax:
Properties:
- Outputs are positive and sum to one โ a valid categorical distribution.
- It preserves order: the largest logit gets the largest probability.
- It is invariant to adding a constant to all logits (so one weight vector is redundant; regularisation resolves this).
- For $K = 2$, softmax reduces to the sigmoid of the difference of logits.
The name comes from being a smooth ("soft") version of the argmax. Scaling logits by a temperature $T$, $\text{softmax}(\mathbf{z}/T)$, makes the distribution sharper ($T < 1$) or flatter ($T > 1$) โ a knob we will meet again in knowledge distillation and language-model sampling.
Loss: categorical cross-entropy#
With a one-hot label $\mathbf{y}$ and predicted probabilities $\mathbf{p}$:
โ the negative log-probability of the correct class, i.e. the negative log-likelihood under a categorical model.
The gradient#
Remarkably, the gradient with respect to the logits is again "prediction minus target":
so for the weight matrix $\mathbf{W} \in \mathbb{R}^{d \times K}$ (columns $\mathbf{w}_k$):
The loss is convex in $\mathbf{W}$, so training finds the global optimum.
import numpy as np
from sklearn.datasets import load_digits
from sklearn.model_selection import train_test_split
X, y = load_digits(return_X_y=True) # 8x8 digit images, 10 classes
X = X / 16.0
X_tr, X_te, y_tr, y_te = train_test_split(X, y, test_size=0.25, random_state=0, stratify=y)
def softmax(Z):
Z = Z - Z.max(axis=1, keepdims=True) # numerical stability
E = np.exp(Z)
return E / E.sum(axis=1, keepdims=True)
K, d = 10, X.shape[1]
W, b = np.zeros((d, K)), np.zeros(K)
Y = np.eye(K)[y_tr] # one-hot labels
for epoch in range(500):
P = softmax(X_tr @ W + b)
G = (P - Y) / len(X_tr) # dJ/dZ
W -= 0.5 * (X_tr.T @ G + 1e-4 * W)
b -= 0.5 * G.sum(axis=0)
if epoch % 100 == 0:
loss = -np.log(P[np.arange(len(y_tr)), y_tr] + 1e-12).mean()
print(epoch, round(loss, 4))
acc = (softmax(X_te @ W + b).argmax(1) == y_te).mean()
print("test accuracy:", round(acc, 3))A simple linear softmax model reaches roughly 95% accuracy on this small digits dataset.
Alternative strategies: reducing multiclass to binary#
Any binary classifier can be extended to $K$ classes:
One-vs-Rest (OvR)#
Train $K$ binary classifiers, each separating one class from all others; predict the class with the highest score. Simple and scalable, but each classifier faces an imbalanced problem, and scores from separately trained classifiers may not be comparable.
One-vs-One (OvO)#
Train a classifier for every pair of classes โ $K(K-1)/2$ of them โ and predict by majority vote. Each classifier trains on less data, which suits algorithms that scale poorly with dataset size (such as kernel SVMs), but the number of models grows quadratically.
| Strategy | Models trained | Probabilities | Typical use |
|---|---|---|---|
| Softmax (multinomial) | 1 | Coherent, sum to 1 | Default for linear models and neural nets |
| One-vs-Rest | $K$ | Need normalisation | Linear SVMs, very many classes |
| One-vs-One | $K(K-1)/2$ | From votes | Kernel SVMs |
Multi-label is different#
In multi-label classification, each example can belong to several classes at once (a photo containing both "beach" and "dog"). Softmax is wrong here because it forces classes to compete. Use independent sigmoids โ one binary logistic output per label โ with binary cross-entropy summed over labels.
Evaluating multiclass models#
- Confusion matrix โ which classes are confused with which.
- Per-class precision/recall/F1, combined with macro averaging (treats classes equally โ good for imbalanced data) or micro/weighted averaging.
- Top-k accuracy โ whether the true class is among the $k$ most probable (standard for ImageNet with 1,000 classes).