注意力机制详解:从QKV到多头注意力与PyTorch实现
在深度学习里,注意力机制(Attention Mechanism)几乎是近年来最值得讲清楚的概念之一。它是 Transformer 模型的基石,也是 GPT、BERT、ViT 等一系列模型的底层核心。很多人一开始接触 Transformer 时,最难跨过去的门槛并不是代码,而是“模型为什么要设计 Query、Key、Value 这三样东西”以及“注意力权重到底在计算什么”。本文就从最基础的“什么是注意力机制”开始,先把概念讲透,再逐步推导到缩放点积注意力(Scaled Dot-Product Attention)和多头注意力(Multi-Head Attention),最后用 PyTorch 写一个最小可运行的实现,帮助你真正理解输入输出、矩阵形状和实际效果。
这篇文章适合正在学习深度学习基础、已经知道 RNN/CNN 但还没弄懂 Transformer 的读者。读完以后,你不仅能解释注意力机制的含义,还能手写一个简单的自注意力模块,知道怎么验证结果,也能在模型效果不理想时沿着正确的方向排查问题。
1. 为什么深度学习模型需要注意力机制
1.1 从人类的注意力行为说起
“注意力”这个词首先来自人类认知行为。我们在阅读一段文本时,并不是每个字都同等重要。比如读到“他推开家门,发现一只猫坐在沙发上”,你会自然而然地关注“猫”“沙发”这两个词,而不会对“推开”分配同等权重。这种根据当前目标动态调整关注重点的能力,就是注意力。
深度学习里的注意力机制,模仿的正是这个过程。它让模型在处理某一个位置的输出时,不是只依赖最近邻的信息,而是能够动态地从全部输入里挑选出“与当前位置关系更密切”的内容,并给这些内容分配更高的权重。
用一句话概括:注意力机制是一套“动态加权求和”的规则。它根据当前查询的需求,计算输入序列中每个元素的重要性,再按重要性把输入的信息聚合起来。
1.2 定长向量带来的信息瓶颈
在注意力机制出现之前,处理序列数据主要依靠 RNN、LSTM 这类循环网络。RNN 的核心问题在于:无论输入序列多长,最后通常只保留一个固定长度的隐藏状态向量。这个向量要压缩整句信息,必然会出现信息瓶颈。句子越长,早期信息在长期传递过程中越容易被遗忘。
LSTM 通过门控机制缓解了长距离遗忘问题,但仍然要顺序计算。每个时间步的隐藏状态都依赖上一个时间步,导致训练时无法并行。更重要的是,当模型要回答“张三在杭州工作,李四在上海工作,那谁在杭州?”这种问题时,RNN 虽然理论上能记住全部信息,实际训练中却很难让最后一步的表示精准定位到“张三”和“杭州”的关联。
注意力机制的出现改变了这一局面。它让模型可以直接访问序列中的所有位置,不需要依赖上一个时间步逐步传递信息。当前输出需要什么内容,就直接从输入位置中去找,找到以后把相关内容加权取出来。
1.3 注意力机制的本质:动态加权求和
注意力机制可以拆成三步:
- 计算查询(Query)与每个输入位置(Key)之间的相似度。
- 用相似度作为未归一化的权重,经过 softmax 转换为概率分布。
- 按这个概率分布对每个输入位置对应的值(Value)进行加权求和。
这个过程中,最重要的一点是“权重是动态计算的”,而不是像卷积核那样固定的。对于同一个输入序列,查询不同,注意力的分布就不同;对于同一个查询,输入序列不同,注意力的分布也会不同。这正是注意力机制区别于标准卷积、池化和全连接的地方。
为了避免歧义,可以把注意力机制理解为一张“查表”操作:你有一个问题,表格里每一行有一个键值对,先根据问题匹配键,然后取出对应值。注意力机制只是把匹配过程改成可导的相似度计算,把取值过程改成加权求和,从而让整个操作能够放进神经网络里用梯度下降训练。
1.4 与全连接、卷积、池化的对比
把注意力机制和常见网络层放在一起对比,会更清楚它解决了什么问题。
| 网络层 | 处理方式 | 感受范围 | 权重是否动态 | 主要问题 |
|---|---|---|---|---|
| 全连接 | 每个输出与所有输入连接 | 全局 | 否 | 参数随输入长度爆炸,无法处理变长序列 |
| 卷积 | 局部窗口内加权求和 | 局部 | 否 | 需要多层堆叠才能覆盖长距离依赖 |
| 池化 | 固定窗口内取最大值或平均值 | 局部 | 否 | 丢失位置和组合信息 |
| 注意力 | 全局范围内动态加权求和 | 全局 | 是 | 计算量与序列长度平方相关,需要位置编码 |
从这张表可以看出,注意力机制在“全局建模”和“动态选择”这两个维度上都有明显优势。这也是 Transformer 选择它作为基础模块的原因之一。不过这份优势是有代价的:当序列长度为 n 时,注意力矩阵是 n×n,计算复杂度为 O(n²),长文本场景需要专门优化。
2. 注意力机制的计算基础:Query、Key、Value 与相似度打分
2.1 从检索场景理解 Query、Key、Value
注意力机制里的 Query、Key、Value 经常让新手困惑。一个比较直观的理解方式是参照搜索引擎。
假设你想检索“深度学习中的注意力机制是什么”。这句话是你的 Query。数据库里有很多文档,每篇文档都有一个标题作为 Key,正文内容作为 Value。搜索引擎先计算 Query 和每篇文档标题的匹配程度,再把匹配程度高的文档正文返回给你。匹配过程就是 Query 与 Key 的相似度计算,返回过程就是对 Value 的挑选或聚合。
在注意力机制中,Query、Key、Value 都是向量:
- Query 表示“我当前需要什么信息”。
- Key 表示“每个输入位置能提供什么信息”。
- Value 表示“每个输入位置实际携带的内容”。
模型通过 Query 和 Key 的匹配程度,决定从哪些 Value 中取信息,取多少比例。如果 Query 和某个 Key 很相似,对应 Value 的权重就高,这个位置的信息就会更多地进入输出。
2.2 相似度打分函数
有了 Query 和 Key,第一步是计算相似度分数。深度学习中常用几种打分方式:
| 打分方式 | 计算公式 | 特点 |
|---|---|---|
| 点积 | score = Q·K | 实现简单,但数值范围随向量维度增大 |
| 缩放点积 | score = Q·K / sqrt(d_k) | 缓解点积值过大,Transformer 默认方案 |
| 加性注意力 | score = v^T tanh(W_q Q + W_k K) | 表达能力更强,但计算开销更大 |
| 双线性注意力 | score = Q^T W K | 引入可学习矩阵,灵活性更高 |
Transformer 选择缩放点积注意力,除了效果不差之外,更重要的原因是矩阵乘法可以高度优化,所以在 GPU 上计算效率明显高于加性注意力。
需要注意,点积相似度依赖于向量维度和向量方向。当向量维度 d_k 较大时,点积结果可能变得很大,导致后续 softmax 的梯度非常小。缩放因子 sqrt(d_k) 就是用来控制分数范围的。
2.3 softmax 归一化与注意力权重
相似度分数只是未归一化的权重。为了让不同位置的权重可以比较,需要经过 softmax 函数,把所有位置的分数转换成总和为 1 的概率分布。
softmax 的计算方式如下:
softmax 有两个作用:
- 把分数变成非负的、可解释的权重。
- 拉大相对差异,让高相似度位置的权重更突出。
当某些位置需要完全屏蔽时,可以把该位置的分数设置为负无穷(-inf)。经过 softmax 后,exp(-inf) 为 0,对应位置的注意力权重就是 0,该位置不会对输出产生任何影响。这也是注意力掩码的基本原理。
2.4 加权求和得到输出
得到注意力权重后,最后一步就是对 Value 做加权求和:
这一步是线性的,所以整个注意力计算过程对于输入是可导的。模型可以通过反向传播,逐步学习到更合适的 Q、K、V 映射矩阵,从而让注意力分布更符合任务需求。
下面用 NumPy 实现一个最简形式的注意力机制,帮助你形成直观感受:
这个例子中,Query 与第一个 Key 相似度最高,所以第一个 Value 得到的权重最大,输出更接近 [10, 0]。这里最终输出不是直接选中第一个 Value,而是做加权平均,因此保留了可导性和平滑性。
3. 自注意力机制:每个位置都可以参考全局
3.1 什么是自注意力
前面讲到的注意力机制,Query 可以来自外部,也可以来自输入序列本身。当 Query、Key、Value 都来自同一个输入序列时,这种结构就叫自注意力(Self-Attention)。
自注意力的初衷是:对于一个序列中的每个位置,它需要理解自己在整个序列中的上下文,从而决定哪些位置和自己相关。以句子“小明喜欢看猫,因为它很可爱”为例,模型要理解“它”指代“猫”,就需要让“它”这个位置去关注“猫”这个位置。这个关注关系不是预先写死的,而是模型在训练中自动学出来的。
在自注意力中,输入序列的每个向量会先乘以三个不同的权重矩阵,分别得到 Query、Key、Value。也就是说,每个位置既是“提问者”,也是“被检索的内容”。
3.2 为什么 Transformer 选择 Self-Attention
Transformer 论文里使用 Self-Attention 的核心原因有三个。
第一是长距离依赖建模。CNN 需要通过不断堆叠卷积层来扩大感受野,RNN 需要通过时间步传递信息。Self-Attention 一步到位,任何两个位置之间只需要一次计算就能建立关系。
第二是并行计算。RNN 必须按时间顺序计算,Self-Attention 对序列中所有位置同时计算,训练速度显著提升。
第三是稳定的训练动态。相比 RNN 的链式求导,Self-Attention 的路径更短,梯度可以更直接地在长距离之间传播。
当然 Self-Attention 也有自己的缺点,主要是 O(n²) 的计算复杂度。后面发展的稀疏注意力、线性注意力、FlashAttention 等方法,都是在解决这个复杂度问题。
3.3 缩放点积注意力公式
自注意力的标准计算公式如下:
其中:
- Q 的形状为 [batch_size, seq_len, d_k]
- K 的形状为 [batch_size, seq_len, d_k]
- V 的形状为 [batch_size, seq_len, d_v]
- K^T 表示 K 的最后两个维度转置,得到 [batch_size, d_k, seq_len]
- Q 与 K^T 矩阵相乘后,得到 [batch_size, seq_len, seq_len] 的注意力分数
这里有一个细节:为什么要除以 sqrt(d_k)?
如果 d_k 很大,两个向量点积后的方差会随之变大。假设向量每个分量均值为 0、方差为 1,那么 d_k 维向量的点积均值是 0,方差是 d_k。方差越大,点积值的分布越分散,softmax 后某些位置的权重会接近 1,其余位置接近 0,梯度会非常小。除以 sqrt(d_k) 后,点积的方差被拉回到 1 附近,softmax 区域更平滑,梯度更稳定。
这里特别容易记错的是:缩放因子是 sqrt(d_k),而不是 d_k。d_k 表示每个注意力头的维度,不是整个模型的隐藏维度。
3.4 位置编码的必要性
Self-Attention 本身对位置没有感知。如果把序列的位置打乱,注意力分数不会变化,因为注意力计算只依赖向量内容,不依赖位置信息。但语言和图像中的顺序往往很重要,因此必须在输入中注入位置信息。
Transformer 的解决方案是位置编码。最经典的位置编码使用正弦和余弦函数:
位置编码与词向量相加后送入注意力层,模型才能区分不同位置的向量。后续工作也提出了可学习位置编码、相对位置编码、旋转位置编码(RoPE)等变体,但它们要解决的问题都一样:让自注意力知道“谁在哪个位置”。
学习自注意力时,不要忽略位置编码。很多人只记住 QKV 计算,却忘了 Transformer 之所以能建模顺序信息,是因为额外注入了位置信号。
3.5 掩码注意力的两种常见类型
实际训练中,注意力矩阵不是永远完整的。有两种常见掩码:
- Padding Mask:用于忽略序列中补零的无意义位置。通常把 padding 位置对应的 Key 分数设为
-inf,从而让这些位置不参与 attention。 - Casual Mask(因果掩码):用于自回归模型。在预测第 t 个位置时,不允许模型看到第 t 个位置之后的信息,因此注意力矩阵的上三角部分被设置为
-inf。
掩码的实现方式是在 softmax 之前对分数矩阵做 masked_fill。一个常见的坑是掩码的维度没有扩展正确,比如 Padding Mask 形状是 [batch, seq_len],而注意力分数是 [batch, num_heads, seq_len, seq_len],需要先扩展成四维再做填充。
4. 从自注意力到多头注意力
4.1 多头注意力的动机
如果只做一个自注意力,模型只能从一种角度建立词与词之间的关系。但语言中的关系是多种多样的:有的位置依赖需要关注句法关系,有的需要关注近义替换,有的需要关注指代关系。只用一套 Q、K、V 映射,会让这些不同关系互相干扰。
多头注意力(Multi-Head Attention)把单个注意力过程复制多份,每一份使用不同的线性映射,形成多个“头”。每个头可以学习不同的注意力模式,最后把所有头的输出拼接起来,再经过一个输出投影层,融合不同子空间的信息。
类比来说,单头注意力像是让一个审查员从头到尾只看一个角度;多头注意力则像同时派出多个审查员,每人侧重不同方面,最后把他们的意见汇总。
4.2 多头拆分的计算流程
假设模型维度 d_model 为 512,头数 num_heads 为 8,那么每个头的维度是 d_k = d_model / num_heads = 64。
计算流程如下:
- 输入 X 分别通过三个线性层 W_Q、W_K、W_V,得到 Q、K、V,形状都是 [batch, seq_len, d_model]。
- 把 Q、K、V 的最后一维拆成 num_heads 份。例如把 [batch, seq_len, 512] 拆成 [batch, seq_len, 8, 64]。
- 调整维度顺序为 [batch, num_heads, seq_len, d_k],相当于把每个头单独拿出来。
- 对每个头独立计算缩放点积注意力,得到 [batch, num_heads, seq_len, d_v]。
- 将头维度转回最后一维,得到 [batch, seq_len, num_heads, d_k]。
- 拼接成 [batch, seq_len, d_model]。
- 经过输出投影 W_O,得到多头注意力的最终结果。
整个过程中,每个头拥有自己的线性映射矩阵。模型训练时,不同头会自动分化,关注不同的特征。
4.3 参数数量和计算量分析
多头注意力并不会显著增加总参数量,因为每个头的维度变小了。
以 d_model=512、num_heads=8 为例:
- Q、K、V 三个线性层都是 [512, 512],总参数为 3 × 512 × 512。
- 输出投影层也是 [512, 512]。
- 总参数约 4 × 512 × 512,与输入维度相关,不受头数影响。
虽然每个头的维度是 64,但所有头合起来仍覆盖完整的 512 维空间。因此,增加头数不会线性增加参数,但会增加内部拆分和拼接的计算。
需要明确的是,多头注意力的计算复杂度仍与序列长度平方相关。每个头计算的注意力矩阵都是 seq_len × seq_len,头数变化不会改变这种平方关系。
4.4 为什么需要输出投影层
每个头独立计算后,得到的是不同子空间中的表示。把多个头的输出拼接起来,维度变成 [batch, seq_len, d_model],但拼接操作只是简单的空间堆叠,没有对头与头之间的信息进行融合和变换。
输出投影层 W_O 的作用就是对拼接结果再做一次线性变换,让不同头的信息能够交互融合,同时也把输出维度恢复到模型内部统一的 d_model。没有这个投影层,多头注意力就退化成一个分块独立注意力,模型的表达能力会明显受限。
在实现时,输出投影层通常就是一个 nn.Linear(d_model, d_model),不要遗漏。
5. 用 PyTorch 从零实现缩放点积注意力与多头注意力
5.1 环境准备与依赖版本
这里使用 PyTorch 实现,安装命令如下:
建议使用 PyTorch 1.13 以上版本,因为从更早版本开始 PyTorch 才完整支持后续要介绍的高效注意力算子。开发环境建议使用 Python 3.9 或 3.10。如果你使用 GPU,还需要按对应 CUDA 版本安装 PyTorch,具体安装命令参考 PyTorch 官网。
学习环境不需要太复杂,CPU 上也能运行本文的小样例。只要矩阵维度不大,CPU 完全够用。下面的代码直接在一个 Python 文件中就能验证。
5.2 构造输入示例
自注意力的输入通常是一组序列向量。为了便于演示,假设 batch_size=2,序列长度 seq_len=4,模型维度 d_model=8,注意力头数 num_heads=2。
这里的 x 可以理解为已经完成词嵌入并加入了位置编码的输入。实际任务中,x 来自 Embedding 层和位置编码层。
5.3 ScaledDotProductAttention 类
先实现单头缩放点积注意力。这个类接收形状为 [batch, heads, seq_len, d_k] 的 Q、K、V,输出注意力结果和注意力权重。
关键点:
k.transpose(-2, -1)转置的是 K 的最后两个维度,也就是把 [seq_len, d_k] 变成 [d_k, seq_len]。q.size(-1)是 d_k,所以缩放因子是每个头的维度,而不是 d_model。- 掩码必须在 softmax 之前完成,否则被掩码的位置仍然是 0 而不是
-inf,还会参与归一化。
5.4 MultiHeadAttention 类
接下来实现多头注意力。这里使用三个独立线性层分别生成 Q、K、V,和一个输出投影层。
这段代码中有两个容易写错的地方:
view的顺序必须是(batch, seq_len, num_heads, d_k),然后transpose(1, 2),不能直接view(batch, num_heads, seq_len, d_k),因为那样会把连续内存按错误顺序拆分。- 经过
transpose后,张量内存可能不连续。在做view之前必须先调用contiguous(),否则 PyTorch 会报错。
5.5 前向传播验证
实例化模型并运行一次前向传播:
预期输出:
输出形状与输入形状一致,说明多头注意力保持了序列维度和模型维度不变。注意力矩阵的第二个维度是头数,所以可以看到每个头都有自己的 4×4 注意力权重矩阵。
5.6 检查输出形状和统计量
除了看形状,还要检查数值是否合理:
对于每一个 head,每一行的注意力权重总和应该接近 1,因为 softmax 对最后一维做了归一化。如果求和结果明显不是 1,说明掩码或 softmax 维度写错了。
也可以检查输出是否有 NaN:
出现 NaN 时,优先检查是否把 -inf 放到了 softmax 之前,以及掩码位置是否覆盖了所有需要屏蔽的位置。
6. 从输出反推注意力机制在做什么:可视化与验证
6.1 如何观察注意力权重
注意力权重矩阵是自注意力中唯一能直观解释模型行为的中间产物。通过可视化,可以看到某个 token 在预测另一个 token 时“看了”哪些位置。
在上一节的实现中,attn 的形状是 [batch, num_heads, seq_len, seq_len]。最后一个维度的第 j 列表示当前位置 i 对位置 j 的注意力权重。可以取出一个样本的一个头来看:
如果要把矩阵画成热力图,可以用 matplotlib:
6.2 一个更直观的文本例子
仅仅使用随机向量,注意力矩阵很难看出规律。为了演示,可以构造一个简单的词向量输入,让模型在训练之前先进行一次前向传播。由于权重是随机初始化的,矩阵大致均匀,但能观察到 softmax 的效果:每行权重虽然不同,但分布偏均匀。
要看到有意义的注意力模式,必须经过训练。这也是初学者容易误解的地方:注意力机制本身不“理解”语义,它只是提供了一种建模能力,真正的语义依赖训练数据和损失函数。
所以文章开头需要的直观例子:如果要把“猫”和“它”关联起来,需要训练数据中出现足够多的指代关系,模型才会在对应位置分配更高权重。
6.3 验证输出和期望是否一致
验证自注意力实现正确,可以从几个角度检查:
- 输出形状是否与输入一致。
- 注意力权重每行归一化。
- 掩码位置的权重是否为 0。
- 当 Q、K、V 完全相同时,输出分布是否符合 softmax 加权平均的预期。
- 将多头注意力输出与 PyTorch 官方
nn.MultiheadAttention对比,作为参考基准。
下面是一个简易对比方法:
这里的官方模块输出与手写版本并非完全一致,因为默认初始化不同,但可以作为维度检查的参考。
6.4 注意力矩阵的解读
注意力矩阵的每一行代表“当前位置 i 关注其他位置的分布”。行中某个值越大,说明位置 i 的表示越依赖位置 j。
在多头注意力中,不同头可能关注不同语义。例如:
- 一个头可能主要关注相邻词,负责局部语法。
- 另一个头可能关注相隔较远的词,负责长距离指代。
- 第三个头可能关注句子结束符,负责全局信息聚合。
不要指望每个头都有清晰可解释的模式。注意力权重只是模型内部的一种软对齐,并不等同于用户可解释的“原因”。这一点在写论文或做分析时要格外谨慎。
7. 常见理解误区与排查思路
7.1 容易与“注意力”混淆的概念
很多文章会把 Attention 和 Self-Attention、Multi-Head Attention 混用。实际上:
- Attention 是通用概念,包括外部记忆、Encoder-Decoder 注意力等。
- Self-Attention 是 Attention 在同一个序列内部的特例。
- Multi-Head Attention 是 Self-Attention 的一种具体实现方式,通过多头并行增强表达能力。
还有一个容易混淆的概念是“通道注意力”,例如 Squeeze-and-Excitation 网络中的 SE 注意力。它计算的是通道维度的权重,与 Transformer 里基于 QKV 的注意力机制不是同一套算法。不要把两者混在一起理解。
7.2 只实现了一个注意力,为什么结果不好
注意力机制只是 Transformer 的一个子层。实际模型里,注意力层前后通常还有残差连接、LayerNorm、前馈网络。如果只把输入送入注意力层就直接输出,模型能力非常有限。
常见的错误是:
- 缺少残差连接,导致深层梯度不稳定。
- 缺少 LayerNorm,导致数值变化过大。
- 缺少前馈网络,导致非线性表达能力不足。
- 学习率设置不合适,导致注意力矩阵无法收敛。
如果手写注意力模块后任务效果不佳,优先检查完整模型结构,而不是怀疑注意力实现本身。
7.3 数值稳定性:为什么除以 sqrt(d_k)
前面提到除以 sqrt(d_k) 是为了控制点积方差。下面用一个小实验说明:
当 d_k 为 64 时,未缩放的 scores 标准差接近 8,缩放后接近 1。如果不缩放,softmax 的输入过大,很多位置的梯度会极其小,训练会变慢甚至停滞。
7.4 掩码实现错误的表现
掩码错误通常有以下现象:
- 注意力权重求和小于 1。
- 输出出现 NaN。
- 模型在训练时 loss 不下降,或预测时看到未来信息导致结果异常。
排查顺序:
- 确认掩码形状。Padding Mask 应为 [batch, 1, 1, seq_len] 或 [batch, 1, seq_len, seq_len],需要能广播到注意力分数形状。
- 确认掩码填充值。通常用 0 表示屏蔽,用 1 表示保留。
- 确认
masked_fill使用的是mask == 0而不是mask == 1。 - 确认 softmax 在掩码之后执行。
如果使用因果掩码,还需要注意上三角矩阵的构造:
输出是一个下三角矩阵,上三角为 0。将上三角位置的分数填充为 -inf,即可阻止当前位置看到未来信息。
7.5 自注意力与 RNN、CNN 的对比
经常被问到的对比表如下:
| 模型 | 并行性 | 长距离依赖 | 计算复杂度 | 位置建模 |
|---|---|---|---|---|
| RNN/LSTM | 差 | 较弱 | O(n) 时间步 | 天然包含位置 |
| CNN | 好 | 依赖堆叠层数 | O(k*n) | 局部位置 |
| Self-Attention | 好 | 强 | O(n²) | 需要位置编码 |
自注意力以更高的计算复杂度换来了更强的并行性和长距离建模能力。后半段学习 Transformer 时,要记住这种取舍不是免费的。
8. 在 Transformer 大框架中的位置及下一步学习路径
8.1 注意力机制之上还有什么模块
注意力机制只是 Transformer 的一个子层。一个完整的 Transformer Encoder Block 通常包含:
- 多头注意力层
- 残差连接
- LayerNorm
- 前馈神经网络(Feed-Forward Network,FFN)
- 第二个残差连接和 LayerNorm
Transformer Decoder 在此基础上还会增加一层交叉注意力(Cross-Attention),用来让解码器关注编码器输出的信息。
学习注意力机制时,如果只停在 QKV 计算上,后面看 Transformer 代码时会觉得松散。建议先理解单头注意力,再理解多头,最后把注意力放回 Block 中理解它如何与残差、LayerNorm、FFN 配合。
8.2 学习顺序建议
按下面的顺序学习会顺畅很多:
- 理解 Embedding 和输入表示。
- 理解位置编码。
- 理解单头缩放点积注意力。
- 理解多头注意力。
- 构建一个完整的 Encoder Block。
- 理解 Mask 在 Encoder 和 Decoder 中的区别。
- 读一遍 PyTorch 官方
nn.Transformer源码。 - 用一个小文本分类任务验证模型。
不要一上来就尝试手写 GPT 或 BERT。先把最小注意力模块跑通,再逐步叠加模块。
8.3 工程落地时的注意事项
在实际工程中,除了理解原理,还需要注意以下问题:
- 使用已优化的注意力算子。PyTorch 2.0 以后提供了
F.scaled_dot_product_attention,它会自动选择内存高效实现,长序列下比手写循环快很多。 - 注意浮点精度。FP32 与 BF16 下,注意力权重的分布会有差异。训练和推理要保持一致的精度。
- 长序列场景要使用稀疏注意力或 FlashAttention,否则显存会随序列长度平方增长。
- 训练时记得设置
model.train(),推理时设置model.eval()。Dropout 在两种模式下行为不同。 - 注意掩码在训练和推理时的区别。训练时可以使用 Teacher Forcing,推理时要逐步生成并维护 KV Cache。
8.4 可复用的实现检查清单
写一个注意力模块时,可以用下面的清单自查:
- [ ] 输入形状是否为 [batch, seq_len, d_model]。
- [ ] d_model 是否能被 num_heads 整除。
- [ ] Q、K、V 是否都是通过独立线性层生成。
- [ ] 拆分多头时是否正确使用
view+transpose,并在还原时使用contiguous。 - [ ] 缩放因子是否为
math.sqrt(d_k)。 - [ ] 掩码是否在 softmax 之前应用。
- [ ] 掩码填充值是否使用
-inf。 - [ ] 注意力权重每行求和是否为 1。
- [ ] 输出形状是否等于输入形状。
- [ ] 是否包含输出投影层。
- [ ] 是否已经加入残差连接和 LayerNorm。
- [ ] 是否区分了训练和推理模式。
把这些检查点过一遍,注意力模块的常见错误基本都能暴露出来。
注意力机制的核心并不神秘:它先通过相似度计算生成动态权重,再把信息按权重聚合。理解了这个过程,再看 Transformer 中的 QKV、多头、掩码和位置编码,就不会被一堆术语吓倒。下一步可以继续学习 Transformer 的完整架构,或者在代码库里实现一个 Encoder Block,把这里的注意力模块接上残差、LayerNorm 和 FFN,你会发现自己已经能看懂真正的大模型基础结构了。