Course Content
Transformer Architecture Q&A
6 sections · 60 lessons
Why does attention weight matrix have shape (seq_len × seq_len), not embedding size?
What you need to know
Follow the shapes
Q (T, d_k)K (T, d_k)Q Kᵀ (T, d_k) × (d_k, T) = (T, T) d_k is summed over and disappearssoftmax (T, T) each row sums to 1A V (T, T) × (T, d_v) = (T, d_v) back to one vector per tokenThe rule of matrix multiplication decides it: the inner dimension (d_k) must match and vanishes; the outer dimensions (T and T) remain. With batches and heads the full tensor is (B, h, T, T).
For cross-attention the matrix is (T_query, T_key) — for example (20, 1500) when 20 text tokens attend to 1,500 audio frames. It is only square in self-attention.
Why this shape is a cost problem
Memory for one head's matrix in fp16 (2 bytes per number):
Sequence length T | Entries T² | Memory per head per layer |
|---|---|---|
| 4,096 | 16.8 million | 32 MB |
| 32,768 | 1.07 billion | 2 GB |
| 128,000 | 16.4 billion | 30.5 GB |
Multiply by heads and layers and a naive implementation cannot run long contexts at all. Doubling T quadruples this memory and the score compute.
How it is handled in 2026
- FlashAttention computes attention in tiles held in fast on-chip memory, using a running softmax, and never writes the full
T × Tmatrix to GPU memory. Memory becomes linear inT; compute is still quadratic, but much faster because memory traffic drops. - Sliding-window or local attention limits each token to nearby keys (for example the last 4,096), cutting the matrix to a band.
- Sparse and linear-attention variants, and hybrids that mix attention with state-space layers, attack the
T²compute itself. - GQA and MQA do not change this matrix; they shrink the KV cache, which is a separate, linear-in-
Tmemory cost.
A real-life example
A legal-tech company upgrades its document classifier from a 512-token BERT to a long-context encoder (ModernBERT supports 8,192 tokens) so a whole contract fits in one pass. Their first attempt uses an attention implementation that builds the full matrix. With 12 heads and a batch of 16, the fp16 scores at 8,192 tokens need about 26 GB for a single layer, before counting the other layers kept for backpropagation, and the job runs out of memory.
Switching to the FlashAttention path (attn_implementation="flash_attention_2" in Hugging Face, or PyTorch's scaled_dot_product_attention) removes the stored matrix; the same batch now fits. The model's maths — and its predictions — did not change, only how the T × T step was computed.
Follow-up questions to expect
- "Is the weight matrix a parameter?" — No. It is computed fresh for every input; the parameters are
W_q,W_k,W_v,W_o. - "Does FlashAttention make attention linear?" — In memory, yes; in compute, no — it still does
O(T²)work. - "Why not reduce
d_modelto save memory?" — TheT × Tmatrix does not depend ond_model; only shortening or sparsifying the sequence dimension helps.