第 3 章 LLMGPT注意力

第 3 章 编码注意力机制

第 3 章 编码注意力机制

本章来源:本文翻译整理自 LLMs-from-scratch 仓库的 ch03/01_main-chapter-code/ch03.ipynb,原书为 Sebastian Raschka《Build a Large Language Model (From Scratch)》。

本章要做什么

本章覆盖注意力机制,也就是 LLM 的引擎。具体包括:长序列建模的问题、用注意力机制捕获数据依赖、用自注意力关注输入的不同部分、实现带可训练权重的自注意力、用因果注意力隐藏未来的词、把单头注意力扩展成多头注意力。

本 notebook 使用的包

from importlib.metadata import version

print("torch version:", version("torch"))

3.1 长序列建模的问题

本节没有代码。

逐词翻译文本是不可行的,因为源语言和目标语言的语法结构存在差异:

在 transformer 模型引入之前,编码器-解码器 RNN 常用于机器翻译任务。在这种设置里,编码器处理来自源语言的 token 序列,用一个隐藏状态(hidden state,一种神经网络里的中间层)生成整个输入序列的压缩表示:

3.2 用注意力机制捕获数据依赖

本节没有代码。

通过网络里的注意力机制,生成文本的解码器部分能够选择性地访问所有输入 token,这意味着在生成某个特定输出 token 时,某些输入 token 比其他 token 更重要:

Transformer 里的自注意力(self-attention)是一种旨在增强输入表示的技术,它让序列里的每个位置都能与同一序列里的每个其它位置互动,并确定彼此的相关性:

3.3 用自注意力关注输入的不同部分

3.3.1 一个没有可训练权重的简单自注意力机制

本节解释一个非常简化的自注意力变体,它不包含任何可训练权重。这纯粹是为了说明,不是 transformer 里使用的那种注意力机制。下一节(3.3.2)会把这种简单注意力机制扩展成真正的自注意力机制。

假设给定一个输入序列 $x^{(1)}$ 到 $x^{(T)}$:

  • 输入是一段文本(比如 "Your journey starts with one step" 这样的句子),已经按第 2 章描述转换成了 token 嵌入。
  • 例如,$x^{(1)}$ 是表示单词 "Your" 的 d 维向量,等等。

目标: 为 $x^{(1)}$ 到 $x^{(T)}$ 里的每个输入序列元素 $x^{(i)}$ 计算上下文向量 $z^{(i)}$(其中 $z$ 和 $x$ 维度相同)。

  • 上下文向量 $z^{(i)}$ 是输入 $x^{(1)}$ 到 $x^{(T)}$ 的加权和。
  • 上下文向量是特定于某个输入的"上下文"。
    • 不用 $x^{(i)}$ 当任意输入 token 的占位符,我们考虑第二个输入 $x^{(2)}$。
    • 继续用一个具体例子,不用占位符 $z^{(i)}$,我们考虑第二个输出上下文向量 $z^{(2)}$。
    • 第二个上下文向量 $z^{(2)}$ 是所有输入 $x^{(1)}$ 到 $x^{(T)}$ 关于第二个输入元素 $x^{(2)}$ 加权的加权和。
    • 注意力权重(attention weights)就是那些决定每个输入元素在计算 $z^{(2)}$ 时对加权和贡献多少的权重。
    • 简而言之,把 $z^{(2)}$ 想成 $x^{(2)}$ 的一个修改版本,它同时还融入了所有与手头任务相关的其它输入元素的信息。

(请注意,这张图里的数字被截断到小数点后一位以减少视觉杂乱;同样,其它图可能也包含截断的值。)

按照惯例,未归一化的注意力权重被称为**"注意力分数"(attention scores),而归一化后总和为 1 的注意力分数被称为"注意力权重"(attention weights)**。

下面的代码一步步走完上面的图。

第 1 步: 计算未归一化的注意力分数 $\omega$。

假设我们用第二个输入 token 作为查询(query),即 $q^{(2)} = x^{(2)}$,我们通过点积计算未归一化的注意力分数:

  • $\omega_{21} = x^{(1)} q^{(2)\top}$
  • $\omega_{22} = x^{(2)} q^{(2)\top}$
  • $\omega_{23} = x^{(3)} q^{(2)\top}$
  • ...
  • $\omega_{2T} = x^{(T)} q^{(2)\top}$

