理解并实现大语言模型中的自注意力、多头注意力、因果注意力与交叉注意力
本文将为你讲解Transformer架构与GPT-4、Llama等大语言模型(LLM)中所使用的自注意力机制。自注意力及相关机制是大语言模型的核心组成部分,在从事相关模型开发工作时,理解这一主题非常有价值。
不过,本文不会只停留在理论层面讲解自注意力机制,而是会带你用Python和PyTorch从零开始编码实现。在我看来,从零手写实现算法、模型与技术,是最高效的学习方式!
补充说明:本文是《从零理解并编码实现大语言模型的自注意力机制》的更新扩展版,那篇文章发布在我旧博客上,距今刚好差不多一年。我本人非常喜欢写(也喜欢读)这类“从零实现”的文章,因此特意将这篇内容更新优化,发布在Ahead of AI平台。
此外,这篇文章也促使我动笔撰写《从零构建大语言模型》一书,目前该书仍在创作中。下方是梳理全书脉络的思维模型,展示了自注意力机制在整个大语言模型体系中的位置。
《从零构建大语言模型》一书涵盖主题概览
为控制文章篇幅,本文默认你已经对大语言模型有基础了解,也掌握了注意力机制的基本概念。本文的核心目标,是通过Python与PyTorch的代码逐行讲解,帮你彻底理解注意力机制的运行原理。
自注意力简介
自注意力机制自《Attention Is All You Need》这篇Transformer开山论文提出以来,已成为众多顶尖深度学习模型的基石,在自然语言处理(NLP)领域尤其如此。如今自注意力的应用无处不在,理解其工作原理至关重要。

原始Transformer架构图,来源:https://arxiv\.org/abs/1706\.03762
深度学习中“注意力”的概念,最早源于对循环神经网络(RNN)的改进——解决RNN难以处理长序列、长句子的问题。举个例子,把一句话从一种语言翻译成另一种语言时,逐词直译通常行不通,因为它忽略了每种语言独有的复杂语法结构与习语表达,最终会得到不准确、甚至不通顺的译文。
错误的逐词直译(上)与正确翻译(下)对比
为解决这个问题,研究者提出了注意力机制:让模型在每个时间步都能访问序列中的所有元素,核心是做到“有选择地关注”,判断特定上下文中哪些单词最重要。2017年问世的Transformer架构,提出了独立的自注意力机制,彻底摆脱了对循环神经网络的依赖。
(为精简篇幅、聚焦自注意力的技术细节,本文只对背景动机做简要介绍,把重点放在代码实现上。)

