Md. Asif Uddin
    Problem I.4.B07

    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 (B,T,∣V∣)(B, T, |V|). 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 B=4B = 4, sequence length T=512T = 512, vocabulary ∣V∣=32,000|V| = 32{,}000. Targets are token indices, shape (B,T)(B, T). 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 33; the loss is rank 00. 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.

    logits:(B,T,∣V∣)=(4,512,32000),targets:(B,T)=(4,512)\text{logits} : (B, T, |V|) = (4, 512, 32000), \qquad \text{targets} : (B, T) = (4, 512)

    The shapes do not match, and they should not. The targets are indices, not one-hot vectors: integers in [0,∣V∣)[0, |V|), one per position. A one-hot target of shape (4,512,32000)(4, 512, 32000) would hold the same information in 65,536,00065{,}536{,}000 numbers instead of 2,0482{,}048, of which all but 2,0482{,}048 are zero. Frameworks index rather than multiply for exactly this reason.

    Step 2 — the element count.

    B×T×∣V∣=4×512×32,000=65,536,000B \times T \times |V| = 4 \times 512 \times 32{,}000 = 65{,}536{,}000

    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:

    (4,512,32000)⟶(2048, 32000),(4,512)⟶(2048,)(4, 512, 32000) \longrightarrow (2048,\ 32000), \qquad (4, 512) \longrightarrow (2048,)

    2048=4×5122048 = 4 \times 512, 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 2,0482{,}048 rows, take the row’s 32,00032{,}000 logits and its one target index and produce one number:

    (2048,32000)×(2048,)⟶(2048,)(2048, 32000) \times (2048,) \longrightarrow (2048,)

    The vocabulary axis is gone. This is where −zc+log⁡∑jezj-z_c + \log\sum_j e^{z_j} is evaluated per row — never by forming probabilities first, for the reasons of I.4.B04.

    Step 5 — the reduction. (2048,)⟶()(2048,) \longrightarrow (), rank 00. This is the one number ∇θ\nabla_\theta is taken of, and Step 7 shows it is the step most often got wrong.

    The complete trace.

    StageShapeValues
    logits(4,512,32000)(4, 512, 32000)65,536,00065{,}536{,}000
    targets(4,512)(4, 512)2,0482{,}048
    flattened logits(2048,32000)(2048, 32000)65,536,00065{,}536{,}000
    flattened targets(2048,)(2048,)2,0482{,}048
    per-token loss(2048,)(2048,)2,0482{,}048
    reduced loss()()11

    Step 6 — memory.

    fp32:65,536,000×4 bytes=262.1 MB\text{fp32:}\quad 65{,}536{,}000 \times 4\ \text{bytes} = 262.1\ \text{MB}bf16:65,536,000×2 bytes=131.1 MB\text{bf16:}\quad 65{,}536{,}000 \times 2\ \text{bytes} = 131.1\ \text{MB}

    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 524.3524.3 MB in fp32 and 262.1262.1 MB in bf16.

    This single tensor is often the largest in the model. At B=4B = 4 it already rivals the activations of an entire transformer block, and it grows linearly in both BB and ∣V∣|V|. It is why the loss is computed in chunks over the batch for large vocabularies, and why ∣V∣|V| appears in memory budgets as prominently as depth does.

    Step 7 — masking, and the reduction it changes. With 15%15\% padding:

    total positions=2,048,valid positions=round(2048×0.85)=1,741\text{total positions} = 2{,}048, \qquad \text{valid positions} = \text{round}(2048 \times 0.85) = 1{,}741

    Padded positions contribute a loss of 00, having been masked. But there are two different things one can then divide by.

    Mean over all positions. 12048∑iℓi\displaystyle \frac{1}{2048}\sum_i \ell_i — divides by the number of slots.

    Mean over valid positions. 11741∑iℓi\displaystyle \frac{1}{1741}\sum_i \ell_i — divides by the number of real tokens.

    The ratio is

    20481741=1.1763\frac{2048}{1741} = 1.1763

    so the first understates the loss by 17.6%17.6\% 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 17.6%17.6\% 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

    (4,512,32000)→(2048,32000)→(2048,)→()(4, 512, 32000) \to (2048, 32000) \to (2048,) \to ()

    Logits hold 65,536,00065{,}536{,}000 values: 262.1262.1 MB in fp32, 131.1131.1 MB in bf16, and double that if probabilities are materialised too.

    With 15%15\% padding, 1,7411{,}741 of 2,0482{,}048 positions are valid. Reducing over all positions rather than valid ones understates the loss and the gradient by a factor of 2048/1741=1.17632048/1741 = 1.1763, or 17.6%\mathbf{17.6\%}.

    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 1.17631.1763 ratio.

    Executed in CI. The digits above are the digits it printed.

    Check — sanity

    The reshape conserves elements. 4×512×32000=65,536,0004 \times 512 \times 32000 = 65{,}536{,}000 and 2048×32000=65,536,0002048 \times 32000 = 65{,}536{,}000. A reshape that changed the count would be a copy or a truncation, not a reshape.

    Rank falls exactly twice. Rank 3→23 \to 2 by the flatten, 2→12 \to 1 by indexing away the vocabulary, 1→01 \to 0 by the reduction. Three drops, three identifiable operations, no rank lost silently.

    bf16 is exactly half of fp32. 262.1/131.1=2.0262.1 / 131.1 = 2.0, as two bytes against four requires.

    The masking ratio bounds correctly. 1≤2048/1741≤1/0.85=1.17651 \le 2048/1741 \le 1/0.85 = 1.1765. The computed 1.17631.1763 sits just inside, the difference being the rounding of 1741.1741. At zero padding the ratio would be exactly 11 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 (2048,32000)(2048, 32000) rather than (2048,)(2048,). 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 KK independent sigmoids with binary cross-entropy is the right structure instead. The shape (2048,32000)(2048, 32000) survives but the reduction is over ∣V∣|V| as well as over positions, and the loss is a sum of 32,00032{,}000 binary terms per position rather than one categorical term.

    Variation

    Recompute every figure for ∣V∣=128,000|V| = 128{,}000 at the same BB and TT. State the new fp32 memory, and find the batch size at which the logit tensor alone exceeds 1010 GB.

    Draws on