Transformer Architecture Q&A

Course Content

Transformer Architecture Q&A

6 sections · 60 lessons

Describe scaled dot-product attention and why divide scores by √d_k?


Attention weights for the worked example0.5060.1860.3070.1860.5060.3070.3840.2330.384key 1key 2key 3query 1query 2query 3Q K-transpose divided by sqrt(4) = 2, then softmax along each row.
Each row is a soft preference only because the scores were halved first; at d_k = 128 unscaled scores make every row nearly one-hot and gradients vanish.

What you need to know

Text
Attention(Q, K, V) = softmax( Q Kᵀ / sqrt(d_k) ) V

Worked by hand: 3 tokens, d_k = 4

Take queries, keys and values (as if already projected):

Text
Q = [1 0 1 0]    K = [1 0 1 0]    V = [1 0]    [0 1 0 1]        [0 1 0 1]        [0 1]    [1 1 1 0]        [1 1 0 0]        [1 1]Step 1: Q Kᵀ            Step 2: divide by sqrt(4) = 2[2 0 1]                 [1.0 0.0 0.5][0 2 1]                 [0.0 1.0 0.5][2 1 2]                 [1.0 0.5 1.0]Step 3: softmax each row        Step 4: weights times V[0.506 0.186 0.307]             [0.814 0.494][0.186 0.506 0.307]             [0.494 0.814][0.384 0.233 0.384]             [0.767 0.616]

Check one entry: row 1 scores [1.0, 0.0, 0.5] give e^1 = 2.718, e^0 = 1, e^0.5 = 1.649; the sum is 5.367, so the weights are 0.506, 0.186, 0.307. Output row 1 is 0.506·[1,0] + 0.186·[0,1] + 0.307·[1,1] = [0.814, 0.494].

Why sqrt(d_k): the variance argument

A dot product q · k adds up d_k products. If each entry has mean 0 and variance 1 and they are independent, each product has variance 1, so the sum has variance d_k and standard deviation sqrt(d_k). We checked this with 10,000 random pairs:

d_kstd of raw q · kstd after / sqrt(d_k)
42.01.0
648.11.0
12811.31.0

Why large scores hurt

Softmax of [2, 1, 0.5] is [0.63, 0.23, 0.14] — a soft preference. Multiply those scores by 11.3 (what an unscaled d_k = 128 head can produce) and softmax gives about [1.0, 0.0, 0.0]. The output is then a hard pick. Worse, the gradient of softmax at a near one-hot point is almost zero, so the query and key weights stop learning.

Scaling is not a modelling choice; it is a normalisation that keeps softmax in its useful range as head size changes.

A real-life example

A team building an in-house code-completion model wrote a custom attention kernel to try a new idea, and forgot the / sqrt(d_k) with d_head = 128. Training loss dropped for a few hundred steps and then flattened far above the baseline. Attention-entropy plots showed almost every head putting all its weight on one token from early in training.

Adding the scale fixed it. In PyTorch, torch.nn.functional.scaled_dot_product_attention applies 1 / sqrt(d_k) by default, which is a good reason to use it (it also dispatches to FlashAttention kernels) rather than hand-written maths.

Follow-up questions to expect

  • "Why not divide by d_k instead?" — That over-shrinks: variance becomes 1 / d_k, scores cluster near 0, and softmax becomes almost uniform.
  • "Is the scale ever different?" — Some models use a learned or tuned scale, or normalise Q and K (QK-norm) to control score size; the goal is the same.
  • "What does FlashAttention change?" — Not the maths: it computes the same result in tiles without writing the T × T matrix to GPU memory.