Transformer架构详解(3):自注意力机制深入理解
一、引入:Self-Attention
1.1 Attention 的作用
Attention 中最核心的三个概念是 Query(Q)、Key(K)和 Value(V)。这三个词直接借鉴了信息检索领域的术语,下面先从一个真实的语言难题出发:
句子 A:The bank of the river.(河边) 句子 B:Money in the bank.(银行)
同一个词 “bank” 在两句话里含义截然不同。机器翻译时如何判断?答案是看上下文中的其他词:句子 A 里 “river”(河流)权重最高,句子 B 里 “money”(钱)权重最高。这种”让每个词去关注其他词、用相关程度加权汇总语义”的机制,就是 Self-Attention。
用大白话说:Attention 就是让句子中的每个词去”关注”其他所有词,然后根据关注程度加权汇总信息,从而得到一个融合了上下文的”新含义”。
更进一步,打个检索的比方:你在查一本百科全书里关于”苹果”的信息,你会先生成一个查询(Query):“苹果是什么?“,然后翻阅每一页的标题(Key)来判断相关性,最后把相关页面的正文内容(Value)按相关程度加权合并。Attention 做的正是这件事:
- Query(Q):我想查什么——当前词想要获取的信息
- Key(K):索引/标题——每个词用来被别人匹配的标识
- Value(V):实际内容——匹配成功后要传递的信息
在 Self-Attention 中,序列里的每个 token 都同时扮演这三种角色:它既是提问者(生成 Q),也是被查询的索引(生成 K),还是信息的提供者(生成 V)。通过 Q 和 K 的匹配来计算”注意力权重”,再用这些权重对 V 做加权求和,就完成了信息的聚合。
“银行”还是”河岸”的歧义问题,RNN 时代需要把前面所有词的信息按顺序一步步传递过来才能解决;Self-Attention 则让所有词在同一步内直接互相”比对”,长距离依赖不再衰减。
1.2 三元组的形式化定义
回到 Self-Attention 的语境。给定输入序列 X(形状为 ),三个线性变换将 X 映射到不同的空间:
- :每个 token 生成自己的”查询向量”——“我需要什么信息”
- :每个 token 生成自己的”索引向量”——“我能提供什么信息的线索”
- :每个 token 生成自己的”内容向量”——“我实际携带的信息”
Q、K、V 之所以要从同一个 X 做三次不同的线性变换,而不是直接用 X 本身,原因在于解耦不同的角色。一个 token “需要什么”和它”能提供什么”往往是不同的。比如在一个句子里,动词可能需要关注它的主语和宾语(Query 的方向),但它作为被别人关注的对象时,提供的是动作语义(Key/Value 的方向)。三个独立的投影矩阵让模型有自由度去学习这些不同的映射。
二、计算过程:从输入到输出
假设我们有一个长度为 N 的序列,每个 token 用一个 d 维向量表示,输入矩阵 X 的形状为 。
2.1 线性投影生成 Q、K、V
输入 X 分别乘以三个权重矩阵,得到 Query、Key、Value:
其中 是可学习的参数矩阵,形状都是 (d,d)。这三次矩阵乘法就是三次 GEMM 操作——后续 CUDA 优化和张量并行的核心对象之一。
在实际实现中,为了提高 GPU 利用率,通常会把 合并成一个大矩阵 (形状 ),做一次 GEMM 然后 split,这样能更好地利用 GPU 的算力。
2.2 计算注意力分数
用 Q 和 K 的内积来衡量每对 token 之间的”匹配度”:
得到的 S 是一个 N×N 的矩阵, 表示第 i 个 token 对第 j 个 token 的关注程度(原始分数)。
2.3 缩放(Scale)
将分数除以 ( 是每个头的维度,后面会解释):
为什么要缩放?
直觉上说,当维度 很大时,Q 和 K 的内积值会变得很大(因为是 个分量相加),导致 softmax 的输入值差异悬殊。softmax 对大数值非常敏感——输入差距一大,输出就会”极化”成接近 one-hot 的分布,梯度几乎为零,训练就卡住了。除以 能把方差拉回到 1 附近,让 softmax 工作在一个梯度比较健康的区间。
2.4 Softmax 归一化
对每一行做 softmax,把原始分数变成概率分布(每行之和为 1):
现在表示:第 i 个 token 分配给第 j 个 token 的注意力权重,每行之和为 1。
Softmax 的作用是双重的:一方面把任意实数映射到 (0,1) 区间,使其可以作为权重;另一方面保证每行的权重之和为 1,形成一个合法的概率分布。
AI Infra 关联:Softmax 是一个看似简单但在高性能场景下需要精心优化的算子。标准实现需要对每行做两遍扫描(第一遍求最大值和指数和,第二遍归一化),Online Softmax 算法将其合并为一遍扫描,FlashAttention 正是基于此实现了 Attention 的高效融合。
2.5 加权求和
用注意力权重 A 对 Value 矩阵 V 做加权求和:
最终每个 token 得到一个 维向量,其中融合了它”应该关注”的所有其他 token 的信息。关注程度由 A 的权重决定。
完整公式(一行总结)
2.6 输出投影
最后还要过一个输出投影矩阵 :
将多头拼接后的结果映射回模型的隐藏维度(后面会详述多头注意力)。这个投影在 Multi-Head Attention 中尤其重要——它负责将多个头拼接后的表示重新混合。
三、代码实现
3.1 单头注意力:Single-Head Self-Attention
1 | |
3.2 多头注意力:Multi-Head Self-Attention
1 | |
四、Multi-Head Attention
4.1 多头注意力的引入
多头注意力机制是在自注意力机制的基础上发展起来的,是自注意力机制的变体,旨在增强模型的表达能力和泛化能力。它通过使用多个独立的注意力头,分别计算注意力权重,并将它们的结果进行拼接或加权求和,从而获得更丰富的表示。
单头 Attention 只有一组 QKV 投影,意味着模型只能学习一种”关注模式”。但语言中 token 之间的关系是多维度的——同一个词和其他词之间可能同时存在句法关系(主谓一致)、语义关系(同义替换)、位置关系(相邻词的局部模式)等。
打个比方:在一次项目评审会议上,只派一个评审员去审阅整个项目,他只能从自己擅长的角度提出意见。如果派出一个评审团——一位看技术架构,一位看代码质量,一位看测试覆盖率,一位看文档完整性——每个人独立给出评分和建议,最后汇总成一份综合评审报告,覆盖面就远比单人评审要全面得多。
Multi-Head Attention 就是这个”评审团”机制。每个头有自己独立的 投影参数,在不同的子空间中捕捉不同类型的关系。
4.2 数学原理
假设模型隐藏维度 ,头数 ,则每个头的维度 。
Multi-Head Attention 的完整公式:
其中每个头的投影矩阵 、、、 的形状是 。但实际实现中,并不会真的维护 h 组小矩阵——而是用一个大矩阵 (形状 (512,512))做一次投影,然后 reshape 成 来切分。这样做是等价的:大矩阵可以看作 8 个小矩阵纵向拼接。
通用公式:MHA 的参数量为 (不算 bias 的情况)。不管有多少个头,总参数量不变——头数只影响切分方式,不影响总参数。这是因为每增加一个头,每个头的维度相应减小,两者乘积(即总投影维度)始终等于 。
4.3 多头的意义
多头结构天然适合并行化。8 个头的计算完全独立,可以:
- GPU 内并行:利用 batch 维度,在一次 CUDA kernel 启动中同时处理所有头
- 多卡张量并行(Tensor Parallelism):将不同头分配到不同 GPU 上,每张 GPU 只计算自己负责的若干头。比如 8 个头分到 4 张 GPU,每张处理 2 个头。最后通过一次 AllReduce 通信汇总输出投影的结果
五、Masked Multi-Head Attention
在《Attention is all you need》中,Decoder 里有一个 Masked Multi-Head Attention,即具有掩码的多头注意力机制。
由于在 Decoder 中的 Embedding 层和 Positional Encoding 层是对目标序列所有单词都进行了嵌入,但我们利用 Transformer 模型在生成序列时,是一个个单词往外蹦的;那么,确保在预测序列的特定位置时,模型只能使用到该位置之前的信息,其后面的信息就不能被注意力机制看到(也就是保证模型在训练的时候只能看到当前单词之前的单词,不能看到之后的),从而防止信息泄露。就需要想办法去遮挡后面的信息。
比如说一个字符串序列 ,我们在使用自注意力机制的时候,会计算当前单词 和其它单词()的关系。
但是,当我们在 Transformer 中的 Decoder 中开始需要生成序列时,当前位置需要预测的单词,只能允许和前面的信息有关,就比如我们要预测 是什么单词,那我们只能使用 和 的信息,并不能用到 的信息。
那如何在 Transformer 网络中去遮挡后面的信息呢?
主要是通过在自注意力机制中应用一个掩码(Mask)来实现的。准确地说,还是通过矩阵来进行操作。
例如,现在有段话:“I am fine”, 我们提前计算好了他们之间的注意力得分,如下图所示:

但当在计算单词 “am” 的注意力得分时,只能访问它自己和它之前的单词,就不能获得 “am”之后的单词信息,比如 “fine”,因为 “fine” 是之后才会生成的单词;如下图所示:

- 当我们访问到
位置时,只能获取 <start>和它自己的注意力得分,其他的不能获取 - 当我们访问到 “I“ 位置时,只能获取 “I” 和 <start>, 以及 “I” 和 它自己的注意力得分,其他的不能获取
- 当我们访问到 ”am“ 位置时,只能获取 ”am“ 和 <start>, “am” 和 “I”, 以及 “am” 和 它自己的注意力得分,其他的不能获取
- 以此类推......
为了防止解码器看到未来的信息,就需要在已得到的注意力分数矩阵式加上一个 mask 机制(或者叫 mask 矩阵),如下图所示:

当以上得到的橙色矩阵再经过 sigmoid 函数时,相对“当前单词”的“未来单词”的注意力得分就会变为0,这样就不会访问到未来信息。
