Transformer Architecture Q&A

Course Content

Transformer Architecture Q&A

6 sections · 60 lessons

What is GPT output tensor shape per forward pass, and why that shape?


What you need to know

Why each axis exists

  • Batch — independent sequences processed together.
  • Sequence length — the causal mask means position t has seen tokens 0..t, so it can predict token t+1. Every position is a separate prediction.
  • Vocabulary — the LM head gives one score per vocabulary entry; softmax turns them into a probability for each possible next token.

These are logits, not probabilities. Softmax is applied inside the loss function (for numerical stability) or by the sampler at generation time.

How training uses it

Text
logits: (B, T, V)      targets = ids shifted left by one: (B, T)loss = cross_entropy(logits.view(B*T, V), targets.view(B*T))

All B × T predictions are trained in parallel. This is the main reason decoder-only pretraining is efficient.

How big it gets

Text
GPT-2 small, batch 4 × 1,024 × 50,257  = 206M numbers → 412 MB fp16, 824 MB fp32Llama-3-8B, batch 8 × 8,192 × 128,256  = 8.4B numbers → 33.6 GB in fp32

The second line is more memory than the 8B model's weights. Training code therefore avoids materialising all logits at once: it computes the LM head and cross-entropy in chunks, or uses fused "linear plus cross-entropy" kernels (for example Cut Cross-Entropy, Wijmans et al., 2024) that never store the full tensor.

At inference

Only logits[:, -1, :] is needed, because earlier positions are already decided. Serving engines slice the hidden state to the final position before the LM head, so the output is (B, 1, V). In decode with a KV cache, T is 1 anyway.

A real-life example

A team fine-tunes an 8B chatbot on long support transcripts at 8K context. The run crashes with out-of-memory errors even though the weights, gradients and optimiser state fit. Profiling shows the spike at the loss: batch 8 × 8,192 positions × 128,256 vocabulary entries in fp32, 33.6 GB for one tensor, plus the same again for its gradient.

They switch to a chunked cross-entropy that processes 1,024 positions at a time. Peak memory drops by tens of gigabytes, the training loss is unchanged, and the run fits. At serving time the model code already computed logits only for the last position — which is why nobody noticed the problem during inference testing.

Follow-up questions to expect

  • "Why logits and not probabilities?" — Cross-entropy combines log-softmax and the loss in one numerically stable step; exposing logits also lets the sampler apply temperature and filters before softmax.
  • "What is the shape during decode with a KV cache?" — (B, 1, V): one new position per sequence.
  • "Why does vocabulary size matter for cost?" — The LM head is a d × V matmul at every position that needs logits; with a 128K–256K vocabulary it is one of the largest single matmuls in the model.