Skip to main content
Create your own

Scaled Dot-Product Self-Attention

Hello! Welcome back to our journey into sequence modeling.

In our last lesson, we enhanced RNN-based seq2seq models with Bahdanau and Luong attention. We saw how allowing the decoder to selectively look at different parts of the encoder's output solved the information bottleneck problem. However, those models still relied on the sequential processing of RNNs, which can be slow and struggles to capture very long-range dependencies.

Today, we take a revolutionary step by asking: what if we build a model entirely out of attention, removing the recurrence altogether? This brings us to self-attention. Your learning outcome for this lesson is to derive and implement scaled dot-product self-attention, the fundamental mechanism that powers the Transformer architecture and nearly all modern large-scale AI models.

1. From Cross-Attention to Self-Attention: The Q, K, V Framework

In the previous lesson, we saw attention as a mechanism where a decoder queries an encoder. To generalize this, we introduce the concepts of Queries (Q), Keys (K), and Values (V).

  • Query: A representation of the current word, asking "Who should I pay attention to?"
  • Key: A representation of another word, advertising "This is what I am." The query is matched against all keys.
  • Value: A representation of that same other word, providing "This is the information I contain."

In the seq2seq model, the decoder's hidden state was the Query, while the encoder's hidden states served as both the Keys and Values. This is called cross-attention because Q comes from a different sequence than K and V.

Self-attention flips this on its head: the Queries, Keys, and Values all come from the same sequence. Each word in the input sequence generates its own Q, K, and V vectors. It then uses its Query to score against every other word's Key, and the resulting attention weights are used to create a weighted sum of all words' Values.

The result is a new, context-rich representation for each word that incorporates information from the entire sequence, weighted by relevance. This allows the model to capture dependencies like "The cat, which was on the mat, ... was sleeping", directly connecting "was" back to "cat" regardless of the distance.

2. The Scaled Dot-Product Attention Formula

The most common implementation of self-attention is Scaled Dot-Product Attention. The formula is compact and highly efficient.

Let's break this down step-by-step.

Scaled Dot Product Attention Mechanism Diagram
This diagram shows the flow of the scaled dot-product attention mechanism. An input sequence `X` is projected into Queries (Q), Keys (K), and Values (V) using linear layers. The attention scores are computed, scaled, and normalized via softmax, then used to create a weighted sum of the Values.

The process is as follows:

  1. Compute Scores (QK^T): We perform a matrix multiplication between the Queries Q and the transpose of the Keys K. If Q has shape (num_queries, d_k) and K has shape (num_keys, d_k), the resulting score matrix will have shape (num_queries, num_keys). Each element (i, j) in this matrix is the dot product of the i-th query with the j-th key, representing their similarity.
  2. Scale (/ sqrt(d_k)): We divide all scores by the square root of the dimension of the key vectors, d_k. This is a critical step for stabilizing the training process, and we will derive the reason for it shortly.
  3. Normalize (softmax): We apply a softmax function along each row of the scaled score matrix. This converts the raw scores into a probability distribution (the attention weights), ensuring they are all positive and sum to 1.
  4. Compute Output (* V): Finally, we multiply the attention weights matrix by the Values V matrix. This produces an output where each row is a weighted sum of all the value vectors, with the weights determined by the attention scores. The output is a new representation for each query position, now infused with context from the entire sequence.

To see a detailed walkthrough of this calculation, let's turn to a video that explains the mechanics.

Attention is all you need (Transformer) - Model explanation (including math), Inference and Training

This video by Umar Jamil provides a clear, step-by-step visualization of the self-attention calculation. It's a great way to build an initial intuition for how the formula works in practice.

Watch from 20:11 to 25:30. Focus on how the Q, K, and V matrices are used to compute the attention score matrix, how softmax is applied, and how the final output is a new set of embeddings that captures relationships between words.

3. Derivation: Why Scale by ?

Now for the core of today's lesson: deriving the necessity of the scaling factor, .

Without scaling, the formula is just softmax(QK^T)V. What's wrong with that?
Let's assume, for simplicity, that the components of our Q and K vectors are independent random variables with a mean of 0 and a variance of 1. A property of statistics states that the dot product of two such vectors, , will have a mean of 0 but a variance of .

This means that as the embedding dimension d_k grows, the variance of the dot products also grows. A high variance means the dot product scores will be more spread out. Some will be very large, and others very small.

When these large-magnitude scores are fed into the softmax function, it saturates. The largest score gets pushed towards 1, while all others are pushed towards 0. This creates a nearly one-hot attention distribution, where the model pays attention to only one word and ignores all others. The gradient of the softmax function for these saturated inputs becomes extremely small, effectively halting learning. This is the classic vanishing gradient problem.

To counteract this, we need to bring the variance of the dot products back down to 1, regardless of the dimension d_k. The video below provides a fantastic mathematical and intuitive explanation of why dividing by accomplishes exactly this.

