Skip to content
Léonel Vodounou
Chapters

Build a Large Language Model (from Scratch) · Chapter 3

Attention mechanisms

A deep dive into the heart of LLM architecture: simplified self-attention, trainable weights, causal attention and multi-head attention.

Léonel VODOUNOU

July 3, 2026 · 52 min read

Introduction

This chapter covers attention mechanisms, which are the core of the LLM architecture. The goal is to understand how they work, in isolation and mechanically, before integrating them into the full model in chapter 4.

Figure 3.1
Figure 3.1: The three main stages of building an LLM. This chapter focuses on step 2 of stage 1: implementing the attention mechanisms. Source: Raschka (2024).

Four variants of the attention mechanism will be implemented step by step, each building on the previous one:

  1. Simplified self-attention: a stripped-down version, with no trainable weights, to grasp the fundamental logic.
  2. Self-attention with trainable weights: the complete, trainable version.
  3. Causal attention: a mask is added so that the model cannot “see” future tokens while generating text token by token.
  4. Multi-head attention: several attention mechanisms run in parallel, letting the model capture different aspects of the relationships between tokens at the same time.
Figure 3.2
Figure 3.2: The four attention variants implemented in this chapter, from simplified self-attention to multi-head attention. Source: Raschka (2024).

The final multi-head attention implementation will be reused as is in the LLM architecture in the next chapter.


The Problem with Modeling Long Sequences

Before introducing the self-attention mechanism, it helps to understand the problem it solves, and why the architectures that came before it fell short.

The context: machine translation

Take a translation model. Translating word by word is impossible: the grammatical structures of the source and target languages differ too deeply.

Figure 3.3
Figure 3.3: Translating word by word is not enough. Translation requires understanding the overall context and realigning the grammar between the two languages. Source: Raschka (2024).

The classic solution is an encoder-decoder architecture: the encoder reads and compresses the input sequence, and the decoder produces the translated sequence from that compressed representation.

The encoder-decoder architecture with RNNs

Before Transformers, recurrent neural networks (RNNs) dominated this task. An RNN processes the sequence token by token, maintaining a hidden state that is updated at each step: a kind of internal memory carried along the sequence.

In an encoder-decoder RNN:

  • The encoder goes through the entire input sequence and tries to condense its meaning into a single final hidden state vector.
  • The decoder takes this vector as its starting point and generates the translation token by token, maintaining its own hidden state at each step.
Figure 3.4
Figure 3.4: In an encoder-decoder RNN, the encoder compresses the whole source sequence into a single final hidden state, which the decoder then uses to generate the translation token by token. Source: Raschka (2024).

The fundamental limitation

This final hidden state vector is the bottleneck of the architecture. All the information in the input sequence (however long it is) has to fit into this single vector. During decoding, the model no longer has access to the encoder’s intermediate hidden states: all it has is this compressed representation.

On short sentences, this works. On long sequences with distant dependencies, the context gets diluted and lost: this is long-range context loss.

This is exactly the limitation that motivated the invention of attention mechanisms: rather than forcing all the information into a single vector, let the decoder directly access all of the encoder’s hidden states, weighting their relevance according to the token being generated.


Capturing Dependencies with Attention Mechanisms

Bahdanau attention (2014), the first breakthrough

To get around the bottleneck of the single hidden vector, Bahdanau et al. proposed in 2014 a modification of the encoder-decoder RNN: instead of passing on only the final hidden state, all of the encoder’s intermediate hidden states remain accessible. At each decoding step, the decoder computes a relevance score for each of them and weights what it reads accordingly: these are called attention weights.

Figure 3.5
Figure 3.5: With an attention mechanism, the decoder can selectively access all the input tokens. Some tokens are more relevant than others for generating a given output token. This relevance is quantified by the attention weights. Source: Raschka (2024).

The Transformer (2017), the second breakthrough

Three years later came a more radical discovery: the RNN itself is not needed. The Transformer architecture, proposed in 2017, keeps the principle of attention weights but drops recurrence entirely. It introduces self-attention: a mechanism through which each token in a sequence can directly weigh the relevance of every other token in the same sequence, without going through intermediate hidden states.

Self-attention in one sentence

To compute its representation, each token looks at the whole sequence and decides, through learned weights, which other tokens it should “pay attention” to.

This mechanism is at the heart of modern LLMs such as GPT, and it is what this chapter implements from scratch.

Figure 3.6
Figure 3.6: Self-attention lets each position in the sequence interact with all the others and weigh their importance. This chapter codes the complete implementation. Source: Raschka (2024).

Self-Attention: Attending to Different Parts of the Input

The “self” in self-attention

The word self refers to the fact that the mechanism operates within one and the same sequence: each token computes its attention weights by comparing itself with all the other tokens of the same sequence. This is what sets it apart from classic Bahdanau attention, where the weights are computed between two different sequences (input and output).


Simplified Self-Attention, Without Trainable Weights

Before introducing trainable weights, we implement a stripped-down version to grasp the fundamental logic. The goal is to compute, for each input token x(i)x^{(i)}, a context vector z(i)z^{(i)}, that is, an enriched version of its embedding that incorporates information from all the other tokens in the sequence.

Figure 3.7
Figure 3.7: For each input token x(i)x^{(i)}, self-attention computes a context vector z(i)z^{(i)} by combining all the input vectors, weighted by the attention weights α21\alpha_{21} to α2T\alpha_{2T}. Here, we illustrate the computation of z(2)z^{(2)} for the token “journey”. Source: Raschka (2024).

We work on the sentence “Your journey starts with one step.”, already encoded as 3-dimensional vectors:

import torch

inputs = torch.tensor(
    [[0.43, 0.15, 0.89],  # Your     (x^1)
     [0.55, 0.87, 0.66],  # journey  (x^2)
     [0.57, 0.85, 0.64],  # starts   (x^3)
     [0.22, 0.58, 0.33],  # with     (x^4)
     [0.77, 0.25, 0.10],  # one      (x^5)
     [0.05, 0.80, 0.55]]  # step     (x^6)
)

Computing the context vector z(2)z^{(2)} takes three steps.


Step 1. Compute the attention scores (dot product)

For the query token x(2)x^{(2)} (“journey”), we compute its dot product with each token of the sequence. This score measures similarity: the higher it is, the more aligned the two tokens are in the vector space, and the more one should “pay attention” to the other.

query = inputs[1]  # x^(2): "journey"

attn_scores_2 = torch.empty(inputs.shape[0])
for i, x_i in enumerate(inputs):
    attn_scores_2[i] = torch.dot(x_i, query)

print(attn_scores_2)
tensor([0.9544, 1.4950, 1.4754, 0.8434, 0.7070, 1.0865])
Figure 3.8
Figure 3.8: The attention scores ω21\omega_{21} to ω2T\omega_{2T} are computed as the dot product between the query vector x(2)x^{(2)} and each input vector. Source: Raschka (2024).

Dot product. It multiplies two vectors element by element and sums the results. Writing the explicit loop is equivalent to using torch.dot:

res = 0.
for idx, element in enumerate(inputs[0]):
    res += inputs[0][idx] * query[idx]
print(res)                          # tensor(0.9544)
print(torch.dot(inputs[0], query))  # tensor(0.9544)

Step 2. Normalize to get the attention weights (softmax)

We normalize the scores to get attention weights that sum to 1, which makes them interpretable as relative importances.

A naive normalization by the sum does produce valid weights:

attn_weights_2_tmp = attn_scores_2 / attn_scores_2.sum()
print("Attention weights:", attn_weights_2_tmp)
print("Sum:", attn_weights_2_tmp.sum())
Attention weights: tensor([0.1455, 0.2278, 0.2249, 0.1285, 0.1077, 0.1656])
Sum: tensor(1.0000)

In practice, softmax is always preferred, for two deep reasons.

Why softmax rather than simple normalization?

Reason 1. Numerical stability. The naive softmax computes:

softmax(xi)=exi∑jexj\text{softmax}(x_i) = \frac{e^{x_i}}{\sum_j e^{x_j}}

In float32, exe^x gives inf as soon as x≈89x \approx 89 (overflow) and rounds to zero as soon as x≈−104x \approx -104 (underflow). On long sequences, attention scores can easily cross these thresholds:

torch.exp(torch.tensor(89.0))    # → tensor(inf): overflow
torch.exp(torch.tensor(-104.0))  # → tensor(0.): underflow

We then get inf/inf = NaN or 0/0 = NaN, and the gradient becomes unusable. PyTorch gets around this with the log-sum-exp trick: subtract the maximum score before exponentiating,

softmax(xi)=exi−max⁡(x)∑jexj−max⁡(x)\text{softmax}(x_i) = \frac{e^{x_i - \max(x)}}{\sum_j e^{x_j - \max(x)}}

which is mathematically identical (the e−max⁡(x)e^{-\max(x)} cancels out between numerator and denominator) but numerically stable: the largest value becomes e0=1e^0 = 1, and all the others fall within (0,1](0, 1].

Reason 2. Favorable gradient properties. Let’s compare the Jacobians of the two normalizations.

Simple normalization wi=xi/Sw_i = x_i / S, S=∑jxjS = \sum_j x_j:

∂wi∂xi=1−wiS∂wi∂xj=−wiS(j≠i)\frac{\partial w_i}{\partial x_i} = \frac{1 - w_i}{S} \qquad \frac{\partial w_i}{\partial x_j} = \frac{-w_i}{S} \quad (j \neq i)

The gradient depends on SS, the raw sum of the scores: if SS is large, the gradient collapses and the update signal propagates poorly. Moreover, this formula requires positive scores, which the dot product does not guarantee.

Softmax wi=exi/Zw_i = e^{x_i} / Z, Z=∑kexkZ = \sum_k e^{x_k}:

Deriving the Jacobian. We apply the quotient rule to wi=exi/Zw_i = e^{x_i}/Z.

Case i=ji = j:

∂wi∂xi=exiZ−exiexiZ2=exiZ−(exiZ)2=wi−wi2  ⇒  wi(1−wi)\frac{\partial w_i}{\partial x_i} = \frac{e^{x_i} Z - e^{x_i} e^{x_i}}{Z^2} = \frac{e^{x_i}}{Z} - \left(\frac{e^{x_i}}{Z}\right)^2 = w_i - w_i^2 \;\Rightarrow\; \boxed{w_i(1 - w_i)}

