holehouse.org Blog Machine learning notes

21: Attention

A note on this chapter

Why attention?

The problem with this

Without attention - one fixed vector carries everything h1 h2 h3 h4 c dec everything via one c With attention - the decoder looks back at every state h1 h2 h3 h4 dec weights the states it needs
The bottleneck, and what attention replaces it with.

Attention as a soft lookup

query q k1 k2 k3 k4 score score score score softmax 0.7 0.2 0.05 0.05 sum to 1 v1 v2 v3 v4 output
Score against every key, softmax the scores, average the values.

Scaled dot-product attention

Why divide by the square root of dk?

The whole thing in matrix form

Attention(Q, K, V) = softmax( QKT / √dk ) V
Q, K QK T ÷ √d k mask optional softmax × V out
Scaled dot-product attention, end to end.

Self-attention

Cross-attention

Positional encoding

Masking

keys (what we attend to) queries 1 2 3 4 1 2 3 4 allowed masked to -inf query 3 may use keys 1, 2 and 3 but never 4
A causal mask - each query sees only itself and what came before.

Multi-head attention

X head 1 head 2 head 3 head h d/h d/h d/h d/h concat W O out
Several heads, each in a smaller dimension, concatenated and projected.

Attention inside a transformer block

The cost of attention

Summary