上面,$\omega$ 是希腊字母 "omega",用来象征未归一化的注意力分数。$\omega_{21}$ 里的下标 "21" 表示输入序列元素 2 被用作查询,去和输入序列元素 1 计算。

假设我们有下面的输入句子,已经按第 3 章描述嵌入成 3 维向量(这里我们用非常小的嵌入维度来说明,这样它能排进页面而不换行):

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)
)

(在这本书里,我们遵循常见的机器学习和深度学习惯例:训练样本表示为行,特征值表示为列;在上面这个张量的情况下,每行代表一个词,每列代表一个嵌入维度。)

本节的主要目标是演示如何使用第二个输入序列 $x^{(2)}$ 作为查询来计算上下文向量 $z^{(2)}$。上图描绘了这个过程的初始步骤:通过点积运算计算 $x^{(2)}$ 和所有其它输入元素之间的注意力分数 ω。

我们用输入序列元素 2,$x^{(2)}$,作为例子来计算上下文向量 $z^{(2)}$;稍后在本节中,我们会把它推广到计算所有上下文向量。第一步是计算查询 $x^{(2)}$ 和所有其它输入 token 之间的点积,从而得到未归一化的注意力分数:

query = inputs[1]  # 2nd input token is the query

attn_scores_2 = torch.empty(inputs.shape[0])
for i, x_i in enumerate(inputs):
    attn_scores_2[i] = torch.dot(x_i, query) # dot product (transpose not necessary here since they are 1-dim vectors)

print(attn_scores_2)

补充说明:点积本质上就是把两个向量逐元素相乘并把得到的乘积求和的简写:

res = 0.

for idx, element in enumerate(inputs[0]):
    res += inputs[0][idx] * query[idx]

print(res)
print(torch.dot(inputs[0], query))

第 2 步: 归一化未归一化的注意力分数("omegas",$\omega$),让它们加起来等于 1。这里有一个简单的方法,把未归一化的注意力分数归一化到总和为 1(这是一个惯例,对解释有用,对训练稳定性也很重要):

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())

然而,在实践中,使用 softmax 函数做归一化很常见,也值得推荐,因为它更能处理极端值,训练时有更理想的梯度性质。下面是一个朴素的 softmax 函数实现,它也把向量元素归一化到总和为 1:

def softmax_naive(x):
    return torch.exp(x) / torch.exp(x).sum(dim=0)

attn_weights_2_naive = softmax_naive(attn_scores_2)

print("Attention weights:", attn_weights_2_naive)
print("Sum:", attn_weights_2_naive.sum())

上面这个朴素实现可能因为溢出和下溢问题,在输入值很大或很小时遭遇数值不稳定。因此,实践中推荐用 PyTorch 的 softmax 实现,它为性能做了高度优化:

attn_weights_2 = torch.softmax(attn_scores_2, dim=0)

print("Attention weights:", attn_weights_2)
print("Sum:", attn_weights_2.sum())

第 3 步: 把嵌入的输入 token $x^{(i)}$ 与注意力权重相乘并把得到的向量求和,计算上下文向量 $z^{(2)}$:

query = inputs[1] # 2nd input token is the query

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)

3.3.2 计算所有输入 token 的注意力权重

推广到所有输入序列 token:

上面,我们计算了输入 2 的注意力权重和上下文向量(如下图高亮行所示)。接下来,我们把这个计算推广到计算所有注意力权重和上下文向量。

(请注意,这张图里的数字被截断到小数点后两位以减少视觉杂乱;每行的值加起来应该是 1.0 或 100%;同样,其它图里的数字也被截断。)

在自注意力里,过程从计算注意力分数开始,随后归一化得到总和为 1 的注意力权重。然后,这些注意力权重通过输入的加权求和用来生成上下文向量。

把之前的第 1 步应用到所有成对元素,计算未归一化的注意力分数矩阵:

attn_scores = torch.empty(6, 6)

for i, x_i in enumerate(inputs):
    for j, x_j in enumerate(inputs):
        attn_scores[i, j] = torch.dot(x_i, x_j)

print(attn_scores)

我们可以通过矩阵乘法更高效地达到同样的效果:

attn_scores = inputs @ inputs.T
print(attn_scores)