Case i≠ji \neq j: ∂exi∂xj=0\frac{\partial e^{x_i}}{\partial x_j} = 0, but ∂Z∂xj=exj\frac{\partial Z}{\partial x_j} = e^{x_j}, so:

∂wi∂xj=0⋅Z−exiexjZ2=−exiZ⋅exjZ  ⇒  −wiwj\frac{\partial w_i}{\partial x_j} = \frac{0 \cdot Z - e^{x_i} e^{x_j}}{Z^2} = -\frac{e^{x_i}}{Z}\cdot\frac{e^{x_j}}{Z} \;\Rightarrow\; \boxed{-w_i w_j}

Both cases combine using the Kronecker delta δij\delta_{ij}:

∂wi∂xj=wi(δij−wj)\boxed{\frac{\partial w_i}{\partial x_j} = w_i(\delta_{ij} - w_j)}

This Jacobian has three concrete advantages:

  • Gradient independent of the raw scale: it only depends on wi∈(0,1)w_i \in (0,1). The signal stays within a controlled range whatever the magnitude of the scores.
  • Strictly nonzero gradient: since wi∈(0,1)w_i \in (0,1) strictly, wi(1−wi)>0w_i(1-w_i) > 0. The correction signal never silently dies out.
  • Amplified differences: the exponential accentuates the gaps between scores. A small advantage ϵ\epsilon in score produces a more pronounced advantage in weight, giving sharper attention patterns and more informative gradients.
Implementation
# Naive version: unstable on large values
def softmax_naive(x):
    return torch.exp(x) / torch.exp(x).sum(dim=0)

# PyTorch version: built-in log-sum-exp, use this in practice
attn_weights_2 = torch.softmax(attn_scores_2, dim=0)
print("Attention weights:", attn_weights_2)
print("Sum:", attn_weights_2.sum())
Attention weights: tensor([0.1385, 0.2379, 0.2333, 0.1240, 0.1082, 0.1581])
Sum: tensor(1.)
Figure 3.9
Figure 3.9: The attention scores ω21\omega_{21} to ω2T\omega_{2T} are normalized with softmax to produce the attention weights α21\alpha_{21} to α2T\alpha_{2T}, which sum to 1. Source: Raschka (2024).

Step 3. Compute the context vector (weighted sum)

The context vector z(2)z^{(2)} is the weighted sum of all the input vectors, each multiplied by its attention weight:

z(2)=∑iα2i⋅x(i)z^{(2)} = \sum_{i} \alpha_{2i} \cdot x^{(i)}
context_vec_2 = torch.zeros(query.shape)
for i, x_i in enumerate(inputs):
    context_vec_2 += attn_weights_2[i] * x_i

print(context_vec_2)
tensor([0.4419, 0.6515, 0.5683])
Figure 3.10
Figure 3.10: z(2)z^{(2)} is the weighted sum of all the input vectors x(1)x^{(1)} to x(T)x^{(T)}, weighted by the corresponding attention weights. Source: Raschka (2024).

Unlike the raw embedding of “journey” [0.55, 0.87, 0.66], the context vector z(2)z^{(2)} [0.4419, 0.6515, 0.5683] encodes not only the meaning of “journey”, but also its weighted relationship with every other token in the sentence.

The next step is to generalize this computation to produce all the context vectors z(1)z^{(1)} to z(T)z^{(T)} at once.


Generalization: Attention Weights for All Tokens

We apply the same pipeline to all tokens at once, replacing the loops with matrix operations.

Figure 3.11
Figure 3.11: We extend the computation from the previous section to every row: one context vector z(i)z^{(i)} per token. Source: Raschka (2024).

Step 1. Attention scores for all pairs

attn_scores = inputs @ inputs.T
print(attn_scores)
tensor([[0.9995, 0.9544, 0.9422, 0.4753, 0.4576, 0.6310],
        [0.9544, 1.4950, 1.4754, 0.8434, 0.7070, 1.0865],
        ...])

Why inputs @ inputs.T is equivalent to the double loop. inputs is a matrix of shape (6, 3): 6 tokens, each represented by a 3-dimensional vector. inputs.T is its transpose, of shape (3, 6). Their matrix product gives a (6, 6) tensor whose element (i,j)(i, j) equals ∑kinputs[i,k]×inputs[j,k]\sum_k \text{inputs}[i,k] \times \text{inputs}[j,k], which is exactly the dot product between tokens ii and jj, what the double loop computed explicitly.

Step 2. Normalization with softmax

attn_weights = torch.softmax(attn_scores, dim=-1)

dim=-1 tells PyTorch to apply softmax along the last dimension, here the columns. Each row is normalized independently, so that its values sum to 1.

Step 3. Context vectors

all_context_vecs = attn_weights @ inputs
print(all_context_vecs)
tensor([[0.4421, 0.5931, 0.5790],
        [0.4419, 0.6515, 0.5683],
        [0.4431, 0.6496, 0.5671],
        [0.4304, 0.6298, 0.5510],
        [0.4671, 0.5910, 0.5266],
        [0.4177, 0.6503, 0.5645]])

attn_weights is (6, 6) and inputs is (6, 3): the product is (6, 3), where each row ii is the sum of all the input vectors weighted by token ii‘s attention weights.

This version still has no trainable parameters. The next section introduces the matrices WQW_Q, WKW_K, WVW_V to make the mechanism truly trainable.

Implementing Self-Attention with Trainable Weights

Figure 3.13
Figure 3.13: We extend the previous mechanism with trainable weight matrices. The extensions (causal mask, multiple heads) come next. Source: Raschka (2024).

The structural difference from the simplified version is the introduction of three trainable weight matrices WQW_Q, WKW_K, WVW_V, updated through backpropagation. They are what give the model the ability to learn which kind of similarity matters for the task, rather than measuring raw similarity between embeddings.


Computing the Attention Weights Step by Step

The three projection matrices and their roles

Each input token x(i)x^{(i)} is projected into three distinct subspaces through these matrices:

q(i)=x(i)WQk(i)=x(i)WKv(i)=x(i)WVq^{(i)} = x^{(i)} W_Q \qquad k^{(i)} = x^{(i)} W_K \qquad v^{(i)} = x^{(i)} W_V

The terms query, key and value come from databases and information retrieval:

  • Query q(i)q^{(i)}: the request issued by token ii to question all the other tokens: it lets token ii find, among them, the ones that matter for building its contextual representation.

  • Key k(j)k^{(j)}: the index key of token jj: each token exposes a key that is compared with the other tokens’ queries to determine its relevance.

  • Value v(j)v^{(j)}: the actual information content of token jj, like the value in a key-value pair of a dictionary: it is what actually gets passed on if the token is judged relevant.

The mechanism then works in three steps: compare the query q(i)q^{(i)} with all the keys through the dot product ωij=q(i)⋅k(j)\omega_{ij} = q^{(i)} \cdot k^{(j)}, identify the most relevant tokens through softmax, then retrieve the corresponding values weighted by these scores to form the context vector z(i)z^{(i)}.

It is a bit like a search engine: you type “best pizzerias in Paris” (query), Google compares it with the titles and metadata of every indexed page (keys), then returns the actual content of the most relevant pages (values).

In the simplified version seen earlier, q=k=v=xq = k = v = x: there was no learned projection, and the similarity measured was the raw closeness between embeddings. Introducing WQW_Q, WKW_K, WVW_V lets the model learn specialized representations for each of these three roles.

Matrix weights vs. attention weights. The entries of WQW_Q, WKW_K, WVW_V are learned parameters, that is, scalars optimized by gradient descent and fixed once training is over. The attention weights αij\alpha_{ij}, on the other hand, are dynamic: recomputed on every forward pass from the current input. These are two different uses of the word “weight”.


Implementation

x_2 = inputs[1]          # token "journey", shape: (3,)
d_in  = inputs.shape[1]  # 3
d_out = 2                 # output dimension (in GPT, d_in == d_out)

torch.manual_seed(123)
W_query = torch.nn.Parameter(torch.rand(d_in, d_out), requires_grad=False)
W_key   = torch.nn.Parameter(torch.rand(d_in, d_out), requires_grad=False)
W_value = torch.nn.Parameter(torch.rand(d_in, d_out), requires_grad=False)

requires_grad=False turns off gradient computation for these matrices here, only to keep the printed output short. In real training, we would set requires_grad=True so that they get updated through backpropagation.

We project x(2)x^{(2)} and all the tokens:

query_2 = x_2 @ W_query   # shape: (2,)
keys    = inputs @ W_key   # shape: (6, 2)
values  = inputs @ W_value # shape: (6, 2)

print(query_2)
tensor([0.4306, 1.4551])

We have projected the 6 three-dimensional tokens into a two-dimensional space. We only compute query_2 for the current token, but we need the keys and values of all the tokens to weight their contribution to the context vector of x(2)x^{(2)}.


Step 1. Attention scores

Figure 3.15
Figure 3.15: The scores are now computed between the query and key projections, and no longer directly between the raw embeddings. Source: Raschka (2024).
attn_scores_2 = query_2 @ keys.T
print(attn_scores_2)
tensor([1.2705, 1.8524, 1.8111, 1.0795, 0.5577, 1.5440])

Step 2. From scores to attention weights: the 1/√dₖ factor

Figure 3.16
Figure 3.16: The scores are scaled before the softmax. Source: Raschka (2024).
d_k = keys.shape[-1]  # dimension of the keys = 2
attn_weights_2 = torch.softmax(attn_scores_2 / d_k**0.5, dim=-1)
print(attn_weights_2)
tensor([0.1500, 0.2264, 0.2199, 0.1311, 0.0906, 0.1820])

Why divide by dk\sqrt{d_k}?

This is the justification that gives the architecture its name: scaled dot-product attention.

Assume the components of qq and kk are approximately i.i.d. with distribution N(0,1)\mathcal{N}(0, 1). The dot product q⋅k=∑l=1dkqlklq \cdot k = \sum_{l=1}^{d_k} q_l k_l is then a sum of dkd_k independent random variables with mean 0 and variance 1, so:

Var(q⋅k)=dk⇒σ(q⋅k)=dk\text{Var}(q \cdot k) = d_k \qquad \Rightarrow \qquad \sigma(q \cdot k) = \sqrt{d_k}
Proof

