Seyed Masoud Hosseini · Overview · Study log · Weekly summaries · Ideas · Search · Transcript · RSS feed

Deep Learning Systems · Lecture 22 of 25 · 1:01:12

Lecture 21: Transformer Implementation

Lecture 21 - Transformer Implementation on YouTube

Study guide

What this lecture covers

This lecture answers how the Transformer block described conceptually in the previous lecture translates into actual code. It builds a NumPy implementation of self-attention, extends it to mini-batches, adds multi-head attention, and finally assembles a full Transformer encoder layer, checking each step against PyTorch's reference implementation.

After watching, you can implement self-attention, batched self-attention, multi-head attention, and a Transformer encoder block in plain array code, explain why Transformers need a genuine batch matrix multiplication rather than a reshape-based trick, and describe how PyTorch packs the key/query/value weights and multiple heads into single matrices internally.

Key ideas

  • Combined K, Q, V weights: rather than three separate projections, WK, WQ, WV are concatenated into one matrix so X @ W_KQV is computed in a single, larger matrix multiplication, then split into three components; this mirrors how PyTorch stores its in_proj_weight.
  • Output projection: after the attention-weighted sum of V, a further linear layer W_out is applied; this is not strictly necessary for single-head attention but matters once multiple heads are combined.
  • Mask by addition: the causal mask is added (with -inf on masked entries) rather than subtracted, matching PyTorch's convention.
  • Batch matrix multiplication is genuinely different: unlike convolution batching, which can be reduced to one big 2D matrix multiply via reshaping, batched self-attention requires computing K_i @ Q_i.T independently for each batch element, which is a true batch matrix multiply (bmm), not a reshape trick.
  • Batch-first layout for Transformers: unlike RNNs, which need (time, batch, hidden) for contiguous per-timestep slices, Transformers should use (batch, time, hidden) since attention multiplies over the trailing two dimensions.
  • Multi-head attention: instead of one large attention computation per layer, K, Q, V are split along the feature dimension into H heads, attention is computed independently per head with the score scaled by sqrt(D/H), and the outputs are concatenated back together; the intuition is that a single large softmax wastes the non-linearity's capacity.
  • PyTorch quirk: nn.MultiheadAttention returns the average attention matrix across heads, not each head's individual attention matrix.
  • Transformer block in code: self-attention, then residual add and layer norm, then a two-layer ReLU feed-forward network, then another residual add and layer norm, implementable in under 20 lines of NumPy.

Walkthrough

Basic self-attention as a layer (0:00)

The lecture reframes self-attention as a module with its own weights WK, WQ, WV, and an output projection W_out, computing softmax(X@WK @ (X@WQ).T / sqrt(D)) @ (X@WV) @ W_out. Biases are set to zero for simplicity. A NumPy softmax helper is written first since NumPy has no built-in.

Implementing and testing single-head attention (2:00)

WK, WQ, WV are combined into one W_KQV matrix so the three projections happen in a single matrix multiplication, then split via np.split along the last axis. The resulting attention() function is compared against PyTorch's nn.MultiheadAttention with one head; the weights are pulled from PyTorch's in_proj_weight and out_proj.weight, and the outputs match to numerical precision.

Mini-batching and why it needs a real batch matmul (15:10)

The lecture argues Transformers should use (batch, T, D) layout, unlike the (T, batch, D) layout used for RNNs, since attention multiplies over the trailing two dimensions and this keeps memory contiguous. It then demonstrates in NumPy that ordinary matrix multiplication of higher-rank tensors (as used for convolution batching) is not the same operation as batch matrix multiplication: multiplying a tensor by a plain 2D matrix flattens the leading dimensions, while true batch matmul computes an independent matrix product per batch element. Self-attention needs the latter. The attention function is then generalized to work in both batched and unbatched form by splitting and transposing on the last axes.

Multi-head attention (30:31)

The motivation given is that a single large dot product per position wastes the softmax non-linearity's expressive power, so K, Q, V are split along the feature dimension into H heads, each of size D/H, attention is computed per head, and results are concatenated. The implementation reshapes each of K, Q, V from (B, T, D) to (B, H, T, D/H) using reshape and swapaxes, runs the same attention computation batched over heads, then swaps back and reshapes to (B, T, D) before applying W_out. This is checked against PyTorch with 4 heads, matching to numerical precision; the lecture notes PyTorch's returned attention matrix is actually the average across heads, unlike this implementation which returns every head's matrix.

