R
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 T=(Ti)1≤i≤NT\mathcal{T} = (T_i)_{1\le i\le N_T} = the NTN_T 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 E∈RNT×dE \in \mathbb{R}^{N_T \times d} — token ii → row vector EiE_i (dim dd).
  • In LLMs: EE 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):

Modelvocabembedding dim dd
Llama 3 70B~128k8,192
Gemma 3 27B~262k5,376

2. Attention — the core formalism

The database analogy

A database = set of (key, value) tuples; query qq returns the value of the matching key. Attention relaxes "exact match" to "similarity-weighted average".

The math (general kernel view)

Setup: queries XQ∈RNQ×dinX_Q \in \mathbb{R}^{N_Q \times d_{in}}, keys XK∈RNKV×dinX_K \in \mathbb{R}^{N_{KV} \times d_{in}}, values XV∈RNKV×dinX_V \in \mathbb{R}^{N_{KV} \times d_{in}} (same dind_{in} — all come from the same embeddings). Chat example: NQN_Q = tokens of the latest question, NKVN_{KV} = the whole conversation.

Math (general attention):

Q=XQWQ,K=XKWK,V=XVWV,WQ,WK∈Rdin×dQK, WV∈Rdin×doutQ = X_Q W_Q, \quad K = X_K W_K, \quad V = X_V W_V, \qquad W_Q, W_K \in \mathbb{R}^{d_{in}\times d_{QK}}, \ W_V \in \mathbb{R}^{d_{in}\times d_{out}}

A=(κ(Q[i,:], K[j,:]))1≤i≤NQ1≤j≤NKV∈RNQ×NKV,Z[i,i]=∑1≤j≤NKVκ(Q[i,:],K[j,:])A = \big(\kappa(Q[i,:],\, K[j,:])\big)_{\substack{1\le i\le N_Q \\ 1\le j\le N_{KV}}} \in \mathbb{R}^{N_Q\times N_{KV}}, \qquad Z[i,i] = \sum_{1\le j\le N_{KV}} \kappa(Q[i,:], K[j,:])

Y=Z−1AV∈RNQ×dout,dout=dV\boxed{Y = Z^{-1} A V \in \mathbb{R}^{N_Q \times d_{out}}}, \qquad d_{out} = d_V

Takeaway: YY = rows of VV weighted by normalized query–key similarity. The kernel is almost always κ(v,w)=fκ(⟨v,w⟩)\kappa(v,w) = f_\kappa(\langle v,w\rangle) — dot product through a scalar function.

Cost: O(NQNKV(dQK+dV))O(N_Q N_{KV} (d_{QK} + d_V)) — quadratic in token count; dQK,dVd_{QK}, d_V are small by comparison.

Softmax = exp kernel + row normalization

Math (scaled exponential kernel) [Vas+17]:

κ(v,w)=exp⁡ ⁣(⟨v,w⟩dQK),A=exp⁡ ⁣(QKTdQK) (element-wise),(Z−1A)[i,:]=σ ⁣(Q[i,:]KTdQK)\kappa(v,w) = \exp\!\left(\frac{\langle v,w\rangle}{\sqrt{d_{QK}}}\right), \qquad A = \exp\!\left(\frac{QK^T}{\sqrt{d_{QK}}}\right) \ \text{(element-wise)}, \qquad (Z^{-1}A)[i,:] = \sigma\!\left(\frac{Q[i,:]K^T}{\sqrt{d_{QK}}}\right)

Why dQK\sqrt{d_{QK}}: entries of v,wv, w i.i.d. mean-0 var-1 ⇒ ⟨v,w⟩\langle v, w\rangle has variance dQKd_{QK} — 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 hh): Qh=XQWQhQ_h = X_Q W_Q^h, Kh=XKWKhK_h = X_K W_K^h, Vh=XVWVhV_h = X_V W_V^h, Ah=fκ(QhKhT)A_h = f_\kappa(Q_h K_h^T),

Yh=Zh−1AhVh∈RNQ×dhead,Y=concat⁡col(Yh)1≤h≤NheadsWO,WO∈RNheadsdhead×doutY_h = Z_h^{-1} A_h V_h \in \mathbb{R}^{N_Q\times d_{head}}, \qquad Y = \operatorname{concat}_{col}(Y_h)_{1\le h\le N_{heads}} W_O, \quad W_O \in \mathbb{R}^{N_{heads}d_{head}\times d_{out}}

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 XX; heads differ only by their weight matrices:

"linearly project the queries, keys and values hh 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 XWQX W_Q matmul — never a partition of XX.

Projection sizes (the free-parameter question)

  • dQKd_{QK} is a free parameter — no constraint vs dind_{in}. Paper: dQKd_{QK} and dVd_V "can be different in theory, in practice they are often equal to one another" (§2.1); in MHA, dQK=dheadd_{QK} = d_{head} (§2.2).
  • Total Q/K width = Nheads⋅dheadN_{heads} \cdot d_{head}. Table 3: Llama 3 70B → 64×128 = 8,192 = dind_{in}; Gemma 3 27B → 32×128 = 4,096 < 5,376. Equal or smaller in practice; never larger.
  • dQKd_{QK} enters the math in exactly one place: the dQK\sqrt{d_{QK}} kernel scaling.