Let ql∼N(0,1)q_l \sim \mathcal{N}(0,1) and kl∼N(0,1)k_l \sim \mathcal{N}(0,1), independent.

Step 1. Variance of the product qlklq_l k_l

Var(qlkl)=E[ql2kl2]−E[qlkl]2\text{Var}(q_l k_l) = \mathbb{E}[q_l^2 k_l^2] - \mathbb{E}[q_l k_l]^2

Since ql,kl∼N(0,1)q_l, k_l \sim \mathcal{N}(0,1) are independent, their squares follow independent chi-squared distributions with 1 degree of freedom:

ql2∼χ2(1)kl2∼χ2(1)independentq_l^2 \sim \chi^2(1) \qquad k_l^2 \sim \chi^2(1) \quad \text{independent}

For a χ2(1)\chi^2(1) distribution: E[X]=1\mathbb{E}[X] = 1. By independence of ql2q_l^2 and kl2k_l^2:

E[ql2kl2]=E[ql2] E[kl2]=1×1=1\mathbb{E}[q_l^2 k_l^2] = \mathbb{E}[q_l^2]\,\mathbb{E}[k_l^2] = 1 \times 1 = 1

And E[qlkl]=E[ql] E[kl]=0\mathbb{E}[q_l k_l] = \mathbb{E}[q_l]\,\mathbb{E}[k_l] = 0, so:

Var(qlkl)=1−0=1\text{Var}(q_l k_l) = 1 - 0 = 1

Step 2. Variance of the sum q⋅k=∑l=1dkqlklq \cdot k = \sum_{l=1}^{d_k} q_l k_l

The terms qlklq_l k_l are independent of each other (the indices ll are distinct), so the variance of their sum is the sum of their variances:

Var(q⋅k)=∑l=1dkVar(qlkl)=∑l=1dk1=dk\text{Var}(q \cdot k) = \sum_{l=1}^{d_k} \text{Var}(q_l k_l) = \sum_{l=1}^{d_k} 1 = d_k

Step 3. Standard deviation

σ(q⋅k)=Var(q⋅k)=dk\sigma(q \cdot k) = \sqrt{\text{Var}(q \cdot k)} = \sqrt{d_k}

Step 4. After scaling by dk\sqrt{d_k}

Var ⁣(q⋅kdk)=1(dk)2 Var(q⋅k)=dkdk=1\text{Var}\!\left(\frac{q \cdot k}{\sqrt{d_k}}\right) = \frac{1}{(\sqrt{d_k})^2}\,\text{Var}(q \cdot k) = \frac{d_k}{d_k} = 1

The variance is brought back to 1 whatever the value of dkd_k. ■\blacksquare


For dk=1024d_k = 1024 (typical in GPT), the dot products have a standard deviation of about 32. Softmax therefore receives very spread-out inputs: large values push toward 1, small ones toward 0, and the distribution looks like a one-hot vector.

But the gradient of softmax is wi(1−wi)w_i(1 - w_i): when wi→1w_i \to 1 or wi→0w_i \to 0, this gradient tends to 0. We are back to the vanishing gradient problem: training stalls.

Dividing by dk\sqrt{d_k} brings the variance of the score back to 1:

Var ⁣(q⋅kdk)=dkdk=1\text{Var}\!\left(\frac{q \cdot k}{\sqrt{d_k}}\right) = \frac{d_k}{d_k} = 1

The softmax inputs stay within reasonable ranges whatever dkd_k is: the attention weights no longer collapse toward 0 or 1, and the gradients remain usable.


Step 3. Context vector

Figure 3.17
Figure 3.17: The context vector is the weighted sum of the value vectors, and no longer of the raw embeddings. Source: Raschka (2024).
context_vec_2 = attn_weights_2 @ values
print(context_vec_2)
tensor([0.3061, 0.8210])

The fundamental difference from the simplified version: we sum the value vectors v(j)v^{(j)}, not the raw embeddings x(j)x^{(j)}. WVW_V lets the model learn which information to extract from each token and pass on, independently of how that token is indexed (key) or how it queries the others (query).

Explicitly, the context vector of x(2)x^{(2)} is:

z(2)=∑j=1Tα2j⋅v(j)=∑j=1Tα2j⋅x(j)WVz^{(2)} = \sum_{j=1}^{T} \alpha_{2j} \cdot v^{(j)} = \sum_{j=1}^{T} \alpha_{2j} \cdot x^{(j)} W_V

In compact matrix form, for all the tokens at once:

Z=softmax ⁣(QK⊤dk)VZ = \text{softmax}\!\left(\frac{Q K^\top}{\sqrt{d_k}}\right) V

where Q=XWQQ = X W_Q, K=XWKK = X W_K, V=XWVV = X W_V are the projections of the whole sequence, and each row ii of ZZ is the context vector z(i)z^{(i)}.


Implementing a Compact Self-Attention Python Class

The previous code is reorganized into an nn.Module class, the building block of every PyTorch model, which automatically handles registering the parameters, updating them during training, and moving them to the GPU.

import torch.nn as nn

class SelfAttention_v1(nn.Module):
    def __init__(self, d_in, d_out):
        super().__init__()
        self.W_query = nn.Parameter(torch.rand(d_in, d_out))
        self.W_key   = nn.Parameter(torch.rand(d_in, d_out))
        self.W_value = nn.Parameter(torch.rand(d_in, d_out))

    def forward(self, x):
        keys    = x @ self.W_key
        queries = x @ self.W_query
        values  = x @ self.W_value
        attn_scores  = queries @ keys.T
        attn_weights = torch.softmax(attn_scores / keys.shape[-1]**0.5, dim=-1)
        context_vec  = attn_weights @ values
        return context_vec
torch.manual_seed(123)
sa_v1 = SelfAttention_v1(d_in, d_out)
print(sa_v1(inputs))
tensor([[0.2996, 0.8053],
        [0.3061, 0.8210],
        [0.3058, 0.8203],
        [0.2948, 0.7939],
        [0.2927, 0.7891],
        [0.2990, 0.8040]], grad_fn=<MmBackward0>)
Figure 3.18
Figure 3.18: Self-attention in matrix form. XX is projected into QQ, KK, VV through the three weight matrices. The scores QK⊤/dkQK^\top / \sqrt{d_k} are normalized with softmax, then multiplied by VV to produce ZZ. Source: Raschka (2024).

Improved version with nn.Linear

Weight initialization. In SelfAttention_v1, nn.Parameter(torch.rand(...)) draws each weight independently as Wij∼U[0,1)W_{ij} \sim \mathcal{U}[0, 1). This initialization is naive for two reasons:

  • The initial weights are all positive, which restricts the learning dynamics. A good initialization should instead differentiate the neurons effectively from the start by assigning random weights of varying signs (this is called symmetry breaking in the literature);

  • the variance Var⁡(Wij)=112\operatorname{Var}(W_{ij}) = \frac{1}{12} does not depend on dind_{in}, which makes the variance of the output XWXW explode when dind_{in} is large.
Proof

K=XWK = XW is the projection of the input embeddings into key space, with K∈Rn×doutK \in \mathbb{R}^{n \times d_{out}}: each row is a token’s key vector, each column a dimension of key space. We look at the scalar:

y=Ki,j=∑k=1dinxi,k Wk,jy = K_{i,j} = \sum_{k=1}^{d_{in}} x_{i,k}\, W_{k,j}

the jj-th component of token ii‘s key vector. It is the variance of this sum that we analyze to quantify the effect of the initialization.

Working assumptions. We assume xkx_k and WkW_k are independent, E[xk]=0\mathbb{E}[x_k] = 0 and Var⁡(xk)=σ2\operatorname{Var}(x_k) = \sigma^2 (a standard assumption, satisfied after normalizing the inputs).

Variance of one term xkWkx_k W_k. With E[xk]=0\mathbb{E}[x_k] = 0:

Var⁡(xkWk)=E[xk2Wk2]−E[xkWk]2=E[xk2] E[Wk2]−E[xk]2⏟= 0E[Wk]2=σ2 E[Wk2]\operatorname{Var}(x_k W_k) = \mathbb{E}[x_k^2 W_k^2] - \mathbb{E}[x_k W_k]^2 = \mathbb{E}[x_k^2]\,\mathbb{E}[W_k^2] - \underbrace{\mathbb{E}[x_k]^2}_{=\,0}\mathbb{E}[W_k]^2 = \sigma^2\,\mathbb{E}[W_k^2]

The key quantity is E[Wk2]\mathbb{E}[W_k^2], which breaks down as:

E[Wk2]=Var⁡(Wk)+E[Wk]2\mathbb{E}[W_k^2] = \operatorname{Var}(W_k) + \mathbb{E}[W_k]^2

Refresher: the uniform distribution U[a,b]\mathcal{U}[a,b].

Var⁡(W)=(b−a)212,E[W]=a+b2\operatorname{Var}(W) = \frac{(b-a)^2}{12}, \qquad \mathbb{E}[W] = \frac{a+b}{2}

Case torch.rand: Wk∼U[0,1)W_k \sim \mathcal{U}[0, 1).

E[Wk]=12,Var⁡(Wk)=112\mathbb{E}[W_k] = \frac{1}{2}, \qquad \operatorname{Var}(W_k) = \frac{1}{12}E[Wk2]=112+14=13\mathbb{E}[W_k^2] = \frac{1}{12} + \frac{1}{4} = \frac{1}{3}Var⁡(y)=∑k=1dinσ2⋅13=din σ23\operatorname{Var}(y) = \sum_{k=1}^{d_{in}} \sigma^2 \cdot \frac{1}{3} = \frac{d_{in}\,\sigma^2}{3}

The variance grows linearly with dind_{in}: for din=512d_{in} = 512, it is 512 times larger than the input variance.

Kaiming case: Wk∼U ⁣(−1din,1din)W_k \sim \mathcal{U}\!\left(-\dfrac{1}{\sqrt{d_{in}}}, \dfrac{1}{\sqrt{d_{in}}}\right).

