Optimisers, Schedulers and the Training Loop
Logits, softmax and picking the right loss
PyTorch loss functions expect raw scores, not probabilities — and feeding them the wrong one produces a model that trains, badly, without a single error message.
- 7 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
A logit is a model's raw score before it gets converted into a probability — and PyTorch's loss functions want the raw score, not the probability.
Think of judges in a singing contest holding up raw marks: 8.9, 7.2, 9.5. Later, someone converts those marks into a share of the total, so the shares add up to a whole. The raw marks and the shares carry the same ranking, in two different forms.
A model's last layer produces raw marks. The converter that turns them into shares — into percentages that add up to 100 — is called softmax.
Why this trips people up
PyTorch's main classification loss does the converting internally, as part of the loss. It is built to receive raw marks.
So if you convert first, the loss converts again. Converting twice does not crash. It does not print a warning. It produces a model that learns — slowly, weakly, topping out worse than it should. That combination of "runs fine" and "quietly worse" makes this one of the most common bugs in beginner PyTorch code.
How it works
raw scores ──────────────→ [ loss: converts, then scores ] ✓ right
raw scores → [ convert ] → [ loss: converts AGAIN ] ✗ quiet damageThe rule to memorise: the model outputs raw scores; the loss function owns the converting.
Where you have seen this
Think of "is this spam?", "which song is this?", "what does this photo show?". Every one of those features ends in the same pattern: raw scores inside the machine, probabilities shown to the human. The percentage you see on screen was converted at the last moment, for display — not inside the training loss.
Remember this
- Logits are raw scores. Softmax turns them into probabilities.
- PyTorch's classification losses expect logits and convert internally.
- Converting before the loss is a silent bug: the model trains, but worse.
What to learn next
- Writing a validation loop that reports the truth — measuring models without fooling yourself.
- Custom loss functions in PyTorch — when the built-in losses stop being enough.
- Loss functions — the conceptual tour behind these APIs.
Developer — Code and libraries.
Setup
pip install torchOutputs captured with torch 2.5.1 on CPU; these numbers are deterministic.
The bug, measured
import torch
import torch.nn as nn
logits = torch.tensor([[ 3.0, -1.0, 0.5],
[ 0.2, 2.5, -0.3]]) # raw model outputs: two samples, three classes
targets = torch.tensor([0, 1]) # class indices, dtype long
ce = nn.CrossEntropyLoss()
print(f"correct, raw logits in: {ce(logits, targets).item():.4f}")
print(f"bug, softmax before the loss: {ce(logits.softmax(dim=1), targets).item():.4f}")
z = torch.tensor([ 2.0, -1.0, 0.3]) # one output per sample: binary case
t = torch.tensor([ 1.0, 0.0, 1.0]) # binary targets, dtype float
print(f"BCEWithLogitsLoss on raw z: {nn.BCEWithLogitsLoss()(z, t).item():.4f}")
print(f"BCELoss on sigmoid(z): {nn.BCELoss()(z.sigmoid(), t).item():.4f}")correct, raw logits in: 0.1225 bug, softmax before the loss: 0.6285 BCEWithLogitsLoss on raw z: 0.3315 BCELoss on sigmoid(z): 0.3315
Same predictions, same targets — and the softmax-first version reports a loss five times higher. Worse than the wrong number: probabilities are squashed between 0 and 1, so after the loss re-converts them, every prediction looks lukewarm. Confidence can never grow past a ceiling, gradients shrink, and training plateaus early while looking otherwise healthy.
Walkthrough
CrossEntropyLoss = log-softmax + negative log likelihood, fused. The fusion is not only convenience — computing softmax and log separately invites numerical overflow for large logits, and the fused version sidesteps it. That is why the API insists on logits.
Targets are class indices, dtype long. Shape (N,) with values 0..C-1, not one-hot vectors. Handing it floats raises a dtype error — covered with its friends in dtype mismatch errors.
The binary twins. BCEWithLogitsLoss is sigmoid + binary cross-entropy fused; BCELoss expects you to sigmoid first. The output above shows they agree when used correctly. Prefer the fused one: same answer, better numerical safety, one less chance to forget.
At inference time, apply softmax when a human needs probabilities. For picking the winning class, skip it: logits.argmax(dim=1) — softmax never changes the ranking.
Choosing the loss
| Task | Model output shape | Loss | Targets |
|---|---|---|---|
| Multi-class (one right answer) | (N, C) logits | CrossEntropyLoss | (N,) long indices |
| Binary | (N,) logits | BCEWithLogitsLoss | (N,) float 0/1 |
| Multi-label (many can be true) | (N, C) logits | BCEWithLogitsLoss | (N, C) float 0/1 |
| Regression | (N, d) values | MSELoss / HuberLoss | (N, d) float |
Multi-label deserves the highlight: CrossEntropyLoss cannot express "this photo has a dog and a ball", because softmax forces classes to compete for one budget. Independent sigmoids, via BCEWithLogitsLoss, let each label stand alone.
Common mistakes
softmax (or nn.Softmax) as the model's last layer, then CrossEntropyLoss. The measured bug above. Delete the softmax from the model.
sigmoid in the model, then BCEWithLogitsLoss. The binary edition of the same double-conversion. Either drop the sigmoid, or switch the loss to BCELoss — dropping the sigmoid is the better fix.
log_softmax into CrossEntropyLoss. The correct pairing for log-softmax outputs is NLLLoss. Feeding log-probabilities into a loss that log-softmaxes again shifts every loss value and weakens gradients, in the same quiet way.
Binary classification with two output neurons and mismatched machinery. Two logits with CrossEntropyLoss works; one logit with BCEWithLogitsLoss works. Mixing halves of the two setups produces shape errors, or worse, silent nonsense.
Try it yourself
Extend the race from choosing an optimiser: add nn.Softmax(dim=1) as a final layer and retrain. Compare the loss after 200 steps against the clean model, and note that nothing ever crashed.
What to learn next
- Writing a validation loop that reports the truth — measuring models without fooling yourself.
- Custom loss functions in PyTorch — when the built-in losses stop being enough.
- Loss functions — the conceptual tour behind these APIs.
Researcher — Mathematics and papers.
The quantities
For logits $z \in \mathbb{R}^C$ and true class $y$:
$$\mathrm{softmax}(z)_c = \frac{e^{z_c}}{\sum_{j=1}^{C} e^{z_j}}, \qquad \mathcal{L}_{CE}(z, y) = -\log \mathrm{softmax}(z)_y = -z_y + \log \sum_{j} e^{z_j}$$
- $C$ — number of classes; $z_c$ — the logit for class $c$; $y$ — the target index.
The second form is what PyTorch computes, using the log-sum-exp identity $\log\sum_j e^{z_j} = m + \log\sum_j e^{z_j - m}$ with $m = \max_j z_j$, which keeps every exponent non-positive and prevents overflow. Softmax-then-log without this shift produces inf - inf = nan for logits around $\pm 90$ in float32.
Why double softmax flattens gradients
The gradient of cross-entropy with respect to logits is famously clean:
$$\frac{\partial \mathcal{L}_{CE}}{\partial z_c} = \mathrm{softmax}(z)_c - \mathbb{1}[c = y]$$
- $\mathbb{1}[\cdot]$ — the indicator function: 1 when the condition holds, else 0.
Feed probabilities $p$ instead of logits and the loss becomes $\mathcal{L}(p) = -p_y + \log\sum_j e^{p_j}$ with every $p_j \in [0, 1]$: the effective "logit gap" is capped at 1, so the minimum achievable loss is bounded away from zero and gradient magnitudes shrink accordingly — the plateau observed in practice.
Refinements in current use
- Label smoothing (Szegedy et al., 2016, Rethinking the Inception Architecture): replace the hard target with $(1-\alpha)$ on the true class and $\alpha/C$ elsewhere;
CrossEntropyLoss(label_smoothing=0.1). Improves calibration (Müller et al., 2019, When Does Label Smoothing Help?). - Focal loss (Lin et al., 2017, Focal Loss for Dense Object Detection): scales CE by $(1 - p_y)^\gamma$ to down-weight easy examples under extreme class imbalance.
- Soft targets: since v1.10,
CrossEntropyLossalso accepts a full probability distribution as target (shape(N, C)float), which is how knowledge distillation losses are written.
Statistical reading
Minimising CE is maximum likelihood under a categorical model; minimising it against a fixed target distribution $q$ minimises $\mathrm{KL}(q \,|\, \mathrm{softmax}(z))$ up to the constant entropy of $q$. BCE is the Bernoulli special case; the sigmoid is the two-class softmax with one logit pinned to zero.
What to learn next
- Writing a validation loop that reports the truth — measuring models without fooling yourself.
- Custom loss functions in PyTorch — when the built-in losses stop being enough.
- Loss functions — the conceptual tour behind these APIs.