3. Architecture

Encoder layer

(1) multi-head self-attention (XQ=XK=XVX_Q = X_K = X_V, all arrows from the previous layer) → (2) LayerNorm → (3) FFN → (4) skip connections around each sublayer (preserve gradient flow, fight vanishing gradients).

  • FFN: y=f(Wx+b)y = f(Wx + b); ReLU family; GLU variants [Sha20] dominate recent SOTA (Gemma, Llama 3, Qwen3).

Normalization — 4 variants

Math (LayerNorm) [BKH16]: x~=x−μσ2+ϵ\tilde{x} = \frac{x-\mu}{\sqrt{\sigma^2+\epsilon}}, then y=γ⊙x~+βy = \gamma \odot \tilde{x} + \beta — recenter + rescale + learned gain/shift.

Math (RMSNorm) [ZS19]: x~=x/RMS(x)\tilde{x} = x / \mathrm{RMS}(x) with RMS(x)=1d∑ixi2\mathrm{RMS}(x) = \sqrt{\frac{1}{d}\sum_i x_i^2}, 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 ii sees only j≤ij \le i: Si={j:j≤i}S_i = \{j : j \le i\},

A=(1j∈Si κ(Q[i,:],K[j,:]))i,j=κ(QKT+M),Mij={0j∈Si−∞elseA = \big(\mathbb{1}_{j\in S_i}\, \kappa(Q[i,:], K[j,:])\big)_{i,j} = \kappa(QK^T + M), \qquad M_{ij} = \begin{cases}0 & j \in S_i \\ -\infty & \text{else}\end{cases}

(equivalently A=κ(QKT)⊙1Si(j)A = \kappa(QK^T) \odot \mathbb{1}_{S_i}(j) — mask after the kernel; −∞-\infty → zero probability after softmax.)

Cross-attention: keys/values = encoder's final output ENLE_{N_L} (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-onlyDecoder-only
ArchetypeBERT [Dev+19]GPT [Rad+18]
Taskextraction / classificationgeneration (next-token)
Pretrain objectivemasked completionnext-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):

Y[NQ,j]=∑i=1NKVκ(Q[NQ,:],K[i,:])∑ℓ=1NKVκ(Q[NQ,:],K[ℓ,:]) V[i,j]Y[N_Q, j] = \sum_{i=1}^{N_{KV}} \frac{\kappa(Q[N_Q,:], K[i,:])}{\sum_{\ell=1}^{N_{KV}} \kappa(Q[N_Q,:], K[\ell,:])} \, V[i,j]

  • Naive: recompute all K/V per new token → O(Ntokens3d)O(N_{tokens}^3 d) total.
  • KV cache: store K/V; per new token only compute new K/V + the new query row → O(NQNtokensd)O(N_Q N_{tokens} d) per step, O(Ntokens2d)O(N_{tokens}^2 d) overall.
  • Cache memory: 2⋅NL⋅Nheads⋅NKV⋅d2 \cdot N_L \cdot N_{heads} \cdot N_{KV} \cdot d 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: G≤NheadsG \le N_{heads} KV groups, group size s=Nheads/Gs = N_{heads}/G; query head hh uses KV head g(h)=⌊(h−1)/s⌋+1g(h) = \lfloor (h-1)/s \rfloor + 1:

Qh=XQWQh  ∀h,Kg=XKWKg,Vg=XVWVg  (g=1…G),Yh=Zh−1AhVg(h),  Ah=fκ(QhKg(h)T)Q_h = X_Q W_Q^h \ \ \forall h, \qquad K_g = X_K W_K^g, \quad V_g = X_V W_V^g \ \ (g = 1 \dots G), \qquad Y_h = Z_h^{-1} A_h V_{g(h)}, \ \ A_h = f_\kappa(Q_h K_{g(h)}^T)

