Transformer Architecture Q&A

Course Content

Transformer Architecture Q&A

6 sections · 60 lessons

Trace the full GPT model from input tokens to output logits, naming each step?


GPT-2 forward pass with shapesToken ids (B, T)Token + position embeddings (B, T, d)N blocks: x + Attn(LN x), x + MLP(LN x)Final layer norm (B, T, d)LM head, tied to wte: logits (B, T, V)
The shape stays (B, T, d) through every block; only the lookup at the bottom and the head at the top change it.

What you need to know

The pipeline with shapes (GPT-2 style)

Text
1. Tokenize            text -> ids                      (B, T)2. Token embedding     wte[ids]                         (B, T, d)3. Position            + wpe[0:T]                       (B, T, d)   RoPE models skip this4. Embedding dropout   (0.1 in GPT-2, 0 in modern LLMs)5. N blocks, each:     x = x + Attn(LN1(x))    causal attention, then W_o     x = x + MLP(LN2(x))     d -> 4d -> d with GELU6. Final layer norm    ln_f(x)                          (B, T, d)7. LM head             x @ wte.T                        (B, T, V)   logits

The same trace, runnable

Python
import torch, torch.nn as nnV, T_max, d, n_layers, n_heads = 1000, 64, 128, 2, 4   # a toy GPTclass Block(nn.Module):    def __init__(self):        super().__init__()        self.ln1, self.ln2 = nn.LayerNorm(d), nn.LayerNorm(d)        self.attn = nn.MultiheadAttention(d, n_heads, batch_first=True)        self.mlp = nn.Sequential(nn.Linear(d, 4*d), nn.GELU(), nn.Linear(4*d, d))    def forward(self, x):        T = x.size(1)        mask = torch.triu(torch.ones(T, T, dtype=torch.bool), 1)   # True = blocked        h = self.ln1(x)        x = x + self.attn(h, h, h, attn_mask=mask, need_weights=False)[0]        return x + self.mlp(self.ln2(x))wte, wpe = nn.Embedding(V, d), nn.Embedding(T_max, d)blocks, ln_f = nn.ModuleList(Block() for _ in range(n_layers)), nn.LayerNorm(d)ids = torch.randint(0, V, (2, 10))                  # (B, T) token idsx = wte(ids) + wpe(torch.arange(10))                # (B, T, d)for blk in blocks:    x = blk(x)                                      # (B, T, d) after every blocklogits = ln_f(x) @ wte.weight.T                     # (B, T, V), tied headprint(ids.shape, x.shape, logits.shape)# torch.Size([2, 10]) torch.Size([2, 10, 128]) torch.Size([2, 10, 1000])

The shape stays (B, T, d) through every block. Only the embedding lookup (ids to vectors) and the LM head (vectors to logits) change it.

What each part does

  • Attention moves information between positions — the only step that does.
  • MLP transforms each position on its own; most of the parameters and much of the stored knowledge live here.
  • Residual stream — each sublayer adds to x rather than replacing it, so information and gradients flow straight through a deep stack.
  • Pre-norm — the layer norm sits inside each branch, before the sublayer. This keeps the residual path clean and makes deep stacks train more stably than the original post-norm design.
  • Logits — one unnormalised score per vocabulary entry, at every position.

Training versus generation

  • Training uses every position: position t's logits are compared with token t+1 by cross-entropy. One forward pass over 1,024 tokens gives 1,024 predictions.
  • Generation uses only the last position's logits, samples a token, appends it and repeats, with the KV cache avoiding recomputation.

In a Llama-style model the same trace holds with four substitutions: no wpe (RoPE rotates Q and K inside attention), RMSNorm instead of LayerNorm, SwiGLU instead of GELU, and GQA attention.

A real-life example

A code-completion assistant receives the text before the cursor: 3,000 tokens. With a Llama-3-8B-sized model (d = 4096, 32 layers, vocabulary 128,256):

Text
ids                  (1, 3000)after embedding      (1, 3000, 4096)after each block     (1, 3000, 4096)   × 32 blockslast position only   (1, 1, 4096)logits               (1, 1, 128256)

The server slices to the last position before the LM head, because only the next token matters. Computing full logits for all 3,000 positions would create 385 million numbers that are thrown away. The whole pass is the prefill; each later token runs the same trace with T = 1 and the KV cache.

Follow-up questions to expect

  • "Where does position information enter in Llama?" — Not at the input: RoPE rotates each query and key inside every attention layer.
  • "Why a final layer norm?" — In pre-norm, the residual stream itself is never normalised inside the blocks, so its scale grows; ln_f normalises it before the LM head.
  • "Which step is the only one mixing tokens?" — Attention. Everything else works per position.