Why Scaling by the Square Root of Dimensions Matters in Attention | Transformers in Deep Learning

This video by 'Learn With Jay' is dedicated entirely to explaining the mathematical justification for the scaling factor. It's an excellent deep dive that will satisfy the 'derive' part of our learning outcome.

Watch from 03:38 to 19:09. The video is structured to answer our question perfectly: Posing the Question (03:38 - 04:32): Introduction to the scaling factor. Variance and Dimension (04:32 - 11:55): Why the variance of the dot product scales with dimension d_k. High Variance and Vanishing Gradients (12:33 - 15:30): How this high variance causes problems for the softmax function. The Solution (15:30 - 19:09): The mathematical derivation showing that dividing by sqrt(d_k) normalizes the variance.

For a concise, text-based summary of this derivation, please review the first answer on the following StackExchange page. It clearly lays out the argument about variance and standard deviation.

Why use a 'square root' in the scaled dot product

This StackExchange answer provides a succinct and well-written mathematical justification for the scaling factor, reinforcing the concepts from the video.

Read the first and most upvoted answer by user 'pi-tau'. Focus on the argument: Var(α) = d_k, so std(α) = sqrt(d_k), and scaling by this standard deviation normalizes the scores.

4. Implementation in PyTorch

With the theory firmly in place, let's move on to implementation. We will build a DotProductAttention module in PyTorch that encapsulates the formula we've just derived.

Given your background in Python and software engineering, you'll appreciate how the mathematical formula maps cleanly to code, especially using optimized library functions for matrix operations. Two important practical considerations are:

  1. Masking: We often process sentences in batches, padding shorter sentences to match the length of the longest one. We must "mask" these padding tokens so they don't participate in the attention calculation. Similarly, in a decoder context, we must prevent a position from attending to future positions (causal masking). This is typically done by adding a very large negative number (like -1e9) to the attention scores before the softmax step, which effectively zeros them out.
  2. Batch Matrix Multiplication: To process a batch of sequences efficiently, we use batch matrix multiplication (torch.bmm or the @ operator on 3D+ tensors), which performs independent matrix multiplications for each item in the batch.

The following resource provides an excellent, commented implementation of scaled dot-product attention that also handles masking.

Attention Scoring Functions and Implementation

Let's translate theory into practice. This resource from the 'Dive into Deep Learning' book provides a complete, production-quality implementation in PyTorch.

From the page '11.3. Attention Scoring Functions', focus on two key parts: Masked Softmax Operation: Read this section to understand how to handle variable-length sequences by masking padding tokens. The code shows how to add a large negative value before applying softmax. Scaled Dot Product Attention: Study this section carefully. Analyze the DotProductAttention class. Trace how the formula softmax(QK^T/sqrt(d))V is implemented in the forward method using torch.bmm and math.sqrt(d).

Test your understanding!

Consider the DotProductAttention module from the d2l.ai resource. If you pass in queries with shape (batch_size=32, num_queries=50, d_k=64), keys with shape (batch_size=32, num_keys=50, d_k=64), and values with shape (batch_size=32, num_keys=50, d_v=128), what is the shape of the final output tensor?

Show answer

The output shape will be (batch_size=32, num_queries=50, d_v=128).

Let's trace the dimensions:

  1. queries is (32, 50, 64).
  2. keys.transpose(1, 2) is (32, 64, 50).
  3. scores = torch.bmm(queries, keys.transpose(1, 2)) results in (32, 50, 50).
  4. softmax(scores) is also (32, 50, 50). This is the attention weights matrix.
  5. values is (32, 50, 128). Note that the num_keys dimension must match the last dimension of the attention weights.
  6. torch.bmm(attention_weights, values) results in (32, 50, 128).

The output has the same shape as the queries, but with the value dimension d_v as the final dimension.

Conclusion

In this lesson, you've mastered the core computational unit of the Transformer. By moving from RNN-based cross-attention to pure self-attention, we've unlocked a mechanism that is highly parallelizable and powerful at capturing long-range dependencies.

Key Takeaways:

  • Self-attention relates different positions of a single sequence to compute a new representation for each position, framed using Queries, Keys, and Values (Q, K, V).
  • The scaled dot-product attention formula is .
  • The scaling factor is crucial. It normalizes the variance of the dot products to 1, preventing the softmax function from saturating and ensuring stable gradient flow during training.
  • Practical implementations require masking to handle padding and causality, and rely on efficient batch matrix multiplication.

Preview of the Next Lesson:
A single self-attention mechanism is like having one person read a sentence and pick out one type of relationship. But what if we want to capture multiple, complex relationships simultaneously (e.g., syntactic dependencies, co-reference, semantic similarity)? In our next lesson, we will learn how to do just that by building a multi-head attention layer, which runs scaled dot-product attention multiple times in parallel. This is the final component we need before assembling a full Transformer block.

Can't find a good explanation? Sign up and we'll make it for you

Sign up