What it buys:

  • Weights: NheadsN_{heads} copies of WQW_Q, but only GG copies of WK,WVW_K, W_V.
  • Cache: 2Gdhead2 G d_{head} per token/layer vs 2Nheadsdhead2 N_{heads} d_{head} (Table 1: replace NheadsN_{heads} with #KV heads in the WKW_K, WVW_V and cache terms).
  • Limits: G=NheadsG = N_{heads} → MHA; G=1G = 1 → 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 (din=3d_{in}=3, dhead=2d_{head}=2, x=(1,2,3)x=(1,2,3)): WK1=[[1,0],[0,1],[1,1]]W_K^1 = [[1,0],[0,1],[1,1]] → K1=(4,5)K_1 = (4,5); WK2=[[0,1],[1,0],[1,0]]W_K^2 = [[0,1],[1,0],[1,0]] → K2=(5,1)K_2 = (5,1) — same xx, different weights. Share one WKa=WK1W_K^a = W_K^1 between both heads → K1=K2=(4,5)K_1 = K_2 = (4,5) exactly. The "head" = a block of output columns in WK=(WK1⋯WKNheads)W_K = (W_K^1 \cdots W_K^{N_{heads}}) — 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: Q=(Q1⋯QNheads)=XQWQQ = (Q_1 \cdots Q_{N_{heads}}) = X_Q W_Q with WQ=(WQ1⋯WQNheads)W_Q = (W_Q^1 \cdots W_Q^{N_{heads}}), likewise K,VK, V. Then low-rank factorize (latent dims dLQd_{LQ} for queries, dLd_L shared for K/V):

WQ=WLQWLQQ,WK=WLWLK,WV=WLWLV,WL∈Rdin×dLW_Q = W_{LQ} W_{LQQ}, \qquad W_K = W_L W_{LK}, \qquad W_V = W_L W_{LV}, \qquad W_L \in \mathbb{R}^{d_{in}\times d_L}

Cached object: the shared latent L=XWL∈RNKV×dLL = X W_L \in \mathbb{R}^{N_{KV}\times d_L} — NOT the K/Vs. (DeepSeek-V2 trains this directly; it is not factored out of an existing MHA.)

Weight merging (the actual trick):

  1. QK merge per head: WLQKh=WLQQh(WLKh)T∈RdLQ×dLW_{LQK}^h = W_{LQQ}^h (W_{LK}^h)^T \in \mathbb{R}^{d_{LQ}\times d_L} since QhKhT=LQWLQQh(WLKh)TLTQ_h K_h^T = L_Q W_{LQQ}^h (W_{LK}^h)^T L^T — attention logits computed straight on the latent.
  2. OV merge: WLO=blockdiag⁡(WLV1,…,WLVNheads) WO∈RNheadsdL×doutW_{LO} = \operatorname{blockdiag}(W_{LV}^1, \dots, W_{LV}^{N_{heads}})\, W_O \in \mathbb{R}^{N_{heads} d_L \times d_{out}} since Vh=L(WLVh)TV_h = L (W_{LV}^h)^T:

Y=concat⁡col(Zh−1AhL)1≤h≤NheadsWLO,Ah=exp⁡ ⁣(LQWLQKhLT)Y = \operatorname{concat}_{col}\left(Z_h^{-1} A_h L\right)_{1\le h\le N_{heads}} W_{LO}, \qquad A_h = \exp\!\left(L_Q W_{LQK}^h L^T\right)

Only weights kept: WLO,WL,WLQ,WLQKW_{LO}, W_L, W_{LQ}, W_{LQK} + latent vectors. KhK_h / VhV_h are never materialized.

Intuition (why the merging is valid):

  • Rank constraint: WK=WLWLK⇒rank⁡(WK)≤dLW_K = W_L W_{LK} \Rightarrow \operatorname{rank}(W_K) \le d_L — all heads' K/V subspaces live in ONE common dLd_L-dim subspace of Rdin\mathbb{R}^{d_{in}}. That shared-subspace restriction is the only approximation; everything else is exact algebra.
  • Queries are never cached: LQ=XQWLQL_Q = X_Q W_{LQ} has only NQN_Q rows (current step); only LL (the NKVN_{KV} 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
MHA2Nheadsdhead2 N_{heads} d_{head}
GQA2Gdhead2 G d_{head}
MLAdLd_L

DeepSeek-V2: 128 heads × 128 dims = 32,768 floats/token/layer → dL=512d_L = 512 → 64× smaller cache (60 layers, din=5120d_{in} = 5120, MoE).

⚠ RoPE caveat [Su+23]: without positional embeddings, MLA ≡ low-rank MHA exactly. With RoPE, a rotation Rm∈Rdhead×dheadR_m \in \mathbb{R}^{d_{head}\times d_{head}} is applied to K after it's built:

Qh[i,:]RiRjTKh[j,:]T=(LQ[i,:])WLQQhRiRjT(WLKh)T(L[j,:])TQ_h[i,:] R_i R_j^T K_h[j,:]^T = (L_Q[i,:]) W_{LQQ}^h R_i R_j^T (W_{LK}^h)^T (L[j,:])^T

RR sits between the two factors ⇒ WLQQhW_{LQQ}^h and WLKhW_{LK}^h 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)

RefWhy it matters here
[Vas+17] Attention is All You NeedOriginal Transformer + softmax attention (§2)
[BKH16] Layer NormalizationLN definition used in every layer
[ZS19] RMS Layer NormalizationCheaper norm (no recentering/bias) in modern LLMs
[Xio+20] On Layer NormalizationPre-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 VariantsFFN nonlinearity now standard in SOTA
[Dee+24] DeepSeek-V2Source of MLA; 60-layer MoE, dL=512d_L = 512
[Su+23] RoFormer / RoPERotary positional embedding; why MLA must break exactness
[MYZ25] TransMLAGQA/MHA → MLA conversion for other models
[Han+25] Streaming AttentionQuery ⊂ KV; motivates KV caching
[Zha+24] Dive into Deep LearningSource of the database analogy
[Dev+19] BERT / [Rad+18] GPTEncoder-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 QKTQK^T similarity or VV.
  • 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.