Skip to content

Multi-head attention

In the previous lesson we built the context vector by projecting the input into three new matrices, Q, K, and V, each with its own trainable weight matrix that gets optimized during training. We also landed on a key insight: those projection matrices do not have to match the dimensions of our input matrix.

This is where that insight pays off.

If a single self-attention produces a context vector of dimension M, then with multi-head attention we create N of them and concatenate them into a single context matrix.

N attention heads each output a context of width M, concatenated along the feature axis into a single context matrix of shape seqLen by dOut

Reading the diagram:

  • seqLen is the number of tokens in the sequence, one row per token. It never changes.
  • M is the width of a single head’s context vector, the features that one head produces per token.
  • dOut is the width of the final context matrix after all heads are concatenated, so it is M per head stacked across every head.

Why does this matter? Because each head attends to different dimensions with a different emphasis during training. The final concatenated result is richer in the nuance it captures: not one perspective on the relationships between tokens, but many, combined.

As we saw in the diagram, we can define multi-head attention as a mechanism in which we divide the attention mechanism into multiple heads, each producing its own context vector. We will look at two ways of achieving this. The first is stacking attention layers, and the second is attention weight splitting.

The most direct way to build multi-head attention is to take the single-head SelfAttention we already have and create N independent copies of it. Each head has its own Q, K, and V weights, so each one learns to attend differently. We run the same input through every head, then concatenate the results along the feature axis.

import * as tf from "@tensorflow/tfjs-core";
import { SelfAttention } from "@node-llm/core";
// x: [batch, seqLen, dIn]
const numHeads = 4;
const headDim = 2; // each head outputs a context of width M = headDim
// One independent self-attention per head.
const heads = Array.from(
{ length: numHeads },
() => new SelfAttention({ dIn: 4, dOut: headDim, causal: true }),
);
// Run every head on the same input, then concatenate along the last axis.
const context = tf.tidy(() => {
const perHead = heads.map((head) => head.forward(x)); // each [batch, seqLen, M]
return tf.concat(perHead, -1); // [batch, seqLen, numHeads * M] = [batch, seqLen, dOut]
});

Each head.forward(x) returns a context of shape [batch, seqLen, M], and stacking N of them side by side with tf.concat(perHead, -1) gives the full [batch, seqLen, dOut] context matrix from the diagram, where dOut = numHeads * M.

This is exactly what the MultiAttentionStacking class in @node-llm/core does. It holds numHeads independent SelfAttention heads and concatenates their context vectors on every forward pass.

import * as tf from "@tensorflow/tfjs-node";
import { DataLoader, LlmDataset, MultiAttentionStacking } from "@node-llm/core";
const numHeads = 4;
const headDim = 2; // width M of each head; final width dOut = numHeads * headDim = 8
const attn = new MultiAttentionStacking({ dIn: embDim, dOut: headDim, numHeads, causal: true });
for (const batch of loader) {
const x = tf.gather(tokenEmbedding, batch.inputs) as tf.Tensor3D;
const context = attn.forward(x); // [batchSize, contextSize, headDim * numHeads]
break;
}

To understand the weight splitting technique, we need to step back and understand the concept of a view, because weight splitting is built entirely around it.

Higher-dimensional matrices are not stored as matrices in GPU memory. They are stored as a single contiguous 1D array. So when we say a 2×4 or a 3×6 matrix, that is not the actual shape of the data on the hardware. That is a logical view: how we see it, and how we perform operations on it. The physical representation underneath is just a flat 1D sequence of numbers.

Imagine the GPU holds a single flat array of 8 numbers:

Physical memory: [10, 20, 30, 40, 50, 60, 70, 80]

We reshape a tensor with tf.reshape (the tfjs equivalent of PyTorch’s .view()). Depending on the shape we ask, tf partitions how it reads the memory block:

tf.reshape(x, [2, 4]) reads it as 2 rows, 4 columns:

[[10, 20, 30, 40],
[50, 60, 70, 80]]

tf.reshape(x, [4, 2]) reads it as 4 rows, 2 columns:

[[10, 20],
[30, 40],
[50, 60],
[70, 80]]

tf.reshape(x, [2, 2, 2]) reads it as a 3D box, 2 blocks of 2×2:

[[[10, 20], [30, 40]],
[[50, 60], [70, 80]]]

Because reshaping only changes the metadata that says how to read the array, without re-allocating or copying anything in memory, it runs instantaneously no matter how big the tensor is.

This is what makes attention weight splitting more effecient, because it does not move or copy matrices. It only changes how it index into the same flat array. That is why we can carve one big projection into N heads for free. As oppose to
stacking attention mechanism where we run matrix multiplication per head.

So if we descripe our input matrix shape as [batch, seqLen, dOut]. reshaping into [batch, seqLen, numHeads, headDim] does not move or copy a single number.

The MultiHeadAttention class in @node-llm/core takes the approach which real models use: a single set of Q, K, and V projections at the full output dimension dOut then a reshape. That splits each projection into N heads numHeads each of size dOut / numHeads.

These heads never get their own weight matrices. They are views into one big projection, so all N heads are computed in a single batched matrix multiplication.

// Project once at full width: [batch, seqLen, dOut].
// Then split into heads: [batch, seqLen, dOut] -> [batch, numHeads, seqLen, headDim].
const splitHeads = (t: tf.Tensor3D): tf.Tensor4D => {
const reshaped = tf.reshape(t, [batch, seqLen, numHeads, headDim]);
return tf.transpose(reshaped, [0, 2, 1, 3]) as tf.Tensor4D;
};
const q = splitHeads(project(x, wQuery));
const k = splitHeads(project(x, wKey));
const v = splitHeads(project(x, wValue));
// Scaled dot-product per head, all heads at once: [batch, numHeads, seqLen, seqLen].
const scores = tf.div(tf.matMul(q, k, false, true), Math.sqrt(headDim));
const weights = tf.softmax(scores, -1);
const context = tf.matMul(weights, v); // [batch, numHeads, seqLen, headDim]
// Merge heads back and mix them with an output projection Wo: [batch, seqLen, dOut].
const merged = tf.reshape(tf.transpose(context, [0, 2, 1, 3]), [batch, seqLen, dOut]);
const out = project(merged, wOut);

Both approaches produce a context matrix of the same [batch, seqLen, dOut] shape. Stacking is the easier one to reason about; weight splitting is the one you ship, because it does the whole thing in a handful of tensor ops with no per-head operation.

Reference (Python) Here (TypeScript)
nn.ModuleList([SelfAttention(...) for _ in range(n)]) Array.from({ length: numHeads }, () => new SelfAttention(...))
torch.cat([h(x) for h in heads], dim=-1) tf.concat(perHead, -1)
x.view(b, seqLen, numHeads, headDim) tf.reshape(t, [batch, seqLen, numHeads, headDim])
x.transpose(1, 2) tf.transpose(reshaped, [0, 2, 1, 3])
attn_scores @ values tf.matMul(weights, v)
self.out_proj(context) project(merged, wOut)