和之前的第 2 步类似,我们归一化每一行,让每行里的值加起来等于 1:

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

快速验证每行的值确实加起来等于 1:

row_2_sum = sum([0.1385, 0.2379, 0.2333, 0.1240, 0.1082, 0.1581])
print("Row 2 sum:", row_2_sum)

print("All row sums:", attn_weights.sum(dim=-1))

应用之前的第 3 步计算所有上下文向量:

all_context_vecs = attn_weights @ inputs
print(all_context_vecs)

作为健全性检查,之前算出的上下文向量 $z^{(2)} = [0.4419, 0.6515, 0.5683]$ 可以在上面第 2 行找到:

print("Previous 2nd context vector:", context_vec_2)

3.4 实现带可训练权重的自注意力

下面是一个概念框架,说明本节开发的注意力机制如何融入本书和本章的整体脉络与结构:

3.4.1 一步步计算注意力权重

在本节,我们实现原始 transformer 架构、GPT 模型以及大多数其它流行 LLM 里使用的自注意力机制。这种自注意力机制也叫"缩放点积注意力"(scaled dot-product attention)。

整体思路和之前类似:

  • 我们想计算特定于某个输入元素的、作为输入向量加权和的上下文向量。
  • 为此,我们需要注意力权重。

你会看到,和之前介绍的基础注意力机制相比,只有一些细微差别:

  • 最显著的区别是引入了在模型训练期间会更新的权重矩阵。
  • 这些可训练权重矩阵很关键,这样模型(具体说是模型内部的注意力模块)才能学会产生"好的"上下文向量。

一步步实现自注意力机制,我们会从引入三个训练权重矩阵 $W_q$、$W_k$ 和 $W_v$ 开始。这三个矩阵通过矩阵乘法,把嵌入的输入 token $x^{(i)}$ 投影成查询、键和值向量:

  • 查询向量:$q^{(i)} = x^{(i)},W_q $
  • 键向量:$k^{(i)} = x^{(i)},W_k $
  • 值向量:$v^{(i)} = x^{(i)},W_v $

输入 $x$ 和查询向量 $q$ 的嵌入维度可以相同,也可以不同,取决于模型的设计和具体实现。在 GPT 模型里,输入和输出维度通常相同,但为了说明起见、为了更好地跟进计算,这里我们选择不同的输入和输出维度:

x_2 = inputs[1] # second input element
d_in = inputs.shape[1] # the input embedding size, d=3
d_out = 2 # the output embedding size, d=2

下面,我们初始化三个权重矩阵;注意我们设置 requires_grad=False 是为了减少输出的杂乱,便于说明,但如果我们要用这些权重矩阵做模型训练,我们会设置 requires_grad=True 以便在训练期间更新这些矩阵:

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)

接下来我们计算查询、键和值向量:

query_2 = x_2 @ W_query # _2 because it's with respect to the 2nd input element
key_2 = x_2 @ W_key 
value_2 = x_2 @ W_value

print(query_2)

正如我们下面看到的,我们成功地把 6 个输入 token 从 3 维投影到了 2 维嵌入空间:

keys = inputs @ W_key 
values = inputs @ W_value

print("keys.shape:", keys.shape)
print("values.shape:", values.shape)

下一步,第 2 步,我们通过计算查询和每个键向量的点积,计算未归一化的注意力分数:

keys_2 = keys[1] # Python starts index at 0
attn_score_22 = query_2.dot(keys_2)
print(attn_score_22)

因为我们有 6 个输入,所以对给定的查询向量有 6 个注意力分数:

attn_scores_2 = query_2 @ keys.T # All attention scores for given query
print(attn_scores_2)

接下来,在第 3 步,我们用之前用过的 softmax 函数计算注意力权重(归一化后总和为 1 的注意力分数)。和之前的区别是,我们现在通过除以嵌入维度的平方根 $\sqrt{d_k}$(即 d_k**0.5)来缩放注意力分数:

d_k = keys.shape[1]
attn_weights_2 = torch.softmax(attn_scores_2 / d_k**0.5, dim=-1)
print(attn_weights_2)

第 4 步,我们现在为输入查询向量 2 计算上下文向量:

context_vec_2 = attn_weights_2 @ values
print(context_vec_2)

3.4.2 实现一个紧凑的 SelfAttention 类

