21: Attention
A note on this chapter
- Everything up to chapter 19 is my write-up of Professor Ng's 2011 course - this chapter is not
- Attention post-dates the course by some years
- The mechanism in the form used here comes from Neural Machine Translation by Jointly Learning to Align and Translate (Bahdanau et al., 2014) and Attention Is All You Need (Vaswani et al., 2017)
- So there's no lecture behind this one, and the diagrams are my own
- I've added it because attention is the piece that connects these notes to essentially everything that has happened since
- It builds directly on things covered earlier - logistic units, softmax, vectorization, learned weight matrices
- If you've worked through chapters 08 and 09 you already have what you need
Why attention?
- Say we want to map one sequence to another
- Translate a French sentence to English
- Summarize a paragraph
- Predict the next word
- The obvious approach is an encoder-decoder
- Encoder reads the input sequence one step at a time, updating a hidden state
- Final hidden state is a fixed-length vector - the "meaning" of the input
- Decoder generates the output sequence from that vector
The problem with this
- Everything has to squeeze through one fixed-length vector
- This is the bottleneck
- A 5 word sentence and a 50 word sentence get the same number of numbers to describe them
- Performance degrades badly as sentences get longer - which is exactly what you'd expect
- Long-range dependencies get lost
- If word 40 depends on word 2, that information has to survive 38 update steps
- Signal decays, gradients vanish - the same problem we saw with deep networks in chapter 09, just stretched along time instead of depth
- It's inherently sequential
- Step t needs step t-1, so you can't parallelize over the sequence
- Given the whole point of vectorization (chapter 04) was to stop looping, this hurts
- The fix is almost embarrassingly simple
- Don't force everything through one vector
- Keep all the encoder states around
- At each output step, let the decoder decide which of them to look at
- That "decide which to look at, and by how much" is attention
Attention as a soft lookup
- The cleanest way to think about attention is as a dictionary lookup that has been made differentiable
- An ordinary lookup
- You have a set of key-value pairs
- You supply a query
- You find the key that matches, and return its value
- This is a hard lookup - one key wins, everything else returns nothing
- The problem with a hard lookup is that it isn't differentiable
- "Which key matched" is a step function of the query
- Derivative is zero everywhere and undefined at the jump - no gradient, nothing to learn from
- So we soften it
- Score the query against every key - how well does each one match?
- Turn those scores into weights that are positive and sum to 1 (softmax, exactly as in chapter 09)
- Return the weighted average of all the values
- Now everything is smooth
- Nudge the query slightly, the weights shift slightly, the output shifts slightly
- Which means we can backpropagate through it and learn what to attend to
- If one score is much larger than the rest, softmax puts nearly all the weight there
- So a soft lookup can approximate a hard one when it wants to
- We get the behaviour of a lookup and the trainability of a smooth function
Scaled dot-product attention
- Now we need to pick an actual scoring function
- Simplest sensible choice - the dot product
- score(q, ki) = q · ki
- Large when the two vectors point the same way, near zero when orthogonal
- Which is exactly the notion of "similarity" we want
- And it's a matrix multiply, so it's fast - no extra parameters at all
Why divide by the square root of dk?
- There's a problem with raw dot products as vectors get longer
- Suppose the components of q and k are independent with mean 0 and variance 1
- Their dot product is a sum of dk such products
- So it has mean 0 and variance dk - i.e. typical magnitude grows like √dk
- With dk = 64 the scores are routinely in the tens
- Feed those into softmax and it saturates - one weight goes to ~1, the rest to ~0
- In the flat region the gradient is vanishingly small
- Same problem as a saturated sigmoid back in chapter 06, for the same reason
- So divide the scores by √dk
- Puts the variance back to about 1 regardless of dimension
- Keeps softmax in the part of its range where it has useful gradients
- It's a small detail that matters a lot in practice
The whole thing in matrix form
- Stack the queries, keys and values as rows of matrices - Q, K and V
- Q is [n x dk] for n queries
- K is [m x dk], V is [m x dv] for m key-value pairs
- Then the entire operation for every query at once is
Attention(Q, K, V) = softmax( QKT / √dk ) V
- Worth walking through the shapes, since this is where it clicks
- QKT is [n x m] - every query scored against every key
- Divide by √dk, then softmax along each row - so each query's weights sum to 1
- Multiply by V, which is [m x dv] - gives [n x dv]
- One output vector per query, each a weighted average of the values
- Note there are no learned parameters in this function
- It's two matrix multiplies and a softmax
- All the learning lives in how Q, K and V get produced - which is next
Self-attention
- So far Q, K and V were handed to us - in self-attention the sequence produces all three itself
- Start with the sequence as a matrix X, one row per token
- Learn three weight matrices - WQ, WK and WV
- Q = X WQ
- K = X WK
- V = X WV
- Each token therefore emits three different vectors, and it's worth being clear about what each is for
- Query - what this token is looking for
- Key - what this token offers to others looking for something
- Value - what this token actually contributes if attended to
- Splitting key from value is the subtle part, and it's easy to skip past
- What makes a token findable need not be what makes it useful
- A pronoun might match on grammatical role but contribute semantic content
- Keeping them separate lets the model learn the two independently
- Then run scaled dot-product attention on those three
- Every token attends to every token, including itself
- Output is a new sequence, same length, where each position now mixes in information from wherever it found it useful
- Note what we've bought
- Any position can reach any other in one step - no decay over 38 timesteps
- Every position is computed independently, so the whole sequence goes through at once - fully parallel
- Both of the RNN problems from the first section, gone
Cross-attention
- Same machinery, different sources
- Q comes from one sequence, K and V from another
- This is the original translation use - decoder queries the encoder
- Self-attention is just the case where all three come from the same place
Positional encoding
- There's a problem with what we just built, and it's a real one
- Attention is a weighted sum over the whole sequence
- Addition doesn't care about order
- So shuffle the input tokens and you get the same outputs, just shuffled to match
- The mechanism is permutation equivariant - it has no idea what order anything came in
- "the cat sat on the mat" and "the mat sat on the cat" are identical to it
- Clearly not acceptable for language
- Fix is to put the position into the input
- Build a vector that encodes "this is position i"
- Add it to the token's embedding before any attention happens
- Now two identical words at different positions have different representations, and the dot products can pick up on it
- Two broad approaches
- Learned - just an embedding table indexed by position, trained like any other parameter
- Fixed - sinusoids of different frequencies, which is what the original paper used
- Sinusoids have the nice property that they're defined for any position, so they extend past the longest training sequence
- Modern systems mostly encode relative position instead, but the principle is the same - attention needs to be told about order, because it cannot work it out
Masking
- Sometimes a query must not be allowed to see certain positions
- Causal masking - for generation
- When predicting token t we may only use tokens 1 to t
- Otherwise the model can read the answer off the input - it learns nothing and fails completely at generation time when the future isn't there
- This is the same trap as evaluating on your training set (chapter 10), just hidden inside the architecture
- Padding masking - for batching
- Sequences in a batch have different lengths, so short ones get padded
- Padding is meaningless and must not be attended to
- Implementation is a neat trick
- Set the disallowed scores to -∞ before the softmax
- e-∞ = 0, so those positions get exactly zero weight
- And the remaining weights still sum to 1 automatically, because softmax normalizes over what's left
- No renormalization step needed - it falls out
Multi-head attention
- One attention operation has a limitation
- It produces one set of weights per query
- So it can express one notion of relevance at a time
- But a word may relate to different words for entirely different reasons - one syntactic, one semantic
- Averaging those into a single weighting loses both
- So run several attention operations in parallel - each is a head
- Each head gets its own WQ, WK, WV
- So each learns its own idea of what to look for
- Concatenate the outputs and pass through one more learned matrix WO
- The cost is not what you might expect
- With h heads, each head works in dimension dmodel / h rather than dmodel
- e.g. dmodel = 512 with 8 heads gives 64 dimensions per head
- Total computation is about the same as one full-width head
- So we get several independent views essentially for free
- In a trained model heads do specialise, and visibly so
- Some track syntactic relations, some resolve pronouns, some just attend to the previous token
- Though it's easy to over-read this - plenty of heads do nothing interpretable at all, and many can be pruned with little loss
Attention inside a transformer block
- Attention on its own isn't a network - it's one layer in a repeating block
- A block is
- Multi-head self-attention
- A position-wise feed-forward network - an ordinary two-layer net applied to each position separately, with the same weights
- A residual connection around each of those two
- Layer normalization
- Why each of the extra pieces is there
- Residual connections - output = x + f(x)
- Gives gradients a direct path back through the whole stack
- Without them, deep stacks simply don't train - same story as chapter 09
- Layer normalization
- Keeps activations at a sane scale across the layer
- Feature scaling from chapter 04, applied inside the network rather than to the inputs
- The batch-wise sibling of this idea - batch normalization - is in chapter 20
- Feed-forward network
- Attention is a weighted average - it moves information around but is linear in V
- You need somewhere to actually transform it non-linearly, and this is that place
- Usually widened by 4x internally, and it holds most of the parameters
- Residual connections - output = x + f(x)
- Stack these blocks and you have a transformer
- Attention mixes information between positions
- The feed-forward layer processes it at each position
- Repeat - and essentially every large model today is this, made deeper and wider
The cost of attention
- The catch is in the score matrix
- QKT is [n x n] for a sequence of length n
- So compute and memory both scale as O(n2)
- Double the sequence length, quadruple the cost
- Which is fine at n = 512 and ruinous at n = 100,000
- This is the single biggest constraint on how much context a model can take
- It's the reason context length is a headline number people quote
- Broad approaches to it
- Sparse attention - don't attend to everything; use local windows, or a few global tokens
- Low-rank / kernel methods - approximate the softmax so you never form the n x n matrix
- IO-aware exact attention - keep the maths exact but never write the full matrix to slow memory (this is what FlashAttention does, and it's the one that's been most widely adopted)
- Worth noting the third one won not by approximating better but by taking the hardware seriously - a good reminder that the bottleneck isn't always where the maths says it is
Summary
- Attention is a differentiable lookup
- Score a query against every key, softmax to weights, average the values
- Attention(Q, K, V) = softmax(QKT / √dk) V
- Self-attention has the sequence generate its own Q, K and V through learned matrices
- Any position reaches any other in one step
- The whole sequence computes in parallel
- Because it's order-blind, position has to be added to the input explicitly
- Masking with -∞ before the softmax restricts what a query may see, and keeps the weights normalized for free
- Multiple heads give several independent notions of relevance at roughly the cost of one
- The n2 score matrix is the fundamental cost, and most of the engineering effort goes there
- The through-line from the rest of these notes
- Softmax from chapter 09, feature scaling from chapter 04, vectorization from chapter 04, the vanishing gradient problem from chapter 09
- None of the ingredients are new - what's new is arranging them so the model decides for itself what to look at