skip to content

Why does PyTorch's nn.CrossEntropyLoss expect raw logits, not softmax output?

level: middleimportance: must knowfreq 62%

answer

  1. the loss already has the activation in it
  2. log-sum-exp needs the unnormalized scores
  3. apply it twice and the distribution flattens
  4. the binary twin is the with-logits one
  5. softmax belongs at inference, not in the loss

basics

~10 s

nn.CrossEntropyLoss applies log_softmax internally, so it must receive unnormalized scores. Feeding it softmax probabilities applies softmax twice, which flattens the distribution, shrinks gradients and slows or stalls training. No error is raised.

solid answer

~50 s

`nn.CrossEntropyLoss` is a fused `LogSoftmax` plus `NLLLoss`. The fusion is deliberate: computing the log of a softmax directly is numerically unstable for large scores, so PyTorch uses the log-sum-exp trick internally, which requires the raw logits. If you append a softmax to your model and hand the probabilities in, the loss softmaxes them again. Probabilities are already in [0, 1], so the second softmax produces a much flatter distribution — the loss is larger, the gradients are smaller, and the model trains slowly or plateaus, with no exception raised anywhere. The binary counterpart is the same story: `nn.BCEWithLogitsLoss` fuses a sigmoid with binary cross-entropy and takes logits, while `nn.BCELoss` is the version that takes probabilities. If you genuinely have log-probabilities already, `nn.NLLLoss` is the loss to pair with an explicit `nn.LogSoftmax`. Apply softmax or sigmoid yourself only at inference, to report probabilities.

code

python · 9 lines
python
import torch
from torch import nn

logits = torch.tensor([[2.0, 0.5, 0.1]])
target = torch.tensor([0])
ce = nn.CrossEntropyLoss()

print(ce(logits, target))                          # correct: raw logits
print(ce(torch.softmax(logits, dim=1), target))    # double softmax: higher loss

go deeper

for a junior

Recall the pairing rule: with nn.CrossEntropyLoss the model's last layer is a plain Linear, no softmax. Say that the loss applies the softmax for you.

for a middle

Explain the fusion — LogSoftmax plus NLLLoss — and why the log-sum-exp trick needs unnormalized inputs. Describe what double softmax does to gradients rather than just calling it wrong.

for a senior

Spot the bug from a symptom, such as a loss stuck near log(C), and know the full pairing matrix including BCEWithLogitsLoss, pos_weight for imbalance, and where soft-label targets fit.

for a principal

Treat the model-to-loss contract as an interface worth enforcing: decide whether models in your codebase always emit logits, keep the activation out of the module, and add a shape-and-range assertion or test so the double-softmax bug cannot ship silently.