把所有东西放在一起,我们可以这样实现自注意力机制:

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 # omega
        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))

我们可以用 PyTorch 的 Linear 层精简上面的实现。如果我们禁用偏置单元,Linear 层等价于矩阵乘法。用 nn.Linear 而不是我们手写的 nn.Parameter(torch.rand(...) 方法的另一个巨大优势是,nn.Linear 有更受偏好的权重初始化方案,这能带来更稳定的模型训练:

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))

注意 SelfAttention_v1SelfAttention_v2 给出不同的输出,因为它们对权重矩阵使用了不同的初始权重。

3.5 用因果注意力隐藏未来的词

在因果注意力(causal attention)里,对角线以上的注意力权重被掩码,确保对任何给定输入,LLM 在用它计算上下文向量时无法利用未来的 token:

3.5.1 应用因果注意力掩码

在本节,我们把之前的自注意力机制转换成因果自注意力机制。

因果自注意力确保模型对序列中某个位置的预测只依赖之前位置已知的输出,不依赖未来位置。更简单地说,这确保每个下一个词的预测只应该依赖前面的词。

要做到这一点,对每个给定 token,我们掩掉未来的 token(输入文本里在当前 token 之后的那些):

为了说明和实现因果自注意力,我们用上一节的注意力分数和权重来工作:

# Reuse the query and key weight matrices of the
# SelfAttention_v2 object from the previous section for convenience
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)
print(attn_weights)

掩掉未来注意力权重最简单的方法,是用 PyTorch 的 tril 函数创建一个掩码,主对角线以下(包括对角线本身)的元素设为 1,主对角线以上设为 0:

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

然后,我们可以把注意力权重和这个掩码相乘,把对角线以上的注意力分数清零:

masked_simple = attn_weights*mask_simple
print(masked_simple)

然而,如果掩码在 softmax 之后应用,像上面这样,它会破坏 softmax 创建的概率分布。Softmax 确保所有输出值加起来等于 1。在 softmax 之后掩码需要重新归一化输出让它们再次加起来等于 1,这使过程复杂化,并可能导致意想不到的后果。

为了确保行加起来等于 1,我们可以这样归一化注意力权重:

row_sums = masked_simple.sum(dim=-1, keepdim=True)
masked_simple_norm = masked_simple / row_sums
print(masked_simple_norm)

虽然我们在技术上现在已经完成了因果注意力机制的编码,但让我们简单看一个更高效的方法来达到同样效果。所以,与其把对角线以上的注意力权重清零并重新归一化结果,我们可以在它们进入 softmax 函数之前,用负无穷掩掉对角线以上的未归一化注意力分数:

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

正如我们下面看到的,现在每行的注意力权重正确地重新加起来等于 1:

attn_weights = torch.softmax(masked / keys.shape[-1]**0.5, dim=-1)
print(attn_weights)

3.5.2 用 dropout 掩码额外的注意力权重

此外,我们还应用 dropout 来减少训练期间的过拟合。Dropout 可以在几个地方应用:

  • 比如,在计算注意力权重之后;
  • 或者在注意力权重与值向量相乘之后。

这里,我们在计算注意力权重之后应用 dropout 掩码,因为这更常见。

而且,在这个具体例子里,我们用 50% 的 dropout 率,也就是随机掩掉一半的注意力权重。(当我们之后训练 GPT 模型时,我们会用更低的 dropout 率,比如 0.1 或 0.2。)

如果我们应用 0.5(50%)的 dropout 率,未丢弃的值会按 1/0.5 = 2 的因子相应缩放。这个缩放由公式 1 / (1 - dropout_rate) 计算。

torch.manual_seed(123)
dropout = torch.nn.Dropout(0.5) # dropout rate of 50%
example = torch.ones(6, 6) # create a matrix of ones

print(dropout(example))
torch.manual_seed(123)
print(dropout(attn_weights))

注意,得到的 dropout 输出可能因你的操作系统而不同;你可以在这里的 PyTorch issue tracker 读到更多关于这种不一致的信息。

3.5.3 实现一个紧凑的因果自注意力类

现在,我们准备好实现一个能用的自注意力实现,包括因果和 dropout 掩码。

还有一件事是实现处理包含多个输入的批的代码,这样我们的 CausalAttention 类能支持第 2 章实现的数据加载器产生的批输出。