E[Wk]=0,Var⁡(Wk)=(2din)212=13 din\mathbb{E}[W_k] = 0, \qquad \operatorname{Var}(W_k) = \frac{\left(\frac{2}{\sqrt{d_{in}}}\right)^2}{12} = \frac{1}{3\,d_{in}}E[Wk2]=13 din+0=13 din\mathbb{E}[W_k^2] = \frac{1}{3\,d_{in}} + 0 = \frac{1}{3\,d_{in}}Var⁡(y)=∑k=1dinσ2⋅13 din=σ23\operatorname{Var}(y) = \sum_{k=1}^{d_{in}} \sigma^2 \cdot \frac{1}{3\,d_{in}} = \frac{\sigma^2}{3}

The dind_{in} cancels out: the variance stays constant whatever the input dimension.

By default, nn.Linear uses Kaiming initialization (He, 2015), which sets:

Wij∼U ⁣(−1din, 1din)⟹Var⁡(Wij)=13 dinW_{ij} \sim \mathcal{U}\!\left(-\frac{1}{\sqrt{d_{in}}},\ \frac{1}{\sqrt{d_{in}}}\right) \qquad \Longrightarrow \qquad \operatorname{Var}(W_{ij}) = \frac{1}{3\,d_{in}}

which keeps the variance of the signal constant across layers, whatever the dimension.


nn.Linear as a linear operator. In SelfAttention_v1, we define W∈Rdin×doutW \in \mathbb{R}^{d_{in} \times d_{out}} as an nn.Parameter, and the forward pass computes:

K=XW,X∈Rn×din,K∈Rn×doutK = XW, \qquad X \in \mathbb{R}^{n \times d_{in}},\quad K \in \mathbb{R}^{n \times d_{out}}

nn.Linear(d_in, d_out, bias=False) internally stores a matrix W~∈Rdout×din\widetilde{W} \in \mathbb{R}^{d_{out} \times d_{in}}, and calling it on XX computes:

K=X W~⊤K = X\,\widetilde{W}^\top

The two expressions are identical if and only if W~=W⊤\widetilde{W} = W^\top, which is exactly how nn.Linear stores its weights. The two implementations are therefore strictly equivalent.

class SelfAttention_v2(nn.Module):
    def __init__(self, d_in, d_out, qkv_bias=False):
        super().__init__()
        self.W_query = nn.Linear(d_in, d_out, bias=qkv_bias)
        self.W_key   = nn.Linear(d_in, d_out, bias=qkv_bias)
        self.W_value = nn.Linear(d_in, d_out, bias=qkv_bias)

    def forward(self, x):
        keys    = self.W_key(x)
        queries = self.W_query(x)
        values  = self.W_value(x)
        attn_scores  = queries @ keys.T
        attn_weights = torch.softmax(attn_scores / keys.shape[-1]**0.5, dim=-1)
        context_vec  = attn_weights @ values
        return context_vec
torch.manual_seed(789)
sa_v2 = SelfAttention_v2(d_in, d_out)
print(sa_v2(inputs))
tensor([[-0.0739,  0.0713],
        [-0.0748,  0.0703],
        [-0.0749,  0.0702],
        [-0.0760,  0.0685],
        [-0.0763,  0.0679],
        [-0.0754,  0.0693]], grad_fn=<MmBackward0>)

The outputs of v1 and v2 differ only because the initial weights are different; the logic of the forward pass is identical.

Exercise 3.1: Comparing SelfAttention_v1 and SelfAttention_v2

Transfer the weights of SelfAttention_v2 to SelfAttention_v1 and check that both implementations produce the same outputs on inputs.

Solution:

nn.Linear stores its matrix as W~=W⊤\widetilde{W} = W^\top, so all we need is to transpose each weight before copying it.

sa_v1 = SelfAttention_v1(d_in, d_out)
sa_v1.W_query = torch.nn.Parameter(sa_v2.W_query.weight.T)
sa_v1.W_key   = torch.nn.Parameter(sa_v2.W_key.weight.T)
sa_v1.W_value = torch.nn.Parameter(sa_v2.W_value.weight.T)

print(torch.allclose(sa_v1(inputs), sa_v2(inputs)))
True

The next step adds two extensions to this mechanism: the causal mask, which prevents each token from accessing future tokens during generation, and multi-head attention, which runs several attention mechanisms in parallel to capture different kinds of relationships between tokens.


Hiding Future Words with Causal Attention

Causal attention (or masked attention) is a variant of self-attention that restricts each token to the tokens that come before it (and itself) in the sequence. This is the opposite of standard self-attention, which has the whole sequence available as input.

This constraint is fundamental for GPT-style LLMs: when predicting the next token, the model must not “see” future tokens; that would be information leakage.


Applying a Causal Attention Mask

There are two equivalent approaches to obtaining the masked attention weight matrix.

Naive approach: mask after softmax

Step 1. Compute the standard attention weights (softmax)

queries = sa_v2.W_query(inputs)
keys    = sa_v2.W_key(inputs)
attn_scores   = queries @ keys.T
attn_weights  = torch.softmax(attn_scores / keys.shape[-1]**0.5, dim=-1)
tensor([[0.1921, 0.1646, 0.1652, 0.1550, 0.1721, 0.1510],
        [0.2041, 0.1659, 0.1662, 0.1496, 0.1665, 0.1477],
        [0.2036, 0.1659, 0.1662, 0.1498, 0.1664, 0.1480],
        [0.1869, 0.1667, 0.1668, 0.1571, 0.1661, 0.1564],
        [0.1830, 0.1669, 0.1670, 0.1588, 0.1658, 0.1585],
        [0.1935, 0.1663, 0.1666, 0.1542, 0.1666, 0.1529]],
       grad_fn=<SoftmaxBackward0>)

Step 2. Build the lower triangular mask

context_length = attn_scores.shape[0]
mask_simple    = torch.tril(torch.ones(context_length, context_length))
tensor([[1., 0., 0., 0., 0., 0.],
        [1., 1., 0., 0., 0., 0.],
        [1., 1., 1., 0., 0., 0.],
        [1., 1., 1., 1., 0., 0.],
        [1., 1., 1., 1., 1., 0.],
        [1., 1., 1., 1., 1., 1.]])

torch.tril keeps the diagonal and the lower triangle and zeroes out the upper triangle, that is, exactly the future positions to mask.

Step 3. Zeroing: multiply the weights by the mask

