From a logit tensor to one scalar, with every shape named
shape▲▲△Memory to 1 d.p.; counts exact.
STATEMENT
A language model emits logits of shape . Trace every shape on the path from that tensor to the single number the optimiser differentiates. Compute the memory the logits occupy, and determine what masking does to the reduction — including the size of the error if it is done wrongly.
GIVEN
Batch , sequence length , vocabulary . Targets are token indices, shape . Fifteen per cent of positions are padding and must be excluded. Cross-entropy is applied per position and then reduced to one scalar.
FIND
The shape after each stage; the element count and memory of the logits in fp32 and bf16; and the two possible reductions under masking, with the discrepancy between them.
STRATEGY
Follow the rank down. The logits are rank ; the loss is rank . Two things remove a rank — indexing away the vocabulary axis, and reducing away the position axes — and the whole trace is deciding where each happens.
SOLUTION
Step 1 — the incoming pair.
The shapes do not match, and they should not. The targets are indices, not one-hot vectors: integers in , one per position. A one-hot target of shape would hold the same information in numbers instead of , of which all but are zero. Frameworks index rather than multiply for exactly this reason.
Step 2 — the element count.
Step 3 — flatten the position axes. Cross-entropy treats every position independently, so the batch and time axes carry no meaning for it and can be merged:
, and the element count is unchanged — a reshape moves no data. Every framework’s cross-entropy expects exactly this two-dimensional form, which is why calling it on rank-3 logits requires an explicit reshape.
Step 4 — the per-position loss. For each of the rows, take the row’s logits and its one target index and produce one number:
The vocabulary axis is gone. This is where is evaluated per row — never by forming probabilities first, for the reasons of I.4.B04.
Step 5 — the reduction. , rank . This is the one number is taken of, and Step 7 shows it is the step most often got wrong.
The complete trace.
| Stage | Shape | Values |
|---|---|---|
| logits | ||
| targets | ||
| flattened logits | ||
| flattened targets | ||
| per-token loss | ||
| reduced loss |
Step 6 — memory.
And if softmax probabilities are materialised as a separate tensor of the same shape — which the naive route of I.4.B04 requires — that doubles to MB in fp32 and MB in bf16.
This single tensor is often the largest in the model. At it already rivals the activations of an entire transformer block, and it grows linearly in both and . It is why the loss is computed in chunks over the batch for large vocabularies, and why appears in memory budgets as prominently as depth does.
Step 7 — masking, and the reduction it changes. With padding:
Padded positions contribute a loss of , having been masked. But there are two different things one can then divide by.
Mean over all positions. — divides by the number of slots.
Mean over valid positions. — divides by the number of real tokens.
The ratio is
so the first understates the loss by relative to the second.
Why this matters more than it looks. The reported loss is wrong by a fixed factor, which is merely embarrassing. The gradient is wrong by the same factor, which is an unintended reduction in effective learning rate. And the factor is not fixed across batches: it depends on how much padding each batch happens to contain, so the effective learning rate fluctuates from step to step with the batch’s sequence-length distribution. Sorting examples by length into buckets — done for speed — changes the padding fraction systematically, and so silently changes the learning rate schedule.
Dividing by the valid count is the correct choice, and it is not the default in every framework.
Answer
Logits hold values: MB in fp32, MB in bf16, and double that if probabilities are materialised too.
With padding, of positions are valid. Reducing over all positions rather than valid ones understates the loss and the gradient by a factor of , or .
Check — numeric · i-4-b07-loss-shapes.py
logits = B * T * V
valid = round(B * T * (1 - PAD_FRACTION))
print(f"ratio {total / valid:.4f}")Prints the full shape trace, both memory figures, and the ratio.
Executed in CI. The digits above are the digits it printed.
Check — sanity
The reshape conserves elements. and . A reshape that changed the count would be a copy or a truncation, not a reshape.
Rank falls exactly twice. Rank by the flatten, by indexing away the vocabulary, by the reduction. Three drops, three identifiable operations, no rank lost silently.
bf16 is exactly half of fp32. , as two bytes against four requires.
The masking ratio bounds correctly. . The computed sits just inside, the difference being the rounding of At zero padding the ratio would be exactly and the two reductions would agree — which is why this bug is invisible on fixed-length data.
Where this breaks
The trace assumes every position has exactly one target. Two common cases break it, and both are worth recognising by their shapes.
Soft targets. With label smoothing or distillation the target is a full distribution, shape rather than . The indexing step of Step 4 becomes a contraction over the vocabulary axis, the targets now cost the same memory as the logits, and the total doubles again.
Multi-label. When a position may carry several correct classes, softmax is wrong outright — it forces the outputs to sum to one, and independent sigmoids with binary cross-entropy is the right structure instead. The shape survives but the reduction is over as well as over positions, and the loss is a sum of binary terms per position rather than one categorical term.
Variation
Recompute every figure for at the same and . State the new fp32 memory, and find the batch size at which the logit tensor alone exceeds GB.