Transformer Architecture Q&A

Course Content

Transformer Architecture Q&A

6 sections · 60 lessons

Explain combined QKV projection and why it’s faster than three separate projections?


What you need to know

The equivalence

Python
import torch, torch.nn as nntorch.manual_seed(0)d = 64wq, wk, wv = (nn.Linear(d, d, bias=False) for _ in range(3))fused = nn.Linear(d, 3 * d, bias=False)with torch.no_grad():                       # stack the three weights row-wise    fused.weight.copy_(torch.cat([wq.weight, wk.weight, wv.weight], dim=0))x = torch.randn(2, 10, d)q, k, v = fused(x).chunk(3, dim=-1)         # one matmul, then sliceprint(all(torch.allclose(a, b(x), atol=1e-6) for a, b in [(q, wq), (k, wk), (v, wv)]))# True

Multiplying by stacked matrices is the same as stacking the three results. Only the execution changes.

Why it is faster

  • Fewer kernel launches. Every GPU kernel has a fixed launch cost of a few microseconds. In decode at batch size 1, a layer's matmuls are tiny, so this overhead is a real share of the step. Fusing saves two launches per layer; across 32 layers, 64 launches per token.
  • One read of x. Decode is memory-bound. Reading the input activations once instead of three times cuts that traffic.
  • A better-shaped matmul. (B·T, d) × (d, 3d) is one large GEMM that tiles well on tensor cores; three (d, d) GEMMs each leave more of the hardware idle, especially when B·T is small.

The weight bytes read are the same either way (all three matrices must be read), so the gain is in launches, activation reads and efficiency — useful, but not a 3× speedup.

With GQA

K and V are narrower than Q. For Llama-3-8B:

Text
Q: 32 heads × 128 = 4,096K:  8 heads × 128 = 1,024V:  8 heads × 128 = 1,024fused output width = 4,096 + 2 × 1,024 = 6,144   (not 3 × 4,096 = 12,288)

So you split with torch.split(out, [4096, 1024, 1024], dim=-1), not chunk(3).

Where you see it

GPT-2's checkpoint already stores a fused c_attn of shape 768 × 2304. Llama's Hugging Face checkpoint stores q_proj, k_proj and v_proj separately, and serving engines such as vLLM stack them into one fused layer when loading. The same trick is used for the gate and up projections of a SwiGLU FFN.

A real-life example

A team serving a chatbot to a million users profiles one decode step of their 8B model at small batch sizes during quiet hours. The trace shows hundreds of small kernels per token, with the GPU idle between many of them. Loading the checkpoint into a fused QKV layout and a fused gate-up layout cuts the number of matmul launches per layer from 7 to 4.

Combined with CUDA graphs, which record the whole step and replay it with almost no launch overhead, their measured time per token at batch size 1 drops noticeably. At large batch sizes during peak hours the gain is smaller, because the matmuls are big enough that launch overhead matters less — exactly what the theory predicts.

Follow-up questions to expect

  • "Is the fused model a different model?" — No, the outputs are identical; only the weight layout differs, so conversion is a reshape of the checkpoint.
  • "Why not fuse W_o with them too?" — W_o needs the attention output, which depends on Q, K and V, so it cannot run in the same matmul.
  • "How does tensor parallelism handle a fused QKV?" — Each GPU gets its share of Q heads and K/V heads, and the fused weight must be split by head, not by simple thirds.