## What CrossEntropyLoss actually computes `nn.CrossEntropyLoss` is documented as combining `nn.LogSoftmax` and `nn.NLLLoss` in a single class. Given an input of shape `(N, C)` — one row of `C` raw scores per sample — and a target of `N` integer class indices, it computes the log-softmax over the class dimension and then picks out the negative log-probability of the true class, averaging over the batch by default (`reduction='mean'`). Because the log-softmax step lives *inside* the loss, the contract on its input is unambiguous: give it the raw output of your final `nn.Linear`. No activation on the output layer. ## Why the fusion exists Doing it in two steps naively means exponentiating scores, normalizing, then taking a log. With large positive logits the exponentials overflow to infinity; with very negative ones the probability underflows to zero and the log becomes `-inf`. The fused implementation uses the log-sum-exp identity — subtract the row maximum before exponentiating — so the computation stays in a safe range for any input magnitude. That stability trick is only possible if the function can see the unnormalized scores. Hand it probabilities and you have already thrown away the information it needed. ## The double-softmax bug This is the archetypal PyTorch silent-wrong-answer. A model built by someone coming from a framework where the output layer carries the activation ends with `nn.Softmax(dim=1)`, and the loss is `nn.CrossEntropyLoss`. Everything runs. The shapes are right, the loss is a positive finite number, the accuracy even improves — just far more slowly than it should, and it converges to something worse. The mechanism: softmax of a probability vector is still a valid distribution, but far closer to uniform, because the inputs are now squeezed into [0, 1] instead of spanning several units. The gradient signal that distinguishes a confident correct prediction from a confident wrong one is compressed toward zero. The tell in practice is a training loss that drops quickly to a value near `log(C)` and then crawls. ## The binary twin The same trap has a binary form. `nn.BCEWithLogitsLoss` fuses a sigmoid with binary cross-entropy and is numerically stable for large-magnitude logits; `nn.BCELoss` expects values already in [0, 1] and will produce infinities if a probability is exactly 0 or 1. The PyTorch documentation explicitly recommends the with-logits version over a `Sigmoid` layer followed by `BCELoss`. `nn.BCEWithLogitsLoss` also carries a `pos_weight` argument that reweights the positive class, which is the standard lever for class imbalance in multi-label problems. A related mismatch: `nn.NLLLoss` and `nn.KLDivLoss` both expect **log**-probabilities as input, not probabilities. Passing plain probabilities to those is the mirror-image version of the same bug. ## Targets and shapes For `nn.CrossEntropyLoss` the input is `(N, C)` (or `(N, C, d1, d2, ...)` for per-pixel classification), and the target is either: - **class indices**: shape `(N,)`, dtype `long`, values in `[0, C)` — the usual case, and the only form that supports `ignore_index`; or - **class probabilities**: the same shape as the input, dtype float — used for soft labels and distillation. Other useful arguments are `weight` for per-class reweighting, `label_smoothing` for softening the targets, and `reduction` to switch between `'mean'`, `'sum'` and `'none'`. Note that the class dimension is dimension 1, not the last dimension — sequence models often need a `transpose` or a flatten to `(N*T, C)` before the loss. ## At inference Since the model no longer ends in a softmax, remember that its outputs are logits. For a predicted class you do not need the softmax at all — `logits.argmax(dim=1)` gives the same answer because softmax is monotonic. You only need `torch.softmax(logits, dim=1)` when you want calibrated probabilities to show a user or to threshold. For the binary case it is `torch.sigmoid(logits)`. ## How to detect it in someone else's code Grep the model definition for a final `Softmax`, `LogSoftmax` or `Sigmoid` and check which loss consumes it. The valid pairings are: raw logits with `CrossEntropyLoss` or `BCEWithLogitsLoss`; `LogSoftmax` with `NLLLoss`; `Sigmoid` with `BCELoss` (legal but numerically inferior). Anything else is a bug.

  • What shape and dtype does nn.CrossEntropyLoss expect for its target?
    Either class indices of shape (N,) with dtype long and values in [0, C), which is the usual form and the only one supporting ignore_index; or class probabilities with the same float shape as the input, used for soft labels and distillation. The class dimension of the input is dimension 1, so sequence outputs usually need a transpose or a flatten first.
  • When would you use nn.NLLLoss instead of nn.CrossEntropyLoss?
    When the model genuinely emits log-probabilities already — for example a final nn.LogSoftmax you need for another purpose, or a mixture model whose components you combine in log space. NLLLoss just indexes the true class's log-probability and negates it; pairing it with LogSoftmax is mathematically identical to CrossEntropyLoss on logits.
  • What does pos_weight do in nn.BCEWithLogitsLoss?
    It multiplies the positive-class term of the loss, so rare positives count more. It is the standard handling for imbalanced binary or multi-label targets, and is usually set near the ratio of negative to positive examples per label. Unlike the weight argument, which scales whole examples, pos_weight scales only the positive half of each term.

saying these in an interview costs you the question

  • Adding a Softmax layer to the model and using CrossEntropyLoss
  • Saying CrossEntropyLoss will error on normalized inputs
  • Passing raw probabilities to nn.NLLLoss instead of log-probabilities
  • Preferring Sigmoid plus BCELoss over BCEWithLogitsLoss
  • Applying softmax before argmax to get a predicted class

context