摘自《Attention Is All You Need》论文的可视化图:通过注意力权重展示单词“making”对输入中其他单词的依赖与关注程度(颜色深浅与注意力权重的数值大小正相关)。
我们可以把自注意力理解为一种信息增强机制:它通过融入输入的上下文信息,丰富输入嵌入的信息量。换句话说,自注意力机制让模型能够权衡输入序列中不同元素的重要性,并动态调整它们对输出的影响。这在语言处理任务中尤为关键——同一个单词的含义,会随着它在句子、文档中的上下文变化而改变。
需要注意的是,自注意力有很多变体,其中一个主流研究方向是提升自注意力的计算效率。但绝大多数论文仍沿用《Attention Is All You Need》中提出的**缩放点积注意力(scaled-dot product attention)**原始实现;对大多数训练大规模Transformer的公司而言,自注意力本身通常并不是计算瓶颈。
因此,本文将聚焦最经典、实际应用最广泛的原始缩放点积注意力机制(下文统称自注意力)。如果你对其他类型的注意力机制感兴趣,可以参考2020年的《Efficient Transformers: A Survey》、2023年的《A Survey on Efficient Training of Transformers》综述,以及近年提出的FlashAttention与FlashAttention-v2相关论文。
输入句子的嵌入处理
正式开始前,我们先以句子Life is short, eat dessert first为例,看看如何将它输入自注意力机制。和其他文本建模方法(比如循环神经网络、卷积神经网络)一样,我们首先要生成句子的嵌入表示。
为简化演示,这里我们的词典dc只包含输入句子中出现的单词。在实际应用中,词典会覆盖训练数据集中的所有单词,典型的词表规模在3万到5万之间。
输入代码:
sentence = 'Life is short, eat dessert first'
dc = {s:i for i,s
in enumerate(sorted(sentence.replace(',', '').split()))}
print(dc)
输出:
{'Life': 0, 'dessert': 1, 'eat': 2, 'first': 3, 'is': 4, 'short': 5}
接下来,我们用这个词典为每个单词分配一个整数索引:
输入代码:
import torch
sentence_int = torch.tensor(
[dc[s] for s in sentence.replace(',', '').split()]
)
print(sentence_int)
输出:
tensor([0, 4, 5, 2, 1, 3])
得到句子的整数向量表示后,我们就可以通过嵌入层,将输入编码为实数向量形式的嵌入。这里我们使用极小的3维嵌入,也就是每个输入单词对应一个3维向量。
注意:实际场景中的嵌入维度通常在数百到数千维,比如Llama 2的嵌入维度就达到了4096。这里使用3维纯粹是为了演示方便,方便我们查看每个向量的具体数值,不会让页面被数字占满。
句子包含6个单词,因此最终会得到一个6×3的嵌入矩阵:
输入代码:
vocab_size = 50_000 torch.manual_seed(123) embed = torch.nn.Embedding(vocab_size, 3) embedded_sentence = embed(sentence_int).detach() print(embedded_sentence) print(embedded_sentence.shape)
输出:
tensor([[ 0.3374, -0.1778, -0.3035],
[ 0.1794, 1.8951, 0.4954],
[ 0.2692, -0.0770, -1.0205],
[-0.2196, -0.3792, 0.7671],
[-0.5880, 0.3486, 0.6603],
[-1.1925, 0.6984, -1.4097]])
torch.Size([6, 3])
定义权重矩阵
接下来我们讲解应用最广泛的自注意力机制——缩放点积注意力,它是Transformer架构的核心组成部分。
自注意力会用到三个权重矩阵,分别记作$W_q$、$W_k$、$W_v$,它们是模型的可训练参数,会在训练过程中不断更新。这三个矩阵的作用,分别是将输入投影为序列的查询(query)、键(key)和值(value)分量。
对应的查询、键、值序列,通过权重矩阵$W$与嵌入输入$x$做矩阵乘法得到:
-
查询序列:$q(i) = x(i)W_q$,其中$i$为序列中第1到第$T$个位置
-
键序列:$k(i) = x(i)W_k$,其中$i$为序列中第1到第$T$个位置
-
值序列:$v(i) = x(i)W_v$,其中$i$为序列中第1到第$T$个位置
下标$i$代表输入序列中的token位置索引,序列总长度为$T$。
通过输入$x$与权重$W$计算查询、键、值向量
其中,$q(i)$和$k(i)$都是维度为$d_k$的向量。投影矩阵$W_q$和$W_k$的形状为 $d \times d_k$,而$W_v$的形状为 $d \times d_v$。
(注意:$d$代表每个单词向量$x$的维度。)
由于我们要计算查询向量和键向量的点积,因此这两个向量的元素数量必须相等($d_q = d_k$)。在很多大语言模型中,值向量也会采用相同的维度,也就是 $d_q = d_k = d_v$。不过值向量$v(i)$的元素数(决定了最终上下文向量的维度)本身是可以自由设置的。
在接下来的代码演示中,我们设置 $d_q = d_k = 2$,$d_v = 4$,并按如下方式初始化投影矩阵:
输入代码:
torch.manual_seed(123) d = embedded_sentence.shape[1] d_q, d_k, d_v = 2, 2, 4 W_query = torch.nn.Parameter(torch.rand(d, d_q)) W_key = torch.nn.Parameter(torch.rand(d, d_k)) W_value = torch.nn.Parameter(torch.rand(d, d_v))
(和前面的词嵌入向量一样,实际场景中$d_q$、$d_k$、$d_v$的维度会大得多,这里用小数值只是为了方便演示。)
计算未归一化注意力权重
现在,假设我们要计算第二个输入元素的注意力向量——此时第二个输入元素就作为查询。
接下来的内容,我们都以第二个输入$x(2)$为例
代码实现如下:
输入代码:
x_2 = embedded_sentence[1] query_2 = x_2 @ W_query key_2 = x_2 @ W_key value_2 = x_2 @ W_value print(query_2.shape) print(key_2.shape) print(value_2.shape)
输出:
torch.Size([2]) torch.Size([2]) torch.Size([4])
我们可以把这个逻辑推广到所有输入,计算出所有输入对应的键和值,因为后续计算未归一化注意力权重时会用到它们:
输入代码:
keys = embedded_sentence @ W_key
values = embedded_sentence @ W_value
print("keys.shape:", keys.shape)
print("values.shape:", values.shape)
输出:
keys.shape: torch.Size([6, 2]) values.shape: torch.Size([6, 4])
有了全部的键和值,我们就可以进入下一步:计算未归一化的注意力权重$\omega$(omega),如下图所示:
计算未归一化注意力权重$\omega$(omega)
如图所示,$\omega_{i,j}$ 等于查询向量与键向量的点积,即 $\omega_{i,j} = q(i) \cdot k(j)$。
举个例子,我们可以计算当前查询与第5个输入元素(对应索引4)之间的未归一化注意力权重:
输入代码:
omega_24 = query_2.dot(keys[4]) print(omega_24)
(注:$\omega$是希腊字母“omega”,上面代码中的变量名也由此而来。)
输出:
tensor(1.2903)
后续计算正式注意力权重时需要用到这些未归一化的$\omega$值,我们把所有输入token对应的$\omega$都计算出来,就像上图演示的那样:
输入代码:
omega_2 = query_2 @ keys.T print(omega_2)
输出:
tensor([-0.6004, 3.4707, -1.5023, 0.4991, 1.2903, -1.3374])
计算注意力权重
自注意力的下一步,是对未归一化的注意力权重$\omega$做归一化,通过softmax函数得到归一化后的注意力权重$\alpha$(alpha)。此外,在输入softmax之前,我们会先将$\omega$乘以 $1/\sqrt{d_k}$ 做缩放,如下所示:
计算归一化注意力权重$\alpha$
除以$d_k$的平方根进行缩放,是为了保证权重向量的欧氏长度保持在相近的量级,避免注意力权重过大或过小,防止数值不稳定,同时保障模型的训练收敛效果。
为什么偏偏是$\sqrt{d_k}$?因为$q$和$k$的点积是$d_k$个独立项的和,每一项的方差约为1,这意味着原始得分的方差会随$d_k$线性增长。除以$\sqrt{d_k}$后,就能抵消这种增长,让方差回到约1的水平。
代码中的注意力权重计算实现如下:
输入代码:
import torch.nn.functional as F attention_weights_2 = F.softmax(omega_2 / d_k**0.5, dim=0) print(attention_weights_2)
输出:
tensor([0.0386, 0.6870, 0.0204, 0.0840, 0.1470, 0.0229])
最后一步,是计算上下文向量$z(2)$。它是原始查询输入$x(2)$的注意力加权版本,通过注意力权重融合了所有其他输入元素的上下文信息:
注意力权重是针对特定输入元素计算的。这里我们选择的是输入元素$x(2)$。
代码实现如下:
输入代码:
context_vector_2 = attention_weights_2 @ values print(context_vector_2.shape) print(context_vector_2)
输出:
torch.Size([4]) tensor([0.5313, 1.3607, 0.7891, 1.3110])
注意,这个输出向量的维度($d_v = 4$)比原始输入向量的维度($d = 3$)更高,因为我们前面设置了$d_v > d$;实际上,嵌入维度$d_v$的取值是任意的。
自注意力完整实现
结合前面几节的内容,我们可以把自注意力机制的代码整合起来,封装成一个简洁的SelfAttention类:
输入代码:
import torch.nn as nn
class SelfAttention(nn.Module):
def __init__(self, d_in, d_out_kq, d_out_v):
super().__init__()
self.d_out_kq = d_out_kq
self.W_query = nn.Parameter(torch.rand(d_in, d_out_kq))
self.W_key = nn.Parameter(torch.rand(d_in, d_out_kq))
self.W_value = nn.Parameter(torch.rand(d_in, d_out_v))
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 / self.d_out_kq**0.5, dim=-1
)
context_vec = attn_weights @ values
return context_vec
遵循PyTorch的编码惯例,上面的SelfAttention类在__init__方法中初始化自注意力的参数,在forward方法中完成所有输入的注意力权重与上下文向量计算。我们可以这样使用这个类:
输入代码:
torch.manual_seed(123) d_in, d_out_kq, d_out_v = 3, 2, 4 sa = SelfAttention(d_in, d_out_kq, d_out_v) print(sa(embedded_sentence))
输出:
tensor([[-0.1564, 0.1028, -0.0763, -0.0764],
[ 0.5313, 1.3607, 0.7891, 1.3110],
[-0.3542, -0.1234, -0.2627, -0.3706],
[ 0.0071, 0.3345, 0.0969, 0.1998],
[ 0.1008, 0.4780, 0.2021, 0.3674],
[-0.5296, -0.2799, -0.4107, -0.6006]], grad_fn=<MmBackward0>)
可以看到,输出的第二行和我们上一节算出的context_vector_2完全一致:tensor([0.5313, 1.3607, 0.7891, 1.3110])。
多头注意力
在文章开头的第一张图里(这里再放一次方便查看),我们能看到Transformer中使用了一个叫“多头注意力”的模块。
原始Transformer架构中的多头注意力模块,来源:https://arxiv\.org/abs/1706\.03762
这个“多头”注意力模块,和我们前面讲解的自注意力机制(缩放点积注意力)是什么关系呢?
在缩放点积注意力中,输入序列通过查询、键、值三个矩阵做变换。在多头注意力的语境里,这一套矩阵就对应一个注意力头。下图总结了我们前面实现的单个注意力头:
前文实现的自注意力机制总结
顾名思义,多头注意力就是包含多个这样的“头”,每个头都有自己的查询、键、值矩阵。这个概念和卷积神经网络中使用多个卷积核的思路类似,最终会输出多通道的特征图。
多头注意力:包含多个头的自注意力
用代码演示的话,我们可以基于前面的SelfAttention类,写一个MultiHeadAttentionWrapper包装类:
class MultiHeadAttentionWrapper(nn.Module):
def __init__(self, d_in, d_out_kq, d_out_v, num_heads):
super().__init__()
self.heads = nn.ModuleList(
[SelfAttention(d_in, d_out_kq, d_out_v)
for _ in range(num_heads)]
)
def forward(self, x):
return torch.cat([head(x) for head in self.heads], dim=-1)
其中d_*参数和SelfAttention类中的定义完全一致,唯一新增的参数是注意力头的数量:
-
d_in:输入特征向量的维度 -
d_out_kq:查询与键的输出维度 -
d_out_v:值的输出维度 -
num_heads:注意力头的数量
我们用这些参数初始化num_heads个SelfAttention实例,并用PyTorch的nn.ModuleList存储这些实例。
在前向传播时,每个存储在self.heads中的SelfAttention头都会独立处理输入x,然后将所有头的结果在最后一维(dim=-1)拼接起来。我们来实际运行一下:
首先,为了简化演示,我们用单个自注意力头,输出维度设为1:
输入代码:
torch.manual_seed(123) d_in, d_out_kq, d_out_v = 3, 2, 1 sa = SelfAttention(d_in, d_out_kq, d_out_v) print(sa(embedded_sentence))
输出:
tensor([[-0.0185],
[ 0.4003],
[-0.1103],
[ 0.0668],
[ 0.1180],
[-0.1827]], grad_fn=<MmBackward0>)
现在我们把它扩展为4个注意力头:
输入代码:
torch.manual_seed(123)
mha = MultiHeadAttentionWrapper(
d_in, d_out_kq, d_out_v, num_heads=4
)
context_vecs = mha(embedded_sentence)
print(context_vecs)
print("context_vecs.shape:", context_vecs.shape)
输出:
tensor([[-0.0185, 0.0170, 0.1999, -0.0860],
[ 0.4003, 1.7137, 1.3981, 1.0497],
[-0.1103, -0.1609, 0.0079, -0.2416],
[ 0.0668, 0.3534, 0.2322, 0.1008],
[ 0.1180, 0.6949, 0.3157, 0.2807],
[-0.1827, -0.2060, -0.2393, -0.3167]], grad_fn=<CatBackward0>)
context_vecs.shape: torch.Size([6, 4])
从输出可以看到,前面单个自注意力头的结果,正好对应输出张量的第一列。
可以看到,多头注意力的结果是一个6×4的张量:我们有6个输入token,4个自注意力头,每个自注意力头输出1维结果。而前面单头自注意力部分,我们也得到过6×4的张量——那是因为我们把单头的输出维度设成了4,而不是1。
那么问题来了:既然单头自注意力本身就能调整输出嵌入的大小,为什么还需要多头注意力呢?
“提升单个自注意力头的输出维度”和“使用多个注意力头”,两者的本质区别在于模型处理、学习数据的方式不同。虽然两种方式都能提升模型表征不同特征、不同数据维度的能力,但底层逻辑完全不同。
比如,多头注意力中的每个头,都有可能学会关注输入序列的不同部分,捕捉数据中不同层面的关系。这种表征的多样性,正是多头注意力成功的关键。
此外,多头注意力的计算效率更高,尤其适合并行计算。每个头都可以独立处理,非常适配GPU、TPU这类擅长并行运算的现代硬件加速设备。
简而言之,使用多头注意力不只是为了提升模型容量,更是为了增强模型学习数据中多样化特征与关联的能力。举个例子,70亿参数的Llama 2模型就使用了32个注意力头。
因果自注意力
本节我们将前面讲解的自注意力机制,改造为因果自注意力机制——它主要用于GPT这类(解码器风格的)文本生成大语言模型。因果自注意力也常被称为“掩码自注意力”,在原始Transformer架构中,它对应“掩码多头注意力”模块。为了简化讲解,本节我们只看单个注意力头,其原理可以直接推广到多头场景。
原始Transformer架构中的因果自注意力模块,来源:《Attention Is All You Need》, https://arxiv\.org/abs/1706\.03762
因果自注意力的作用是:序列中每个位置的输出,只能依赖当前位置之前的已知输出,不能看到未来位置的信息。简单来说,就是保证下一个单词的预测,只由它前面的单词决定。
在GPT类大语言模型中,为了实现这一点,处理每个token时,我们都会把输入文本中位于当前token之后的“未来token”掩码掉。
下图演示了如何对注意力权重施加因果掩码,隐藏输入中的未来token:
为了演示并实现因果自注意力,我们沿用前面的未归一化注意力得分与注意力权重。首先快速回顾一下自注意力部分的注意力得分计算:
输入代码:
torch.manual_seed(123) d_in, d_out_kq, d_out_v = 3, 2, 4 W_query = nn.Parameter(torch.rand(d_in, d_out_kq)) W_key = nn.Parameter(torch.rand(d_in, d_out_kq)) W_value = nn.Parameter(torch.rand(d_in, d_out_v)) x = embedded_sentence keys = x @ W_key queries = x @ W_query values = x @ W_value # attn_scores 就是前面的 "omegas",即未归一化注意力权重 attn_scores = queries @ keys.T print(attn_scores) print(attn_scores.shape)
输出:
tensor([[ 0.0613, -0.3491, 0.1443, -0.0437, -0.1303, 0.1076],
[-0.6004, 3.4707, -1.5023, 0.4991, 1.2903, -1.3374],
[ 0.2432, -1.3934, 0.5869, -0.1851, -0.5191, 0.4730],
[-0.0794, 0.4487, -0.1807, 0.0518, 0.1677, -0.1197],
[-0.1510, 0.8626, -0.3597, 0.1112, 0.3216, -0.2787],
[ 0.4344, -2.5037, 1.0740, -0.3509, -0.9315, 0.9265]],
grad_fn=<MmBackward0>)
torch.Size([6, 6])
和之前自注意力部分一样,上面的输出是6×6的张量,保存了6个输入token两两之间的未归一化注意力权重(也叫注意力得分)。
之前我们是这样通过softmax计算缩放点积注意力的:
输入代码:
attn_weights = torch.softmax(attn_scores / d_out_kq**0.5, dim=1) print(attn_weights)
输出:
tensor([[0.1772, 0.1326, 0.1879, 0.1645, 0.1547, 0.1831],
[0.0386, 0.6870, 0.0204, 0.0840, 0.1470, 0.0229],
[0.1965, 0.0618, 0.2506, 0.1452, 0.1146, 0.2312],
[0.1505, 0.2187, 0.1401, 0.1651, 0.1793, 0.1463],
[0.1347, 0.2758, 0.1162, 0.1621, 0.1881, 0.1231],
[0.1973, 0.0247, 0.3102, 0.1132, 0.0751, 0.2794]],
grad_fn=<SoftmaxBackward0>)
这个6×6的输出就是注意力权重,和我们之前自注意力部分计算的结果一致。
而在GPT类大语言模型中,模型的训练方式是从左到右,逐个读取、生成token(单词)。如果我们有一条训练文本Life is short eat dessert first,就会形成如下的训练模式:箭头右侧单词的上下文向量,只能融合自身和箭头左侧的单词信息。
-
"Life"→"is" -
"Life is"→"short" -
"Life is short"→"eat" -
"Life is short eat"→"dessert" -
"Life is short eat dessert"→"first"
实现这个模式最简单的方法,就是对注意力权重矩阵的上三角(主对角线以上的部分)做掩码,如下图所示。这样一来,在计算上下文向量(即输入的注意力加权和)时,“未来”的单词就不会被纳入计算。
主对角线以上的注意力权重需要被掩码掉
代码层面,我们可以用PyTorch的tril函数实现:先生成一个由0和1组成的掩码矩阵。
输入代码:
block_size = attn_scores.shape[0] mask_simple = torch.tril(torch.ones(block_size, block_size)) print(mask_simple)
输出:
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.]])
接下来,我们把注意力权重和这个掩码相乘,将主对角线以上的所有注意力权重置零:
输入代码:
masked_simple = attn_weights*mask_simple print(masked_simple)
输出:
tensor([[0.1772, 0.0000, 0.0000, 0.0000, 0.0000, 0.0000],
[0.0386, 0.6870, 0.0000, 0.0000, 0.0000, 0.0000],
[0.1965, 0.0618, 0.2506, 0.0000, 0.0000, 0.0000],
[0.1505, 0.2187, 0.1401, 0.1651, 0.0000, 0.0000],
[0.1347, 0.2758, 0.1162, 0.1621, 0.1881, 0.0000],
[0.1973, 0.0247, 0.3102, 0.1132, 0.0751, 0.2794]],
grad_fn=<MulBackward0>)
这是掩码未来单词的一种实现方式,但你会发现:此时每一行的注意力权重不再求和等于1。为了解决这个问题,我们可以对每一行重新做归一化,让它们的和重新变为1——这是注意力权重的标准约定。
输入代码:
row_sums = masked_simple.sum(dim=1, keepdim=True) masked_simple_norm = masked_simple / row_sums print(masked_simple_norm)
输出:
tensor([[1.0000, 0.0000, 0.0000, 0.0000, 0.0000, 0.0000],
[0.0532, 0.9468, 0.0000, 0.0000, 0.0000, 0.0000],
[0.3862, 0.1214, 0.4924, 0.0000, 0.0000, 0.0000],
[0.2232, 0.3242, 0.2078, 0.2449, 0.0000, 0.0000],
[0.1536, 0.3145, 0.1325, 0.1849, 0.2145, 0.0000],
[0.1973, 0.0247, 0.3102, 0.1132, 0.0751, 0.2794]],
grad_fn=<DivBackward0>)
可以看到,现在每一行的注意力权重之和都为1了。
在Transformer这类模型中,使用归一化的注意力权重主要有两点优势:第一,求和为1的归一化注意力权重类似概率分布,更直观地体现了模型对输入不同部分的关注比例;第二,约束权重和为1,可以控制权重与梯度的量级,优化训练过程的稳定性。
更高效的掩码方式:无需二次归一化
上面我们实现的因果自注意力流程是:先计算注意力得分,再计算注意力权重,接着掩码掉上三角的权重,最后重新归一化。整个流程总结如下图:
前文实现的因果自注意力流程
其实还有一种更高效的方式,可以得到完全相同的结果。这种方法是在输入softmax计算注意力权重之前,就把注意力得分中主对角线以上的值替换为负无穷。流程总结如下:
另一种更高效的因果自注意力实现方案
我们可以用PyTorch写出这个流程,首先对注意力得分的上三角做掩码:
输入代码:
mask = torch.triu(torch.ones(block_size, block_size), diagonal=1) masked = attn_scores.masked_fill(mask.bool(), -torch.inf) print(masked)
上面的代码先生成一个掩码矩阵:主对角线及以下为0,以上为1。其中torch.triu(上三角)会保留矩阵主对角线及以上的元素,将下方元素置零,保留上三角部分;对应的torch.tril(下三角)则保留主对角线及以下的元素。
然后masked_fill方法会把掩码中为1的位置(主对角线以上)的所有元素替换为-torch.inf(负无穷),结果如下:
输出:
tensor([[ 0.0613, -inf, -inf, -inf, -inf, -inf],
[-0.6004, 3.4707, -inf, -inf, -inf, -inf],
[ 0.2432, -1.3934, 0.5869, -inf, -inf, -inf],
[-0.0794, 0.4487, -0.1807, 0.0518, -inf, -inf],
[-0.1510, 0.8626, -0.3597, 0.1112, 0.3216, -inf],
[ 0.4344, -2.5037, 1.0740, -0.3509, -0.9315, 0.9265]],
grad_fn=<MaskedFillBackward0>)
接下来,我们照常应用softmax函数,就能得到经过掩码且归一化的注意力权重:
输入代码:
attn_weights = torch.softmax(masked / d_out_kq**0.5, dim=1) print(attn_weights)
输出:
tensor([[1.0000, 0.0000, 0.0000, 0.0000, 0.0000, 0.0000],
[0.0532, 0.9468, 0.0000, 0.0000, 0.0000, 0.0000],
[0.3862, 0.1214, 0.4924, 0.0000, 0.0000, 0.0000],
[0.2232, 0.3242, 0.2078, 0.2449, 0.0000, 0.0000],
[0.1536, 0.3145, 0.1325, 0.1849, 0.2145, 0.0000],
[0.1973, 0.0247, 0.3102, 0.1132, 0.0751, 0.2794]],
grad_fn=<SoftmaxBackward0>)
为什么这种方法有效?因为最后一步的softmax函数会把输入值转化为概率分布。当输入中存在负无穷时,$e^{-\infty}$趋近于0,因此这些位置对输出概率没有贡献,相当于概率为0。
结语
本文通过逐行编码的方式,拆解了自注意力的内部工作原理。在此基础上,我们进一步学习了多头注意力——这是Transformer大语言模型的核心组件。
随后我们还编码实现了交叉注意力(自注意力的一种变体,在处理两个独立序列时效果突出),以及因果自注意力——这是GPT、Llama等解码器风格大语言模型生成连贯、符合上下文的序列的关键概念。
通过从零编码实现这些复杂机制,希望你已经对Transformer与大语言模型中使用的自注意力机制有了透彻的理解。
(注:本文中的代码仅用于演示原理。如果你要在大语言模型训练中实现自注意力,推荐使用Flash Attention这类优化实现,它们能显著降低内存占用与计算开销。)
附加主题:交叉注意力
在前面自注意力与因果注意力的代码讲解中,我们设置了 $d_q = d_k = 2$,$d_v = 4$。也就是说,查询序列和键序列使用了相同的维度。虽然值矩阵$W_v$通常也会设置成和查询、键矩阵相同的维度(比如PyTorch自带的MultiHeadAttention类就是如此),但值的维度其实可以自由设置……
【以下为付费内容,请跳转至原文付费订阅,6$,两杯咖啡钱】
Understanding and Coding Self-Attention, Multi-Head Attention, Causal-Attention, and Cross-Attention in LLMs,by Sebastian Raschka, 2024-01-14













