Skip to content
samit
Interviews/Transformers & Attention

Transformers & Attention

12 questions

12 questions
Architecture
Walk through the full Transformer architecture block-by-block.▶

Input: token embeddings + positional encoding → [batch, seq_len, d_model]

Each Transformer block (decoder):

  1. Masked Multi-Head Self-Attention + residual + LayerNorm
  2. 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.

Implementation
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.

python
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 @ V
Implement multi-head attention from scratch.frontier▶

Projects Q/K/V to multiple heads, runs attention per head, concatenates and projects back.

python
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.

python
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.