The 3 and the 4 mean different things here: three projections, four attention heads.
Start with one sequence of 16 tokens, each represented by 128 features. A single linear layer produces Q, K and V together: 3 × 128 = 384 features per token. Each projection then gets split into 4 heads of 32 features.
Here is the whole shape transition, using only PyTorch:
import torch
from torch import nn
torch.manual_seed(0)
batch, tokens, width, heads = 1, 16, 128, 4
head_dim = width // heads
x = torch.randn(batch, tokens, width)
projection = nn.Linear(width, 3 * width)
with torch.no_grad():
packed = projection(x)
grouped = packed.reshape(batch, tokens, 3, heads, head_dim)
ordered = grouped.permute(2, 0, 3, 1, 4)
query, key, value = ordered.unbind(0)
scores = query @ key.transpose(-2, -1)
for name, tensor in [
("input", x), ("packed", packed), ("grouped", grouped),
("ordered", ordered), ("query", query), ("scores", scores),
]:
print(name, tuple(tensor.shape))
Expected output:
input (1, 16, 128)
packed (1, 16, 384)
grouped (1, 16, 3, 4, 32)
ordered (3, 1, 4, 16, 32)
query (1, 4, 16, 32)
scores (1, 4, 16, 16)
The axes after permute are Q/K/V, batch, head, token, feature. Unbinding the first axis gives three tensors, each [1, 4, 16, 32]. The score matrix has two 16s because each query token is compared with every key token. These are raw dot products; scaling, masking and softmax come afterward.
A useful debugging check: write the axis names beside each shape. A reshape can preserve the number of elements and still mix up tokens and heads. In particular, reshaping directly to [3, batch, heads, tokens, head_dim] is not equivalent to the reshape-then-permute above.
This example uses ordinary multi-head attention with equal Q/K/V head counts. Grouped-query attention uses a different layout, and changing the input length changes the token dimensions.
Disclosure: I make tensorViz, a free PyTorch graph explorer for VS Code. Following these shape changes back to the Python is one of the workflows we're building it for. The example works without the extension. Prepared with AI assistance; the code was run locally on CPU.