为简单起见,为了模拟这种批输入,我们复制输入文本示例:

batch = torch.stack((inputs, inputs), dim=0)
print(batch.shape) # 2 inputs with 6 tokens each, and each token has embedding dimension 3
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) # New
        self.register_buffer('mask', torch.triu(torch.ones(context_length, context_length), diagonal=1)) # New

    def forward(self, x):
        b, num_tokens, d_in = x.shape # New batch dimension b
        # For inputs where `num_tokens` exceeds `context_length`, this will result in errors
        # in the mask creation further below.
        # In practice, this is not a problem since the LLM (chapters 4-7) ensures that inputs  
        # do not exceed `context_length` before reaching this forward method. 
        keys = self.W_key(x)
        queries = self.W_query(x)
        values = self.W_value(x)

        attn_scores = queries @ keys.transpose(1, 2) # Changed transpose
        attn_scores.masked_fill_(  # New, _ ops are in-place
            self.mask.bool()[:num_tokens, :num_tokens], -torch.inf)  # `:num_tokens` to account for cases where the number of tokens in the batch is smaller than the supported context_size
        attn_weights = torch.softmax(
            attn_scores / keys.shape[-1]**0.5, dim=-1
        )
        attn_weights = self.dropout(attn_weights) # New

        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)
print("context_vecs.shape:", context_vecs.shape)

注意,dropout 只在训练期间应用,不在推理期间应用。

3.6 把单头注意力扩展成多头注意力

3.6.1 堆叠多个单头注意力层

下面是之前实现的自我注意力的总结(为简单起见,因果和 dropout 掩码没有显示)。这也叫单头注意力(single-head attention):

我们只需堆叠多个单头注意力模块,就能得到一个多头注意力模块:

多头注意力背后的主要想法是,用不同的、学习到的线性投影(并行地)多次运行注意力机制。这允许模型联合关注不同位置的不同表示子空间里的信息。

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)


torch.manual_seed(123)

context_length = batch.shape[1] # This is the number of tokens
d_in, d_out = 3, 2
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)

在上面的实现里,嵌入维度是 4,因为我们把 d_out=2 用作键、查询、值向量以及上下文向量的嵌入维度。因为我们有 2 个注意力头,所以输出嵌入维度是 2*2=4。

3.6.2 用权重拆分实现多头注意力

虽然上面是一个直观、功能完整的多头注意力实现(包装了前面单头注意力 CausalAttention 实现),但我们可以写一个叫 MultiHeadAttention 的独立类来达到同样效果。

我们不拼接单独注意力头的输出,而是创建单独的 W_query、W_key 和 W_value 权重矩阵,然后把它们拆分给每个注意力头:

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 # Reduce the projection dim to match desired output dim

        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)  # Linear layer to combine head outputs
        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
        # As in `CausalAttention`, for inputs where `num_tokens` exceeds `context_length`, 
        # this will result in errors in the mask creation further below. 
        # In practice, this is not a problem since the LLM (chapters 4-7) ensures that inputs  
        # do not exceed `context_length` before reaching this forward method.

        keys = self.W_key(x) # Shape: (b, num_tokens, d_out)
        queries = self.W_query(x)
        values = self.W_value(x)

        # We implicitly split the matrix by adding a `num_heads` dimension
        # Unroll last dim: (b, num_tokens, d_out) -> (b, num_tokens, num_heads, head_dim)
        keys = keys.view(b, num_tokens, self.num_heads, self.head_dim) 
        values = values.view(b, num_tokens, self.num_heads, self.head_dim)
        queries = queries.view(b, num_tokens, self.num_heads, self.head_dim)

        # Transpose: (b, num_tokens, num_heads, head_dim) -> (b, num_heads, num_tokens, head_dim)
        keys = keys.transpose(1, 2)
        queries = queries.transpose(1, 2)
        values = values.transpose(1, 2)

        # Compute scaled dot-product attention (aka self-attention) with a causal mask
        attn_scores = queries @ keys.transpose(2, 3)  # Dot product for each head

        # Original mask truncated to the number of tokens and converted to boolean
        mask_bool = self.mask.bool()[:num_tokens, :num_tokens]

        # Use the mask to fill attention scores
        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)

        # Shape: (b, num_tokens, num_heads, head_dim)
        context_vec = (attn_weights @ values).transpose(1, 2) 
        
        # Combine heads, where self.d_out = self.num_heads * self.head_dim
        context_vec = context_vec.contiguous().view(b, num_tokens, self.d_out)
        context_vec = self.out_proj(context_vec) # optional projection

        return context_vec

