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
1import torch, torch.nn as nn2torch.manual_seed(0)3d = 644wq, wk, wv = (nn.Linear(d, d, bias=False) for _ in range(3))5fused = nn.Linear(d, 3 * d, bias=False)6with torch.no_grad(): # stack the three weights row-wise7 fused.weight.copy_(torch.cat([wq.weight, wk.weight, wv.weight], dim=0))89x = torch.randn(2, 10, d)10q, k, v = fused(x).chunk(3, dim=-1) # one matmul, then slice11print(all(torch.allclose(a, b(x), atol=1e-6) for a, b in [(q, wq), (k, wk), (v, wv)]))12# TrueMultiplying 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 whenB·Tis 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:
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_owith them too?" —W_oneeds 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.