Transformer Architecture Q&A

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

Text
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 token

The 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 TEntries T²Memory per head per layer
4,09616.8 million32 MB
32,7681.07 billion2 GB
128,00016.4 billion30.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 × T matrix to GPU memory. Memory becomes linear in T; 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-T memory 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_model to save memory?" — The T × T matrix does not depend on d_model; only shortening or sparsifying the sequence dimension helps.