torch.manual_seed(123)

batch_size, context_length, d_in = batch.shape
d_out = 2
mha = MultiHeadAttention(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)

注意,上面本质上是 MultiHeadAttentionWrapper 的一个更高效的改写版本。得到的输出看起来有点不同,因为随机权重初始化不同,但两者都是完全能用的实现,可以用在我们接下来几章要实现的 GPT 类里。


关于输出维度的说明

  • 在上面的 MultiHeadAttention 里,我用 d_out=2 来使用和之前 MultiHeadAttentionWrapper 类相同的设置。
  • MultiHeadAttentionWrapper 由于拼接,返回的输出维度是 d_out * num_heads(即 2*2 = 4)。
  • 然而,MultiHeadAttention 类(为了更对用户友好)允许我们直接通过 d_out 控制输出维度;这意味着,如果我们设 d_out = 2,输出维度就是 2,无论有多少个头。
  • 事后看来,正如读者指出的,用 d_out = 4 来用 MultiHeadAttention 可能更直观,这样它产生的输出维度和 d_out = 2MultiHeadAttentionWrapper 相同。

注意,此外我们给上面的 MultiHeadAttention 类加了一个线性投影层(self.out_proj)。这只是一个不改变维度的线性变换。在 LLM 实现里用这种投影层是标准惯例,但它不是严格必需的(最近的研究表明,去掉它不影响建模表现;见本章末尾的延伸阅读部分)。

注意,如果你对上面内容的紧凑高效实现感兴趣,你也可以考虑 PyTorch 里的 torch.nn.MultiheadAttention 类。

由于上面的实现乍一看可能有点复杂,让我们看看执行 attn_scores = queries @ keys.transpose(2, 3) 时会发生什么:

# (b, num_heads, num_tokens, head_dim) = (1, 2, 3, 4)
a = torch.tensor([[[[0.2745, 0.6584, 0.2775, 0.8573],
                    [0.8993, 0.0390, 0.9268, 0.7388],
                    [0.7179, 0.7058, 0.9156, 0.4340]],

                   [[0.0772, 0.3565, 0.1479, 0.5331],
                    [0.4066, 0.2318, 0.4545, 0.9737],
                    [0.4606, 0.5159, 0.4220, 0.5786]]]])

print(a @ a.transpose(2, 3))

在这种情况下,PyTorch 里的矩阵乘法实现会处理 4 维输入张量,让矩阵乘法在最后两个维度(num_tokens, head_dim)之间进行,然后对每个头重复。

比如,下面变成一种更紧凑的方式,为每个头分别计算矩阵乘法:

first_head = a[0, 0, :, :]
first_res = first_head @ first_head.T
print("First head:\n", first_res)

second_head = a[0, 1, :, :]
second_res = second_head @ second_head.T
print("\nSecond head:\n", second_res)

总结与要点

  • 参见 ./multihead-attention.ipynb 代码 notebook,它是数据加载器(第 2 章)加我们在本章实现的多头注意力类的精简版,在接下来的章节里训练 GPT 模型时会用到。
  • 你可以在 ./exercise-solutions.ipynb 找到练习解答。

关键概念

  • 注意力分数(attention scores):未归一化的注意力权重,通过查询和键的点积计算。
  • 注意力权重(attention weights):归一化后总和为 1 的注意力分数(softmax)。
  • 上下文向量(context vector):输入关于某个查询的加权和,融入了所有相关输入的信息。
  • 缩放点积注意力(scaled dot-product attention):GPT 等模型使用的自注意力,分数除以 $\sqrt{d_k}$。
  • 因果注意力(causal attention):掩掉未来 token,保证每个预测只依赖前面的词。
  • dropout:训练时随机掩掉一部分注意力权重,减少过拟合。
  • 多头注意力(multi-head attention):并行运行多个注意力头,用不同线性投影,或拆分权重矩阵。

练习

练习解答见仓库的 ch03/01_main-chapter-code/exercise-solutions.ipynb