Assembling the Transformer block (44:38)

layer_norm and relu helpers are defined, then a transformer function combines them: Z = layer_norm(X + multi_head_attention(X, mask, W_KQV, W_out, heads)), followed by layer_norm(Z + relu(Z @ W_ff1) @ W_ff2). This is compared against PyTorch's nn.TransformerEncoderLayer (noting TransformerEncoderLayer, not the decoder variant, is what's normally wanted outside sequence-to-sequence translation setups), with weights copied over from the PyTorch module's attention and linear layers, and the outputs agree to within floating-point precision.

Before you watch

  • Watch the previous lecture on self-attention and Transformer architecture, since this lecture implements exactly the equations it derives.
  • Recall the earlier LSTM implementation lecture's discussion of batching and contiguous memory, which this lecture extends and contrasts with true batch matrix multiplication.
  • Be comfortable reading and writing NumPy array reshaping and axis operations (reshape, swapaxes, split).

Check your understanding

  1. Why can convolution batching be implemented with an ordinary 2D matrix multiply, while self-attention batching requires a genuine batch matrix multiply?
  2. Why does the lecture recommend (batch, T, D) layout for Transformers instead of the (T, batch, D) layout used for RNNs?
  3. What is the purpose of splitting K, Q, V into multiple heads rather than using one large attention computation?
  4. What does PyTorch's nn.MultiheadAttention return as its attention matrix when using multiple heads, and how does that differ from this lecture's implementation?

Vocabulary

self-attention (noun)
An attention mechanism where a sequence attends to itself to mix information across positions.
This lecture implements self-attention in NumPy.
concatenate (verb)
To join two or more arrays together end to end.
WK, WQ, and WV are concatenated into one matrix.
projection (noun)
A transformation that maps data into a new space using a weight matrix.
The output projection is applied after attention.
batch matrix multiplication (noun)
Multiplying many pairs of matrices independently, one pair per batch element.
Self-attention needs a real batch matrix multiplication.
reshape (verb)
To change an array's dimensions without changing its data.
Heads are formed by reshaping K, Q, and V.
flatten (verb)
To combine multiple dimensions of an array into one.
Ordinary matrix multiplication flattens the leading dimensions.
batch-first layout (noun)
An array ordering that puts the batch dimension before time.
Transformers use a batch-first layout instead of RNN's time-first one.
multi-head attention (noun)
Running several smaller attention computations in parallel and combining their results.
Multi-head attention splits K, Q, V into several heads.
head (noun)
One of several parallel attention computations inside multi-head attention.
Each head attends to a different part of the feature space.
expressive power (noun)
The range of different functions or patterns a model can represent.
Multiple heads increase the model's expressive power.
swap axes (phrase)
To exchange the order of two dimensions in an array.
We swap axes to move the head dimension next to time.
residual add (noun)
Adding a layer's input back onto its output.
A residual add follows the attention computation.
layer norm (noun)
A technique that normalizes each example's values to a standard scale.
Layer norm is applied after each residual add.
encoder layer (noun)
One block of a Transformer that processes an input sequence.
PyTorch's TransformerEncoderLayer matches this implementation.
floating-point precision (noun)
The level of accuracy a computer can store numbers with.
The two implementations agree to floating-point precision.
genuinely (adverb)
Truly, in a real and not superficial way.
Batched attention is genuinely different from convolution batching.
assemble (verb)
To put separate pieces together to build something complete.
The transformer function assembles attention and feed-forward layers.
reference implementation (noun)
A trusted, working version of code used to check a new version against.
PyTorch's TransformerEncoderLayer serves as a reference implementation.
bias (noun)
An extra learned number added to a layer's output.
Biases are set to zero for simplicity here.
helper function (noun)
A small piece of code written to support a larger task.
A softmax helper function is written first.

From the YouTube description

This lecture takes you through the implementation of a basic Transformer, including batching, multi-head attention, and the full Transformer block.

← Lecture 20: Transformers and Attention · Lecture 23: Model Deployment →