masked_simple = attn_weights * mask_simple
tensor([[0.1921, 0.0000, 0.0000, 0.0000, 0.0000, 0.0000],
        [0.2041, 0.1659, 0.0000, 0.0000, 0.0000, 0.0000],
        [0.2036, 0.1659, 0.1662, 0.0000, 0.0000, 0.0000],
        [0.1869, 0.1667, 0.1668, 0.1571, 0.0000, 0.0000],
        [0.1830, 0.1669, 0.1670, 0.1588, 0.1658, 0.0000],
        [0.1935, 0.1663, 0.1666, 0.1542, 0.1666, 0.1529]],
       grad_fn=<MulBackward0>

The future positions are now 0, but the rows no longer sum to 1: the probability distribution is broken.

Step 4. Renormalize

row_sums          = masked_simple.sum(dim=-1, keepdim=True)
masked_simple_norm = masked_simple / row_sums
tensor([[1.0000, 0.0000, 0.0000, 0.0000, 0.0000, 0.0000],
        [0.5517, 0.4483, 0.0000, 0.0000, 0.0000, 0.0000],
        [0.3800, 0.3097, 0.3103, 0.0000, 0.0000, 0.0000],
        [0.2758, 0.2460, 0.2462, 0.2319, 0.0000, 0.0000],
        [0.2175, 0.1983, 0.1984, 0.1888, 0.1971, 0.0000],
        [0.1935, 0.1663, 0.1666, 0.1542, 0.1666, 0.1529]],
       grad_fn=<DivBackward0>)

Each row sums to 1. The masked weights are zero, and the remaining weights are redistributed proportionally.


Note: no information leakage One might worry that the future tokens have already “contaminated” the result through the initial softmax, before being zeroed out. They have not. Renormalizing after masking is mathematically equivalent to having computed the softmax over the unmasked positions only from the start. Future tokens have no influence at all on the final distribution.

Proof

Setup and notation

Consider a sequence of TT tokens. Each token at position jj has three vectors, obtained through linear projections of its embedding:

  • qj∈Rdkq_j \in \mathbb{R}^{d_k}: query vector
  • kj∈Rdkk_j \in \mathbb{R}^{d_k}: key vector
  • vj∈Rdvv_j \in \mathbb{R}^{d_v}: value vector

For the current token at position ii, the raw attention score with each token at position jj is:

eij=qi⊤kjdk∈R,j∈{1,…,T}e_{ij} = \frac{q_i^\top k_j}{\sqrt{d_k}} \in \mathbb{R}, \qquad j \in \{1, \dots, T\}

The output of the attention mechanism for token ii is:

zi=∑j=1Tαij vj,αij≥0,∑j=1Tαij=1z_i = \sum_{j=1}^{T} \alpha_{ij}\, v_j, \qquad \alpha_{ij} \geq 0, \quad \sum_{j=1}^{T} \alpha_{ij} = 1

Method 1. Global softmax, then masking and renormalization

Step 1. Softmax over all the tokens.

We compute the attention weights with no restriction:

αij=eeij∑m=1Teeim\alpha_{ij} = \frac{e^{e_{ij}}}{\displaystyle\sum_{m=1}^{T} e^{e_{im}}}

All the tokens contribute to the denominator, including future ones.

Step 2. Applying the triangular mask.

We zero out the weights corresponding to future tokens:

α~ij={αijif j≤i0if j>i\tilde{\alpha}_{ij} = \begin{cases} \alpha_{ij} & \text{if } j \leq i \\ 0 & \text{if } j > i \end{cases}

After this step, the weights no longer sum to 1: some probability mass has been removed.

Step 3. Renormalization.

We divide each weight by the sum of its row:

αij(1)=α~ij∑k=1Tα~ik\alpha_{ij}^{(1)} = \frac{\tilde{\alpha}_{ij}}{\displaystyle\sum_{k=1}^{T} \tilde{\alpha}_{ik}}

Explicit expansion. We simplify the numerator and the denominator separately.

For the numerator, we distinguish two cases depending on the position of jj relative to ii, applying first the definition of α~ij\tilde{\alpha}_{ij} (Step 2) and then that of αij\alpha_{ij} (Step 1):

α~ij={αij=eeij∑m=1Teeimif j≤i0if j>i\tilde{\alpha}_{ij} = \begin{cases} \alpha_{ij} = \dfrac{e^{e_{ij}}}{\displaystyle\sum_{m=1}^{T} e^{e_{im}}} & \text{if } j \leq i \\[10pt] 0 & \text{if } j > i \end{cases}

For the denominator, we expand ∑k=1Tα~ik\displaystyle\sum_{k=1}^{T} \tilde{\alpha}_{ik} by separating the visible positions (k≤ik \leq i) from the future ones (k>ik > i), then apply first the definition of α~ik\tilde{\alpha}_{ik} (Step 2) and then that of αik\alpha_{ik} (Step 1):

∑k=1Tα~ik=∑k≤iα~ik⏟= αik+∑k>iα~ik⏟= 0=∑k≤iαik=∑k≤ieeik∑m=1Teeim\sum_{k=1}^{T} \tilde{\alpha}_{ik} = \sum_{k \leq i} \underbrace{\tilde{\alpha}_{ik}}_{=\,\alpha_{ik}} + \sum_{k > i} \underbrace{\tilde{\alpha}_{ik}}_{=\,0} = \sum_{k \leq i} \alpha_{ik} = \sum_{k \leq i} \frac{e^{e_{ik}}}{\displaystyle\sum_{m=1}^{T} e^{e_{im}}}

The term ∑m=1Teeim\displaystyle\sum_{m=1}^{T} e^{e_{im}} is constant with respect to kk, so we can factor it out:

∑k=1Tα~ik=∑k≤ieeik∑m=1Teeim\sum_{k=1}^{T} \tilde{\alpha}_{ik} = \frac{\displaystyle\sum_{k \leq i} e^{e_{ik}}}{\displaystyle\sum_{m=1}^{T} e^{e_{im}}}

We substitute the numerator and denominator into the expression of αij(1)\alpha_{ij}^{(1)} for j≤ij \leq i:

αij(1)=α~ij∑k=1Tα~ik=eeij∑m=1Teeim∑k≤ieeik∑m=1Teeim\alpha_{ij}^{(1)} = \frac{\tilde{\alpha}_{ij}}{\displaystyle\sum_{k=1}^{T} \tilde{\alpha}_{ik}} = \frac{\quad\dfrac{e^{e_{ij}}}{\displaystyle\sum_{m=1}^{T} e^{e_{im}}}\quad}{\dfrac{\displaystyle\sum_{k \leq i} e^{e_{ik}}}{\displaystyle\sum_{m=1}^{T} e^{e_{im}}}}

The factor ∑m=1Teeim\displaystyle\sum_{m=1}^{T} e^{e_{im}} appears in the numerator and in the denominator, so it cancels out exactly:

αij(1)=eeij∑m=1Teeim×∑m=1Teeim∑k≤ieeik\alpha_{ij}^{(1)} = \frac{e^{e_{ij}}}{\displaystyle\sum_{m=1}^{T} e^{e_{im}}} \times \frac{\displaystyle\sum_{m=1}^{T} e^{e_{im}}}{\displaystyle\sum_{k \leq i} e^{e_{ik}}}αij(1)=eeij∑k≤ieeikfor j≤i,αij(1)=0for j>i\boxed{\alpha_{ij}^{(1)} = \frac{e^{e_{ij}}}{\displaystyle\sum_{k \leq i} e^{e_{ik}}} \quad \text{for } j \leq i, \qquad \alpha_{ij}^{(1)} = 0 \quad \text{for } j > i}

Method 2. Masking the raw scores before the softmax

Step 1. Masking the raw scores.

We replace the scores of future tokens with −∞-\infty before any computation:

e~ij={eijif j≤i−∞if j>i\tilde{e}_{ij} = \begin{cases} e_{ij} & \text{if } j \leq i \\ -\infty & \text{if } j > i \end{cases}

Step 2. A single softmax over the masked scores.

αij(2)=e e~ij∑k=1Te e~ik\alpha_{ij}^{(2)} = \frac{e^{\,\tilde{e}_{ij}}}{\displaystyle\sum_{k=1}^{T} e^{\,\tilde{e}_{ik}}}

Explicit expansion. We split the denominator into visible and future positions, then apply the definition of e~ik\tilde{e}_{ik}:

∑k=1Te e~ik=∑k≤ie e~ik⏟= eeik+∑k>ie e~ik⏟= e−∞ = 0=∑k≤ieeik\sum_{k=1}^{T} e^{\,\tilde{e}_{ik}} = \sum_{k \leq i} \underbrace{e^{\,\tilde{e}_{ik}}}_{=\,e^{e_{ik}}} + \sum_{k > i} \underbrace{e^{\,\tilde{e}_{ik}}}_{=\,e^{-\infty}\,=\,0} = \sum_{k \leq i} e^{e_{ik}}

Likewise for the numerator, applying the definition of e~ij\tilde{e}_{ij}:

e e~ij={eeijif j≤ie−∞=0if j>ie^{\,\tilde{e}_{ij}} = \begin{cases} e^{e_{ij}} & \text{if } j \leq i \\ e^{-\infty} = 0 & \text{if } j > i \end{cases}

We therefore get:

αij(2)=eeij∑k≤ieeikfor j≤i,αij(2)=0for j>i\boxed{\alpha_{ij}^{(2)} = \frac{e^{e_{ij}}}{\displaystyle\sum_{k \leq i} e^{e_{ik}}} \quad \text{for } j \leq i, \qquad \alpha_{ij}^{(2)} = 0 \quad \text{for } j > i}

Equivalence of the two methods

Comparing the two boxed results term by term is immediate:

For j≤i:αij(1)=eeij∑k≤ieeik=αij(2)\text{For } j \leq i: \quad \alpha_{ij}^{(1)} = \frac{e^{e_{ij}}}{\displaystyle\sum_{k \leq i} e^{e_{ik}}} = \alpha_{ij}^{(2)}For j>i:αij(1)=0=αij(2)\text{For } j > i: \quad \alpha_{ij}^{(1)} = 0 = \alpha_{ij}^{(2)}

We conclude:

αij(1)=αij(2)∀ j∈{1,…,T}\boxed{\alpha_{ij}^{(1)} = \alpha_{ij}^{(2)} \qquad \forall\, j \in \{1, \dots, T\}}

This is an exact algebraic equality. In Method 1, the factor ∑m=1Teeim\displaystyle\sum_{m=1}^{T} e^{e_{im}} introduced by the first softmax cancels out exactly during renormalization. The initial presence of the future tokens in the computation leaves no trace in the final result: the denominator only contains ∑k≤ieeik\displaystyle\sum_{k \leq i} e^{e_{ik}}, exactly as in Method 2, where the future tokens were never included. There is therefore no information leakage.


Efficient approach: mask before softmax with −∞

The naive approach applies softmax and then corrects it. We can do better: mask the attention scores before the softmax by replacing the future positions with −∞-\infty.

The justification is immediate: e−∞=0e^{-\infty} = 0, so softmax automatically assigns zero weight to these positions, and the remaining weights sum to 1, with no manual renormalization.

mask   = torch.triu(torch.ones(context_length, context_length), diagonal=1)
masked = attn_scores.masked_fill(mask.bool(), -torch.inf)

torch.triu with diagonal=1 isolates the strictly upper triangle (excluding the diagonal), that is, the future positions. masked_fill replaces these positions with −∞-\infty.

tensor([[0.2899,   -inf,   -inf,   -inf,   -inf,   -inf],
        [0.4656, 0.1723,   -inf,   -inf,   -inf,   -inf],
        [0.4594, 0.1703, 0.1731,   -inf,   -inf,   -inf],
        [0.2642, 0.1024, 0.1036, 0.0186,   -inf,   -inf],
        [0.2183, 0.0874, 0.0882, 0.0177, 0.0786,   -inf],
        [0.3408, 0.1270, 0.1290, 0.0198, 0.1290, 0.0078]],
       grad_fn=<MaskedFillBackward0>)

All that remains is to apply softmax:

attn_weights = torch.softmax(masked / keys.shape[-1]**0.5, dim=-1)
tensor([[1.0000, 0.0000, 0.0000, 0.0000, 0.0000, 0.0000],
        [0.5517, 0.4483, 0.0000, 0.0000, 0.0000, 0.0000],
        [0.3800, 0.3097, 0.3103, 0.0000, 0.0000, 0.0000],
        [0.2758, 0.2460, 0.2462, 0.2319, 0.0000, 0.0000],
        [0.2175, 0.1983, 0.1984, 0.1888, 0.1971, 0.0000],
        [0.1935, 0.1663, 0.1666, 0.1542, 0.1666, 0.1529]],
       grad_fn=<SoftmaxBackward0>)

The result is identical to the naive approach, in two steps instead of four. This is the approach used in practice.


Masking Additional Attention Weights with Dropout

Dropout is a regularization technique proposed by Geoffrey Hinton in 2012. At each training step, each neuron has a probability pp (the dropout rate) of being temporarily switched off: ignored for this pass, but possibly active in the next one. After training, no neuron is ever switched off.

The intuition behind this technique may seem paradoxical. Imagine a company where every employee flips a coin each morning to decide whether to come to work. The company would be forced to adapt: it could no longer rely on a single person for a critical task, all expertise would have to be spread out, and employees would learn to work with many colleagues rather than a fixed handful. It would become far more robust. This is exactly what happens in a neural network: neurons trained under dropout cannot co-adapt with their usual neighbors; each one becomes more useful on its own and less sensitive to small variations in the input.

Another way to look at it: at each training step, the active network is different (one of 2N2^N possible configurations for NN neurons). The final network can be seen as an average ensemble of all these subnetworks.

In the attention mechanism, dropout is applied directly to the attention weights, which is the most common variant in practice. Concretely, the causal mask and the dropout mask are layered on top of each other: the first zeroes out the upper triangle (future tokens), the second randomly zeroes out additional positions among the visible tokens.

Figure 3.22
Figure 3.22: Starting from the causal attention matrix (zeroed upper triangle), a random dropout mask is applied to zero out additional weights among the visible positions, reducing the risk of overfitting during training. Source: Raschka (2024).

Automatic rescaling

With a dropout rate pp, the surviving weights are multiplied by 11−p\frac{1}{1-p}. This correction is necessary: during training, a neuron is connected on average to only a fraction (1−p)(1-p) of its usual inputs. Without this compensation, at inference time (when all neurons are active) the neuron would receive a signal with a very different magnitude from what it learned.

We check this behavior on a matrix of 1s with p=0.5p = 0.5:

torch.manual_seed(123)
dropout = torch.nn.Dropout(0.5)
example = torch.ones(6, 6)
print(dropout(example))
tensor([[2., 2., 0., 2., 2., 0.],
        [0., 0., 0., 2., 0., 2.],
        [2., 2., 2., 2., 0., 2.],
        [0., 2., 2., 0., 0., 2.],
        [0., 2., 0., 2., 0., 2.],
        [0., 2., 2., 2., 2., 0.]])

The elements that remain nonzero after dropout equal 11−0.5=2\frac{1}{1 - 0.5} = 2, which confirms the rescaling.


Applying it to the causal attention weights

torch.manual_seed(123)
print(dropout(attn_weights))
tensor([[2.0000, 0.0000, 0.0000, 0.0000, 0.0000, 0.0000],
        [0.0000, 0.0000, 0.0000, 0.0000, 0.0000, 0.0000],
        [0.7599, 0.6194, 0.6206, 0.0000, 0.0000, 0.0000],
        [0.0000, 0.4921, 0.4925, 0.0000, 0.0000, 0.0000],
        [0.0000, 0.3966, 0.0000, 0.3775, 0.0000, 0.0000],
        [0.0000, 0.3327, 0.3331, 0.3084, 0.3331, 0.0000]],
       grad_fn=<MulBackward0>)

The causal mask (zeroed upper triangle) is preserved. Dropout is layered on top of it, randomly zeroing out additional weights among the visible positions and rescaling the survivors.

Note: dropout results can vary depending on the operating system, regardless of manual_seed. This behavior is documented in the PyTorch issue tracker.

Implementing a Compact Causal Attention Class

We now integrate the causal mask and dropout into a CausalAttention class that replaces SelfAttention. This class must also handle batched inputs, that is, several sequences processed in parallel.

To simulate a batch, we duplicate the input:

batch = torch.stack((inputs, inputs), dim=0)
print(batch.shape)
# torch.Size([2, 6, 3])

A 3D tensor of shape (batch_size, num_tokens, d_in): 2 sequences, 6 tokens each, 3-dimensional embeddings.


The code

class CausalAttention(nn.Module):

    def __init__(self, d_in, d_out, context_length, dropout, qkv_bias=False):
        super().__init__()
        self.d_out = d_out
        self.W_query = nn.Linear(d_in, d_out, bias=qkv_bias)
        self.W_key   = nn.Linear(d_in, d_out, bias=qkv_bias)
        self.W_value = nn.Linear(d_in, d_out, bias=qkv_bias)
        self.dropout = nn.Dropout(dropout)
        self.register_buffer(
            'mask',
            torch.triu(torch.ones(context_length, context_length), diagonal=1)
        )

    def forward(self, x):
        b, num_tokens, d_in = x.shape
        keys    = self.W_key(x)
        queries = self.W_query(x)
        values  = self.W_value(x)

        attn_scores = queries @ keys.transpose(1, 2)
        attn_scores.masked_fill_(
            self.mask.bool()[:num_tokens, :num_tokens], -torch.inf
        )
        attn_weights = torch.softmax(attn_scores / keys.shape[-1]**0.5, dim=-1)
        attn_weights = self.dropout(attn_weights)

        context_vec = attn_weights @ values
        return context_vec
torch.manual_seed(123)
context_length = batch.shape[1]
ca = CausalAttention(d_in, d_out, context_length, 0.0)
context_vecs = ca(batch)
print("context_vecs.shape:", context_vecs.shape)
# context_vecs.shape: torch.Size([2, 6, 2])

Point 1. register_buffer: what is it and why?

In a PyTorch nn.Module, there are two kinds of tensors:

Parameters (nn.Parameter): the learned weights, registered through nn.Linear. PyTorch tracks them automatically: they appear in model.parameters(), receive gradients, and are moved to the GPU with model.to(device).

Buffers (register_buffer): fixed, non-learned tensors that the model needs but that must not be updated by the optimizer. This is exactly the case of the causal mask: it is constant, computed once at initialization.

Without register_buffer, we could simply write self.mask = torch.triu(...). It would work locally, but with two problems:

The mask would not be moved to the GPU automatically. If we call model.to("cuda"), the weights move, but self.mask does not. At the masked_fill_ call, PyTorch would raise a device mismatch error: the mask is on the CPU and the attention scores are on the GPU.

The mask would not appear in model.state_dict(), the dictionary that captures the model’s full state for saving. A model reloaded from a checkpoint would lose its mask.

register_buffer solves both: the mask follows the model everywhere (CPU/GPU) and is included in state_dict().


Point 2. Anatomy of the masking code

register_buffer('mask', torch.triu(...))

The call register_buffer('mask', tensor) does two things: it registers the tensor as a buffer of the module, and it makes it accessible as self.mask. The first argument, 'mask', is simply the name under which the buffer is registered: that name becomes the attribute. We could have written register_buffer('causal_mask', ...) and then accessed it as self.causal_mask.

The registered tensor is:

torch.triu(torch.ones(context_length, context_length), diagonal=1)

torch.triu extracts the upper triangle of a matrix. With diagonal=1, the main diagonal is excluded: only the positions strictly above it are 1:

Mask with diagonal=1:
 tensor([[0., 1., 1., 1., 1., 1.],
        [0., 0., 1., 1., 1., 1.],
        [0., 0., 0., 1., 1., 1.],
        [0., 0., 0., 0., 1., 1.],
        [0., 0., 0., 0., 0., 1.],
        [0., 0., 0., 0., 0., 0.]])

The 1s mark exactly the future positions to mask.

attn_scores.masked_fill_(self.mask.bool()[:num_tokens, :num_tokens], -torch.inf)

Let’s go through it argument by argument.

self.mask.bool() converts the mask from float (0.0 / 1.0) to bool (False / True). masked_fill_ expects a boolean mask: the positions set to True will be replaced.

[:num_tokens, :num_tokens]: the mask was created at initialization with the maximum size context_length × context_length. The current sequence may be shorter. This slicing extracts the submask with the exact size of the sequence being processed, without recreating the tensor on every call.

-torch.inf is the replacement value for the future positions. Since e−∞=0e^{-\infty} = 0, softmax will give them zero weight.

masked_fill_ (with an underscore) is an in-place operation: attn_scores is modified directly in memory with no intermediate copy, which saves memory for potentially large attention matrices.

In short: for each position (i,j)(i, j) where self.mask[i,j] is True (that is, j>ij > i, a future token), the corresponding value in attn_scores is replaced with −∞-\infty.


Point 3. keys.transpose(1, 2) instead of keys.T

In SelfAttention_v1 and SelfAttention_v2, we had keys.T because the inputs were 2D, (num_tokens, d_out). Here keys is 3D: (b, num_tokens, d_out). .T would reverse all the axes, giving (d_out, num_tokens, b), which is wrong. .transpose(1, 2) only swaps the last two dimensions, giving (b, d_out, num_tokens): the batch dimension stays in position 0, and the matrix product works correctly on each sequence of the batch independently.


Figure 3.23
Figure 3.23: How the implementation has progressed, from simplified attention to causal attention with trainable weights, a causal mask and dropout. The next step is multi-head attention. Source: Raschka (2024).

Extending Single-Head Attention to Multi-Head Attention

Multi-head attention runs the attention mechanism several times in parallel, each with its own weight matrices WQW_Q, WKW_K, WVW_V. Each head learns to specialize in a different kind of relationship within the sequence. A single CausalAttention module corresponds to single-head attention: one set of weights processing the input.


Stacking Several Single-Head Attention Layers

The idea

The most direct way to implement multi-head attention is to create several independent CausalAttention modules, run them on the same input, and concatenate their outputs.

Figure 3.24
Figure 3.24: A multi-head attention module with two heads. Each head has its own weight matrices WQW_Q, WKW_K, WVW_V and produces its own context vectors Z1Z_1 and Z2Z_2, which are then concatenated into ZZ. Source: Raschka (2024).

Implementation

class MultiHeadAttentionWrapper(nn.Module):

    def __init__(self, d_in, d_out, context_length,
                 dropout, num_heads, qkv_bias=False):
        super().__init__()
        self.heads = nn.ModuleList(
            [CausalAttention(d_in, d_out, context_length, dropout, qkv_bias)
             for _ in range(num_heads)]
        )

    def forward(self, x):
        return torch.cat([head(x) for head in self.heads], dim=-1)

Two points deserve attention.

nn.ModuleList is a list of CausalAttention modules whose parameters are managed centrally by MultiHeadAttentionWrapper: they are accessible through model.parameters(), moved with model.to(device), and saved in state_dict().

torch.cat(..., dim=-1) concatenates the outputs of each head along the last dimension, that is, the embedding dimension. If each head produces vectors of dimension d_out, the final output has dimension d_out × num_heads.

Example

torch.manual_seed(123)
context_length = batch.shape[1]
mha = MultiHeadAttentionWrapper(d_in, d_out, context_length, 0.0, num_heads=2)
context_vecs = mha(batch)
print(context_vecs)
print("context_vecs.shape:", context_vecs.shape)
tensor([[[-0.4519,  0.2216,  0.4772,  0.1063],
         [-0.5874,  0.0058,  0.5891,  0.3257],
         [-0.6300, -0.0632,  0.6202,  0.3860],
         [-0.5675, -0.0843,  0.5478,  0.3589],
         [-0.5526, -0.0981,  0.5321,  0.3428],
         [-0.5299, -0.1081,  0.5077,  0.3493]],

        [[-0.4519,  0.2216,  0.4772,  0.1063],
         [-0.5874,  0.0058,  0.5891,  0.3257],
         [-0.6300, -0.0632,  0.6202,  0.3860],
         [-0.5675, -0.0843,  0.5478,  0.3589],
         [-0.5526, -0.0981,  0.5321,  0.3428],
         [-0.5299, -0.1081,  0.5077,  0.3493]]], grad_fn=<CatBackward0>)
context_vecs.shape torch.Size([2, 6, 4])

The shape (2, 6, 4) reads as: 2 sequences in the batch, 6 tokens each, context vectors of dimension 2×2=42 \times 2 = 4. The two sequences are identical (the batch was simulated by duplication), hence the identical context vectors.

Figure 3.25
Figure 3.25: With num_heads=2 and d_out=2, each head produces a matrix of 2-dimensional context vectors. The two matrices are concatenated along the column dimension, giving a final dimension of 2×2=42 \times 2 = 4. Source: Raschka (2024).

Exercise 3.2: Reducing the dimension of the context vectors

Change the arguments of MultiHeadAttentionWrapper(..., num_heads=2) so that the output context vectors have dimension 2 instead of 4, while keeping num_heads=2.

Solution:

The output dimension is d_out × num_heads. So we simply pass d_out=1: each head produces 1-dimensional vectors, and their concatenation gives dimension 2.

torch.manual_seed(123)
mha = MultiHeadAttentionWrapper(d_in, 1, context_length, 0.0, num_heads=2)
print(mha(batch).shape)
torch.Size([2, 6, 2])

Limitation of this approach

The heads are processed sequentially in forward through [head(x) for head in self.heads]. This is functionally correct but inefficient: we make num_heads separate passes where a single matrix operation could compute everything in parallel. The next section presents an implementation that takes advantage of this parallelism.


Implementing Multi-Head Attention with Weight Splits

The MultiHeadAttention class merges MultiHeadAttentionWrapper and CausalAttention into a single unit. The guiding idea: perform a single linear projection of dimension doutd_{out}, then implicitly split the result into H=num_headsH = \texttt{num\_heads} subspaces of dimension dh=dout/Hd_h = d_{out}/H. Unlike MultiHeadAttentionWrapper, which kept HH weight matrices WK,j∈Rdh×dinW_{K,j} \in \mathbb{R}^{d_h \times d_{in}} and repeated the matrix multiplication for each head, a single multiplication is enough here.

mha_explained_1
Overview of the multi-head attention mechanism (source: CNRS-FIDLE).

Below, we describe the successive transformations applied to the key tensor KK (the same ones apply identically to QQ and VV).


Step 1. Single projection: (b, num_tokens, d_in) → (b, num_tokens, d_out)

keys = self.W_key(x)   # (b, num_tokens, d_out)

Each row of the output tensor is the key vector ki=WK⊤xi∈Rdout\mathbf{k}_i = W_K^\top x_i \in \mathbb{R}^{d_{out}} of token ii.


Step 2. Splitting by head with .view(): (b, num_tokens, d_out) → (b, num_tokens, H, d_h)

keys = keys.view(b, num_tokens, self.num_heads, self.head_dim)
What .view() actually does

To understand .view(), you first need to understand how PyTorch stores data.

Memory is a hallway of numbered slots. When PyTorch creates a tensor, it reserves a block of consecutive slots in memory (RAM), one slot per number. For example, a 2×32 \times 3 matrix

M=(abcdef)M = \begin{pmatrix} a & b & c \\ d & e & f \end{pmatrix}

is stored in memory as a simple row of 6 slots:

a  b  c  d  e  f\boxed{a}\;\boxed{b}\;\boxed{c}\;\boxed{d}\;\boxed{e}\;\boxed{f}

PyTorch just remembers two pieces of information separately: (1) the address of the first slot, and (2) the shape (2, 3), which lets it compute where to find each element: the element in row ii, column jj is in slot 3i+j3i + j.

.view() changes the reading grid, not where the data physically sits. Calling .view(3, 2) on the same memory amounts to deciding to read it as

M′=(abcdef)M' = \begin{pmatrix} a & b \\ c & d \\ e & f \end{pmatrix}

PyTorch simply updates the shape from (2, 3) to (3, 2) and recomputes the access formula: element (i,j)(i, j) is now in slot 2i+j2i + j. That’s all.

This is why .view() is called copy-free: it allocates no new memory and moves no numbers. It costs almost nothing in time or memory, whatever the size of the tensor.

Applying it to our keys

Notation. For a token ii and a head jj, let ki,j∈Rdh\mathbf{k}_{i,j} \in \mathbb{R}^{d_h} be the associated key subvector.

After step 1, for a fixed token ii, the vector ki∈Rdout\mathbf{k}_i \in \mathbb{R}^{d_{out}} holds the projections of all the heads laid end to end:

K[ b, i, :]=ki=[ki,1,1,…,ki,1,dh⏟ki,1,  ki,2,1,…,ki,2,dh⏟ki,2,  …,  ki,H,1,…,ki,H,dh⏟ki,H]K[\,b,\,i,\,:] = \mathbf{k}_i = \bigl[ \underbrace{k_{i,1,1},\ldots,k_{i,1,d_h}}_{\displaystyle\mathbf{k}_{i,1}},\; \underbrace{k_{i,2,1},\ldots,k_{i,2,d_h}}_{\displaystyle\mathbf{k}_{i,2}},\; \ldots,\; \underbrace{k_{i,H,1},\ldots,k_{i,H,d_h}}_{\displaystyle\mathbf{k}_{i,H}} \bigr]

.view(b, num_tokens, H, d_h) tells PyTorch: “reinterpret the dimension dout=H⋅dhd_{out} = H \cdot d_h as two nested dimensions”. Without moving anything in memory, the same token ii is now accessible as an (H,dh)(H, d_h) block:

K[ b, i, :, :]=(ki,1ki,2⋮ki,H)∈RH×dhK[\,b,\,i,\,:,\,:] = \begin{pmatrix} \mathbf{k}_{i,1} \\ \mathbf{k}_{i,2} \\ \vdots \\ \mathbf{k}_{i,H} \end{pmatrix} \in \mathbb{R}^{H \times d_h}

The jj-th row of this block is exactly ki,j\mathbf{k}_{i,j}, the key vector of token ii for head jj.

view_operation_1
.view() reinterprets the doutd_{out} dimension as (H,dh)(H, d_h). For each token ii, the vector ki∈Rdout\mathbf{k}_i \in \mathbb{R}^{d_{out}} becomes a block K[b,i,:,:]∈RH×dhK[b, i, :, :] \in \mathbb{R}^{H \times d_h} whose jj-th row is ki,j\mathbf{k}_{i,j}. No data is moved in memory.

Step 3. Reorganizing by head with .transpose(1, 2): (b, num_tokens, H, d_h) → (b, H, num_tokens, d_h)

keys = keys.transpose(1, 2)

Notation. Let

Kj=(k1,j⋮kT,j)∈RT×dh(T=num_tokens)K_j = \begin{pmatrix} \mathbf{k}_{1,j} \\ \vdots \\ \mathbf{k}_{T,j} \end{pmatrix} \in \mathbb{R}^{T \times d_h} \qquad (T = \texttt{num\_tokens})

be the key matrix of head jj over the whole sequence, that is, all the projections into the same subspace, one row per token.

What .transpose() actually does, and why it differs from .view()

Like .view(), .transpose() moves nothing in memory: it just changes the access formula. But it creates a problem that .view() did not.

Let’s go back to the 2×32 \times 3 matrix:

M=(123456)M = \begin{pmatrix} 1 & 2 & 3 \\ 4 & 5 & 6 \end{pmatrix}

Physical memory: 1  2  3  4  5  6\boxed{1}\;\boxed{2}\;\boxed{3}\;\boxed{4}\;\boxed{5}\;\boxed{6}

With .view(3, 2), PyTorch reread these same slots as “3 rows of 2 columns”, and each consecutive row really was side by side in memory. No problem.

That is:

M.view(3,2)⇒(123456)M.view(3, 2) \Rightarrow \begin{pmatrix} 1 & 2 \\ 3 & 4 \\ 5 & 6 \end{pmatrix}

With .transpose(), we logically get:

M⊤=(142536)M^\top = \begin{pmatrix} 1 & 4 \\ 2 & 5 \\ 3 & 6 \end{pmatrix}

To read the first row of M⊤M^\top (the elements 11 and 44), PyTorch has to skip slots: it takes slot 0 (1), then slot 3 (4). In between are slots 1 and 2, which logically belong to other rows. The logical order and the memory order diverge:

Memory order (untouched):1  2  3  4  5  6\text{Memory order (untouched)}: \boxed{1}\;\boxed{2}\;\boxed{3}\;\boxed{4}\;\boxed{5}\;\boxed{6} Logical order (after transposing):142536\text{Logical order (after transposing)}: 1\quad 4\quad 2\quad 5\quad 3\quad 6

A tensor is non-contiguous when reading its elements in logical order requires skipping slots in memory.

The two representations of the key tensor

Left representation: organized by token. In the (b, num_tokens, H, d_h) shape produced by .view(), the main iteration dimension is the token. For a fixed token ii, the block K[ b, i, :, :]K[\,b,\,i,\,:,\,:] gathers the projections of all the heads for that token:

K[ b, i, :, :]=(ki,1⋮ki,H)∈RH×dhK[\,b,\,i,\,:,\,:] = \begin{pmatrix} \mathbf{k}_{i,1} \\ \vdots \\ \mathbf{k}_{i,H} \end{pmatrix} \in \mathbb{R}^{H \times d_h}

We read the tensor head by head for a given token. In memory, the data is laid out in this order: all the heads of token 1, then all the heads of token 2, and so on. The matrices K1,…,KHK_1, \ldots, K_H exist conceptually, but their rows are interleaved: k1,1\mathbf{k}_{1,1} followed by k1,2\mathbf{k}_{1,2}, then k2,1\mathbf{k}_{2,1} followed by k2,2\mathbf{k}_{2,2}, and so on. To rebuild KjK_j, you would have to skip slots on every row.

Right representation: organized by head. After .transpose(1, 2), the shape is (b, H, num_tokens, d_h). For a fixed head jj, the block K[ b, j, :, :]=KjK[\,b,\,j,\,:,\,:] = K_j gathers the projections of all the tokens for that head:

K[ b, j, :, :]=Kj=(k1,j⋮kT,j)∈RT×dhK[\,b,\,j,\,:,\,:] = K_j = \begin{pmatrix} \mathbf{k}_{1,j} \\ \vdots \\ \mathbf{k}_{T,j} \end{pmatrix} \in \mathbb{R}^{T \times d_h}

We read the tensor token by token for a given head. Each KjK_j forms an independent logical block. But as seen above, .transpose() has not moved the data: the logical order and the memory order still diverge. The tensor is non-contiguous.

Transposing does not change the data, only the way it is traversed: we switch from reading by token to reading by head.

view_operation_2
.transpose(1, 2) swaps the num_tokens and num_heads dimensions. On the left, the block K[b,i,:,:]K[b, i, :, :] gathers all the heads of a given token. On the right, the block K[b,j,:,:]=KjK[b, j, :, :] = K_j gathers all the tokens of a given head. Physical memory does not move: only the reading order changes.

Step 4. Attention scores through batched multiplication

Since QQ, KK, VV all have shape (b,H,T,dh)(b, H, T, d_h), the product

attn_scores = queries @ keys.transpose(2, 3)

computes QjKj⊤Q_j K_j^\top for all the heads and all the batch elements at once:

attn_scores[ b, j, :, :]=QjKj⊤∈RT×T\texttt{attn\_scores}[\,b,\,j,\,:,\,:] = Q_j K_j^\top \in \mathbb{R}^{T \times T}
mha_explained_2
Computing the attention scores through batched matrix multiplication (source: CNRS-FIDLE).

Step 5. Recombining the heads

After causal masking, softmax and dropout, the context vectors are recombined:

context_vec = (attn_weights @ values).transpose(1, 2)
# (b, H, T, d_h) → (b, T, H, d_h)
context_vec = context_vec.contiguous().view(b, num_tokens, self.d_out)
# (b, T, H, d_h) → (b, T, d_out)
What .contiguous() actually does

.transpose() left the tensor in a state where the logical order and the memory order diverge. .view() cannot handle this gap: it always assumes that reading the elements in logical order means reading them slot after slot in memory. Calling .view() directly would therefore raise an error.

.contiguous() solves this: it allocates a new memory block and copies the data into it in the order that matches the current logical shape.

With our numerical example: after .transpose(), the tensor is logically M⊤M^\top but the memory still holds 1,2,3,4,5,61, 2, 3, 4, 5, 6:

Memory order:1  2  3  4  5  6\text{Memory order}: \boxed{1}\;\boxed{2}\;\boxed{3}\;\boxed{4}\;\boxed{5}\;\boxed{6} Logical order:142536\text{Logical order}: 1\quad 4\quad 2\quad 5\quad 3\quad 6

.contiguous() creates a new block where the two orders match:

New memory:1  4  2  5  3  6\text{New memory}: \boxed{1}\;\boxed{4}\;\boxed{2}\;\boxed{5}\;\boxed{3}\;\boxed{6}

.view() can then be applied safely.

In short. .view() and .transpose() are both copy-free: they only change how PyTorch reads memory. But .transpose() creates a gap between logical order and memory order that .view() cannot handle. .contiguous() is the explicit copy that reconciles them.

The .view(b, num_tokens, self.d_out) that follows concatenates the outputs of the HH heads for each token:

ci=ci,1∥ci,2∥⋯∥ci,H∈Rdout\mathbf{c}_i = \mathbf{c}_{i,1} \| \mathbf{c}_{i,2} \| \cdots \| \mathbf{c}_{i,H} \in \mathbb{R}^{d_{out}}

Finally, a linear projection mixes the information from the different heads:

context_vec = self.out_proj(context_vec)
mha_explained_3
Recombining the heads’ outputs and the final projection WOW_O (source: CNRS-FIDLE).

The complete code

The five steps, brought together in a single class:

class MultiHeadAttention(nn.Module):
    def __init__(self, d_in, d_out, context_length, dropout, num_heads, qkv_bias=False):
        super().__init__()
        assert d_out % num_heads == 0, "d_out must be divisible by num_heads"

        self.d_out = d_out
        self.num_heads = num_heads
        self.head_dim = d_out // num_heads  # dimension of each head

        self.W_query = nn.Linear(d_in, d_out, bias=qkv_bias)
        self.W_key   = nn.Linear(d_in, d_out, bias=qkv_bias)
        self.W_value = nn.Linear(d_in, d_out, bias=qkv_bias)
        self.out_proj = nn.Linear(d_out, d_out)  # final projection that mixes the heads
        self.dropout = nn.Dropout(dropout)
        self.register_buffer(
            "mask",
            torch.triu(torch.ones(context_length, context_length), diagonal=1)
        )

    def forward(self, x):
        b, num_tokens, d_in = x.shape

        # Step 1: a single projection for all the heads
        keys    = self.W_key(x)
        queries = self.W_query(x)
        values  = self.W_value(x)

        # Step 2: split by head
        keys    = keys.view(b, num_tokens, self.num_heads, self.head_dim)
        queries = queries.view(b, num_tokens, self.num_heads, self.head_dim)
        values  = values.view(b, num_tokens, self.num_heads, self.head_dim)

        # Step 3: (b, H, T, d_h), one matrix per head
        keys    = keys.transpose(1, 2)
        queries = queries.transpose(1, 2)
        values  = values.transpose(1, 2)

        # Step 4: scores, causal mask, softmax and dropout, for all heads at once
        attn_scores = queries @ keys.transpose(2, 3)
        mask_bool = self.mask.bool()[:num_tokens, :num_tokens]
        attn_scores.masked_fill_(mask_bool, -torch.inf)
        attn_weights = torch.softmax(attn_scores / keys.shape[-1]**0.5, dim=-1)
        attn_weights = self.dropout(attn_weights)

        # Step 5: recombine the heads, then the final projection
        context_vec = (attn_weights @ values).transpose(1, 2)
        context_vec = context_vec.contiguous().view(b, num_tokens, self.d_out)
        return self.out_proj(context_vec)
torch.manual_seed(123)
batch_size, context_length, d_in = batch.shape
mha = MultiHeadAttention(d_in, d_out=2, context_length=context_length, dropout=0.0, num_heads=2)
context_vecs = mha(batch)
print("context_vecs.shape:", context_vecs.shape)
context_vecs.shape: torch.Size([2, 6, 2])

Unlike MultiHeadAttentionWrapper, where d_out was the dimension of each head, d_out here is the dimension of the total output: with d_out=2 and num_heads=2, each head works in dimension 1. This is the convention GPT models use.

Exercise 3.3: Initializing a GPT-2-sized attention module

Using the MultiHeadAttention class, initialize a multi-head attention module with the same number of heads as the smallest GPT-2 model (12 heads). Also use GPT-2’s input and output embedding dimensions (768). The smallest GPT-2 supports a context length of 1,024 tokens.

Solution:

# Parameters of the smallest GPT-2
d_in = 768              # input embedding dimension
d_out = 768             # output dimension (same as d_in in GPT-2)
num_heads = 12          # number of attention heads
context_length = 1024   # maximum context length

mha = MultiHeadAttention(
    d_in=d_in,
    d_out=d_out,
    context_length=context_length,
    dropout=0.1,
    num_heads=num_heads,
    qkv_bias=False
)

print("Dimension per head:", mha.head_dim)
print("Number of parameters:", sum(p.numel() for p in mha.parameters()))
Dimension per head: 64
Number of parameters: 2360064

Each head works in dimension 768/12=64768 / 12 = 64. The three projections WQW_Q, WKW_K, WVW_V hold 3×7682=1,769,4723 \times 768^2 = 1{,}769{,}472 weights, and the output projection 7682+768=590,592768^2 + 768 = 590{,}592 (weights and bias), for a total of 2,360,064 parameters.


Summary

The key points of this chapter on attention and multi-head attention:

Attention fundamentals

  • Attention mechanisms turn a set of inputs into enriched context representations that incorporate information from every token.
  • Simple attention is a weighted sum: each output element is a linear combination of the inputs, with weights learned by the model.
  • The attention weights are computed through dot products between queries and keys, a compact and efficient formulation.

Dot-product attention

  • We introduce three trainable linear projections: queries (QQ), keys (KK), and values (VV), which let the model learn what kinds of relationships to look for and how to use them.
  • The raw score eij=qi⊤kjdke_{ij} = \frac{q_i^\top k_j}{\sqrt{d_k}} is normalized by the square root of the dimension (scaling) to keep the variance from exploding.
  • Softmax turns these scores into probabilistic weights αij∈[0,1]\alpha_{ij} \in [0, 1] that sum to 1 per row.

Causal attention

  • For language models that read and generate from left to right, we apply a causal mask that prevents each token from accessing future tokens.
  • Two equivalent approaches: (1) apply softmax, then mask and renormalize, or (2) replace the future scores with −∞-\infty before softmax; the second is more efficient in practice.
  • Dropout applied to the attention weights makes the model more robust by forcing it not to over-rely on individual connections.

Multi-head attention

  • A single attention head captures a single way of combining tokens. Several heads running in parallel let the model explore several subspaces and relationships at once.
  • The naive implementation (stacking single-head modules) is conceptually simple but inefficient: the matrix multiplication (the most expensive operation) is repeated HH times.
  • The efficient implementation (the MultiHeadAttention class) performs a single projection of dimension dout=H⋅dhd_{out} = H \cdot d_h, then splits and reorganizes it with .view() and .transpose() without copying the data. The resulting matrix operations (a batched product over all the heads) take advantage of GPU vectorization.

Tensor operations and memory management

  • .view() and .transpose() are copy-free operations: they only change how PyTorch accesses memory.
  • .transpose() creates a divergence between logical order and memory order, which makes it impossible to apply .view() directly afterwards.
  • .contiguous() resolves this gap by reallocating and rewriting the data so that the two orders match, a necessary step before .view().

Real-world scale: GPT-2

  • The smallest GPT-2 (117M parameters) has 12 heads, an embedding dimension of 768 and a context length of 1,024. Each head operates in dimension 64.
  • Modern LLM architectures (GPT-3, LLaMA, etc.) increase the number of heads (up to 96 for the largest GPT-3) and the dimensions (up to 12,288), but the fundamental design stays the same.

Weekly Notes

Every Sunday, I share what I've been learning: papers, ideas, experiments, and questions that stayed with me.

You can unsubscribe at any time with a single click.

0 Likes • 0 Comments

Discussion about this post0

Join the discussion

A secure sign-in link will be sent to your email address.

Loading discussion...