Transformers & Attention
12 questions
Walk through the full Transformer architecture block-by-block.▶
Input: token embeddings + positional encoding → [batch, seq_len, d_model]
Each Transformer block (decoder):
- Masked Multi-Head Self-Attention + residual + LayerNorm
- Feed-Forward Network (Linear → GELU → Linear, 4× width expansion) + residual + LayerNorm
Encoder additionally has cross-attention (decoder attends to encoder output).
Output: final hidden states → language model head (linear → softmax over vocab)
Pre-norm vs Post-norm: modern LLMs (LLaMA, GPT-3) use pre-norm (LayerNorm before attention). Original paper used post-norm. Pre-norm trains more stably.
Parameter count: per block: 4×d² (attention) + 8×d² (FFN) = 12×d². GPT-3 (d=12288, 96 layers): ~175B params.
Encoder vs decoder Transformers - differences and use cases. Why are SOTA models decoder-only?frontier▶
Encoder (BERT-style):
- Bidirectional attention - each token attends to ALL other tokens
- Trained with masked language modeling (predict masked tokens)
- Good for: classification, NER, embeddings, understanding tasks
Decoder (GPT-style):
- Causal attention - each token only attends to past tokens
- Trained with next-token prediction (causal language modeling)
- Good for: text generation, in-context learning, any task via prompting
Why decoder-only dominates SOTA:
- Next-token prediction scales - any text on the internet is free training signal
- In-context learning emerges naturally (few-shot examples in context)
- One model handles generation AND understanding via prompting
- Chinchilla scaling laws favor decoder-only architectures
- Simpler architecture (no cross-attention needed)
What is Flash Attention? How does it reduce memory from O(n²) to O(n)?frontier▶
Standard attention writes the full N×N score matrix to HBM - bandwidth-bound for typical context lengths.
FlashAttention: tile Q,K,V into on-chip SRAM; online softmax rescales partial blocks so the full N×N matrix never materializes in HBM. Output is exact, not approximate.
- Memory: O(N) in HBM vs O(N²)
- Speed: often 2–4× wall-clock (FA2 improves warp scheduling further)
Same IO-aware idea as tiled matmul - minimize HBM round-trips.
What is KV cache? Memory scaling, speedup, and how to reduce it.frontier▶
During autoregressive decode, each step needs K and V for all prior tokens. Without a cache you recompute them every step - O(t²) total work.
With cache: store K,V per layer; each new step only projects the latest token and appends. Per-step compute is O(1) in sequence length (attention still reads the growing cache).
Memory: 2 × layers × kv_heads × head_dim × seq_len × bytes × batch_size.
LLaMA-3-8B (32 layers, GQA with 8 kv heads, d_head=128, FP16): ~131 KB/token → ~512 MB at 4096 tokens. At 128K context the cache alone can exceed weight memory.
Reduce it: GQA/MQA (fewer kv heads), shorter context, FP8/quantized KV, paged allocation (vLLM). See also GQA and Paged Attention questions.
Explain positional encoding - sinusoidal, learned, and RoPE.frontier▶
Transformers have no inherent sense of position (attention is permutation-invariant). Positional encoding adds position information to token embeddings.
Sinusoidal (original): PE(pos, 2i) = sin(pos/100002i/d), PE(pos, 2i+1) = cos(...). Deterministic, generalizes to unseen lengths. But additive - mixed with token embedding.
Learned: each position has a learnable embedding vector. Simple, works well up to training context length. Doesn't generalize beyond training length.
RoPE (Rotary Position Embedding): rotate Q and K vectors based on position before computing attention.
- Rotate 2D slices of Q, K by angle θ = m·θ_base^{-2i/d}
- Key property: (R_m q) · (R_n k) = q^T R_{n-m} k - attention score only depends on relative position (n-m), not absolute positions
- Generalizes to longer sequences via extrapolation/interpolation (YaRN, NTK-aware scaling)
- Used in: LLaMA, Qwen, Mistral, Gemma
What is Grouped Query Attention (GQA)? Why does LLaMA use it?frontier▶
Multi-Head Attention (MHA): each of n_heads has its own K, V. KV cache = n_heads × d_k × seq_len.
Multi-Query Attention (MQA): all heads share a single K, V. KV cache = 1 × d_k × seq_len. Fastest but quality drops.
Grouped Query Attention (GQA): n_kv_heads groups, each shared by n_heads/n_kv_heads query heads.
LLaMA 3 8B: n_heads=32, n_kv_heads=8 → 4x KV cache reduction vs MHA.
Why it matters: KV cache is the memory bottleneck for large batches at long contexts. GQA reduces KV cache × n_kv_heads/n_heads without significantly hurting quality. This allows larger batches (higher throughput).
Training trick: you can convert an MHA model to GQA by averaging (or selecting) K, V heads within each group - "uptrained" MQA.
What is Mixture of Experts (MoE)? Dense vs sparse models.frontier▶
MoE replaces the FFN in each Transformer block with multiple "expert" FFNs. A router network selects top-k experts for each token.
Dense model: every parameter used for every token. 70B Mixtral dense = 70B active params per token.
Sparse MoE: 8 experts, top-2 routing. Each token uses 2/8 experts. Mixtral 8×7B: 46B total params, ~13B active per token. Speed of a 13B model, quality of a 46B model.
Benefits: scale parameters without proportionally scaling compute. More capacity, same FLOP budget.
Challenges:
- Load balancing: without auxiliary loss, all tokens route to same expert (collapse)
- Memory: all experts must fit in GPU memory even if not all active
- Communication overhead: in distributed settings, all-to-all communication for expert routing
DeepSeek-V2 uses fine-grained MoE (many small experts) + shared experts.
Explain scaling laws (Chinchilla). How do you decide model size vs data budget?frontier▶
Kaplan et al. (2020): loss scales as power laws with model size N and data D:
L(N,D) ∝ N-α + D-β
Implied: for a given compute budget C ∝ N×D, optimal to scale N more than D. Led to "too large, too undertrained" models (GPT-3).
Chinchilla (Hoffmann 2022): corrected analysis. Optimal: tokens D ≈ 20× parameters N.
- GPT-3 (175B params, 300B tokens): undertrained. Should have used ~3.5T tokens.
- Chinchilla (70B params, 1.4T tokens): outperforms GPT-3 despite being smaller.
- LLaMA 3 (8B params, 15T tokens): trained far beyond Chinchilla-optimal for inference efficiency.
Practical implication: for production (inference cost matters), train smaller models on more data. For research (care about training compute), follow Chinchilla ratio.
Implement scaled dot-product attention in PyTorch.frontier▶
Core of every Transformer. Q, K, V are projections of input. Output is weighted sum of values.
import torch
import torch.nn.functional as F
def scaled_dot_product_attention(Q, K, V, mask=None):
"""
Q, K, V: (batch, heads, seq_len, d_k)
mask: (batch, 1, seq_len, seq_len) - 0 where we want -inf
"""
d_k = Q.shape[-1]
# (batch, heads, seq_len, seq_len)
scores = (Q @ K.transpose(-2, -1)) / (d_k ** 0.5)
if mask is not None:
scores = scores.masked_fill(mask == 0, float('-inf'))
weights = F.softmax(scores, dim=-1)
# (batch, heads, seq_len, d_k)
return weights @ VImplement multi-head attention from scratch.frontier▶
Projects Q/K/V to multiple heads, runs attention per head, concatenates and projects back.
import torch
import torch.nn as nn
import torch.nn.functional as F
class MultiHeadAttention(nn.Module):
def __init__(self, d_model, n_heads):
super().__init__()
assert d_model % n_heads == 0
self.n_heads = n_heads
self.d_k = d_model // n_heads
self.W_Q = nn.Linear(d_model, d_model, bias=False)
self.W_K = nn.Linear(d_model, d_model, bias=False)
self.W_V = nn.Linear(d_model, d_model, bias=False)
self.W_O = nn.Linear(d_model, d_model, bias=False)
def forward(self, x, mask=None):
B, L, D = x.shape
def proj(linear, x):
return linear(x).view(B, L, self.n_heads, self.d_k).transpose(1, 2)
Q, K, V = proj(self.W_Q, x), proj(self.W_K, x), proj(self.W_V, x)
# Q,K,V: (B, n_heads, L, d_k)
scores = (Q @ K.transpose(-2,-1)) / (self.d_k ** 0.5)
if mask is not None:
scores = scores.masked_fill(mask == 0, float('-inf'))
attn = F.softmax(scores, dim=-1)
out = (attn @ V).transpose(1,2).contiguous().view(B, L, D)
return self.W_O(out)Implement causal mask for decoder-only Transformer.▶
A causal mask ensures each position only attends to earlier positions (and itself). Upper triangle is masked to -inf.
import torch
def causal_mask(seq_len, device='cpu'):
# 1 where attention is allowed, 0 where masked
mask = torch.tril(torch.ones(seq_len, seq_len, device=device))
# shape (1, 1, seq_len, seq_len) for broadcasting over batch and heads
return mask.unsqueeze(0).unsqueeze(0)
# Usage:
# mask = causal_mask(seq_len, device=x.device)
# scores = scores.masked_fill(mask == 0, float('-inf'))Why is tokenization the root of most LLM weirdness?▶
BPE recap: start from bytes/characters, greedily merge the most frequent adjacent pair, repeat to a target vocab size. Frequent words collapse to one token; rare strings fragment into many.
The quirks this causes:
- Non-English tax: under-represented scripts (Hindi, Tamil, even code) are barely merged, so they cost many more tokens - higher price and a smaller effective context window for the same text.
- Bad arithmetic: numbers tokenize inconsistently ("127" might be one token, "128" two). Digits are not atomic, so the model never sees clean place value.
- Can't spell: "strawberry" is a couple of tokens, never characters - so counting its r's is genuinely hard for the model.
- Under-trained tokens: rare merges (the "SolidGoldMagikarp" tokens) have barely-trained embeddings and trigger weird behavior.
- Whitespace/case sensitivity: " the", "the", "The", "THE" are different tokens.
Why BPE at all: char-level → sequences too long for O(n²) attention; word-level → out-of-vocabulary explosions and a huge vocab. BPE is the compromise. Byte-level BPE (GPT-2 onward) operates on UTF-8 bytes so every possible string is representable - no OOV ever.
Note: "tokenization is at the heart of much of the weirdness of LLMs." Being able to trace a specific failure (spelling, math, multilingual cost) back to the tokenizer is a strong signal.