- Published on
Transformers & Attention (Serret 2026) — Paper Study Notes
Understanding Transformers and Attention — Paper Study Notes
Paper: Understanding Transformers and Attention Mechanisms: An Introduction for Applied Mathematicians — Michel Fabrice Serret (Paul Scherrer Institute). arXiv:2604.00965v1 [math.NA], Apr 2026. 13pp, 9 figures. Written for the IPAM RNLA workshop (project: "Randomization in Transformer models").
The whole paper in one line: attention = database lookup with a similarity kernel + row normalization. Everything else is engineering around that core.
Why this paper
- No new results — value is a clean, self-consistent formalism for attention, aimed at applied mathematicians.
- Walks the full modern stack: tokenization → embeddings → attention → MHA → encoder/decoder → KV cache → GQA → MLA.
- Ends on the memory/compute trade-offs that drive modern attention compression — the part that matters for LLM inference.
1. Tokens → Embeddings
Tokenization
- Text → sequence of substrings (tokens); vocabulary = the distinct tokens.
- A text becomes a sequence of token indices.
- Tokenizer tension: fine enough to capture semantics, small enough to not dilute info or blow up memory.
- In practice: tokens ≈ words (paper takes this WLOG).
Embeddings (vectorization)
- One trained linear map: matrix — token → row vector (dim ).
- In LLMs: is trained from scratch as model weights.
- Optional feature embeddings (positional, sentence-id) on top; DeBERTa & DeepSeek V2 keep them separate by design.
Scale check (Table 3 / §1):
| Model | vocab | embedding dim |
|---|---|---|
| Llama 3 70B | ~128k | 8,192 |
| Gemma 3 27B | ~262k | 5,376 |
2. Attention — the core formalism
The database analogy
A database = set of (key, value) tuples; query returns the value of the matching key. Attention relaxes "exact match" to "similarity-weighted average".
The math (general kernel view)
Setup: queries , keys , values (same — all come from the same embeddings). Chat example: = tokens of the latest question, = the whole conversation.
Math (general attention):
Takeaway: = rows of weighted by normalized query–key similarity. The kernel is almost always — dot product through a scalar function.
Cost: — quadratic in token count; are small by comparison.
Softmax = exp kernel + row normalization
Math (scaled exponential kernel) [Vas+17]:
Why : entries of i.i.d. mean-0 var-1 ⇒ has variance — dividing rescales the dot product to unit variance.
⚠ Softmax attention is a special case of the kernel view (exp kernel + row-normalize). The kernel framing is the general object.
Multi-headed attention (MHA)
Math (per head ): , , , ,
Why: parallel attention over different projected subspaces → attends to several types of semantic info simultaneously.
⚠ Heads do NOT split the input. Every head consumes the same full ; heads differ only by their weight matrices:
"linearly project the queries, keys and values times with different, learned linear projections" — Vaswani et al. 2017, §3.2.2
The "split into heads" in diagrams/code = reshape of the output columns of one big matmul — never a partition of .
Projection sizes (the free-parameter question)
- is a free parameter — no constraint vs . Paper: and "can be different in theory, in practice they are often equal to one another" (§2.1); in MHA, (§2.2).
- Total Q/K width = . Table 3: Llama 3 70B → 64×128 = 8,192 = ; Gemma 3 27B → 32×128 = 4,096 < 5,376. Equal or smaller in practice; never larger.
- enters the math in exactly one place: the kernel scaling.
3. Architecture
Encoder layer
(1) multi-head self-attention (, all arrows from the previous layer) → (2) LayerNorm → (3) FFN → (4) skip connections around each sublayer (preserve gradient flow, fight vanishing gradients).
- FFN: ; ReLU family; GLU variants [Sha20] dominate recent SOTA (Gemma, Llama 3, Qwen3).
Normalization — 4 variants
Math (LayerNorm) [BKH16]: , then — recenter + rescale + learned gain/shift.
Math (RMSNorm) [ZS19]: with , then a learned gain — no recentering, no bias ⇒ cheaper. Standard in modern LLMs (Llama, Gemma, DeepSeek).
Pre-LN vs Post-LN [Xio+20]: Post-LN = original Transformer (norm after each sublayer). Pre-LN = norm before → skip connections stay unnormalized, gradients skip through unscathed ⇒ easier training, slight quality cost.
Peri-LN [Kim+25]: norm around each sublayer (before AND after the module) + input/output embedding norms.
- Paper's remark: "state-of-the-art large language models such as Gemma 3 have made use of both Pre-Layer and Post-Layer Normalization" [Kim+25].
- Peri-LN paper (arXiv 2502.02732, ICML 2025): more balanced activation-variance growth, steadier gradient flow (tested ≤3.2B params). Adopters named: Gemma 2/3 and OLMo 2.
- ⚠ Not Gemini, not Kimi. "[Kim+25]" = Jeonghoon Kim et al. — reads like "Kimi", but neither paper mentions Kimi.
Decoder layer
Masked causal self-attention → cross-attention → FFN.
Math (causal mask): query sees only : ,
(equivalently — mask after the kernel; → zero probability after softmax.)
Cross-attention: keys/values = encoder's final output (source sentence), queries = decoder — how the decoder relates its output to the input. ⚠ Cross-attention / shared embedding spaces = root of multimodal LLMs (Flamingo [Ala+22], Latent Diffusion [Rom+22]).
Encoder-only vs Decoder-only
| Encoder-only | Decoder-only | |
|---|---|---|
| Archetype | BERT [Dev+19] | GPT [Rad+18] |
| Task | extraction / classification | generation (next-token) |
| Pretrain objective | masked completion | next-token prediction |
Paradigm: pretrain on huge data → fine-tune on small data [Rad+18].
4. Efficiency: KV cache → GQA → MLA
Why KV caching
The last token's output needs all K/V:
Math (last-token output):
- Naive: recompute all K/V per new token → total.
- KV cache: store K/V; per new token only compute new K/V + the new query row → per step, overall.
- Cache memory: floats → prohibitive in long context.
- Streaming attention [Han+25]: queries ⊂ KV tokens (chat) → only new queries computed, if K/V is cached.
GQA — Grouped Query Attention
⚠ It's the KEY/VALUE heads that get shared — queries are NEVER shared. Q is computed fresh each step and discarded; only K/V get cached — so sharing K/V is what saves memory.
Math: KV groups, group size ; query head uses KV head :
What it buys:
- Weights: copies of , but only copies of .
- Cache: per token/layer vs (Table 1: replace with #KV heads in the , and cache terms).
- Limits: → MHA; → MQA (single shared KV head).
Table 3 examples: Llama 3 70B → 64 query heads / 8 KV heads (8× smaller cache); Gemma 3 27B → 32 / 16 (2×).
⚠ Never-forget version (why shared K/V ⇒ identical outputs):
All heads see the SAME input vector. MHA = same input × different weight matrices. GQA = same input × the SAME weight matrix (within a group) ⇒ the multiplication is literally identical ⇒ identical K/V outputs ⇒ store ONE copy per group. The vectors aren't "similar" — they're the same number, computed twice, stored once.
Numeric check (, , ): → ; → — same , different weights. Share one between both heads → exactly. The "head" = a block of output columns in — the input is never partitioned.
Why it works: queries encode what to look for (must stay diverse per head); K/V encode content of the context (shareable memory). GQA = quality/cache dial. Origin papers (not in Serret's refs): MQA [Sha19, arXiv 1911.02150], GQA [Ain+23, arXiv 2305.13245 — ~8 groups ≈ MHA quality].
MLA — Multi-head Latent Attention (DeepSeek-V2) [Dee+24]
Main idea: cache one low-rank latent vector per token, shared across ALL heads — instead of per-head K/V.
Column-concat view first: with , likewise . Then low-rank factorize (latent dims for queries, shared for K/V):
Cached object: the shared latent — NOT the K/Vs. (DeepSeek-V2 trains this directly; it is not factored out of an existing MHA.)
Weight merging (the actual trick):
- QK merge per head: since — attention logits computed straight on the latent.
- OV merge: since :
Only weights kept: + latent vectors. / are never materialized.
Intuition (why the merging is valid):
- Rank constraint: — all heads' K/V subspaces live in ONE common -dim subspace of . That shared-subspace restriction is the only approximation; everything else is exact algebra.
- Queries are never cached: has only rows (current step); only (the rows) persists across steps.
- Analogy: GQA shares one dictionary across heads; MLA compresses the whole dictionary into a bottleneck that every head reads a linear view of.
Memory comparison (Tables 1–2, floats):
| per-token cache per layer | |
|---|---|
| MHA | |
| GQA | |
| MLA |
DeepSeek-V2: 128 heads × 128 dims = 32,768 floats/token/layer → → 64× smaller cache (60 layers, , MoE).
⚠ RoPE caveat [Su+23]: without positional embeddings, MLA ≡ low-rank MHA exactly. With RoPE, a rotation is applied to K after it's built:
sits between the two factors ⇒ and can't be merged (RoPE would need recomputation per evaluation). Fix: keep the latent rope-free and append a small non-latent K/Q part that carries the position — breaks exactness, keeps the speedup. [MYZ25] (TransMLA) converts pretrained GQA/MHA → MLA, porting DeepSeek's optimization to other models.
Key references (what to look up when a line feels unsupported)
| Ref | Why it matters here |
|---|---|
| [Vas+17] Attention is All You Need | Original Transformer + softmax attention (§2) |
| [BKH16] Layer Normalization | LN definition used in every layer |
| [ZS19] RMS Layer Normalization | Cheaper norm (no recentering/bias) in modern LLMs |
| [Xio+20] On Layer Normalization | Pre-LN variant; unnormalized skip connections |
| [Kim+25] Peri-LN (arXiv 2502.02732, ICML'25) | LN around sublayers (pre + post); Gemma 2/3 + OLMo 2 adopt it |
| [Sha20] GLU Variants | FFN nonlinearity now standard in SOTA |
| [Dee+24] DeepSeek-V2 | Source of MLA; 60-layer MoE, |
| [Su+23] RoFormer / RoPE | Rotary positional embedding; why MLA must break exactness |
| [MYZ25] TransMLA | GQA/MHA → MLA conversion for other models |
| [Han+25] Streaming Attention | Query ⊂ KV; motivates KV caching |
| [Zha+24] Dive into Deep Learning | Source of the database analogy |
| [Dev+19] BERT / [Rad+18] GPT | Encoder-only vs decoder-only + pretrain/finetune paradigm |
Connections & open threads
- RNLA angle (why this paper exists): the kernel + row-normalization framing is exactly where randomized NLA enters attention — sketch / row-sample the similarity or .
- Cross-attention ↔ latent diffusion [Rom+22]: same mechanism the video_diffusion notes build on.
- MLA weight-merge algebra = a clean worked example of low-rank weights (cf. sparsity_tricks).
- Self-consistency checks: softmax = kernel special case; GQA = MHA with tied KV heads; MLA without RoPE = low-rank MHA. All three read off the same formalism.