Transformer自注意力机制详解:从原理到PyTorch实现

Transformer自注意力机制详解:从原理到PyTorch实现 这节内容是深度学习中 Transformer 最核心的组件自注意力机制Self-Attention。不管你看的是 GPT、BERT、LLaMA还是 Vision Transformer、Swin Transformer它们的底层都有一个共同模块——自注意力。这个模块做的事情可以用一句话概括让序列中的每一个位置通过动态计算权重与序列中其他所有位置交换信息。很多人学到递归神经网络RNN和卷积神经网络CNN之后再看 Transformer 会觉得像黑盒。实际上Transformer 的骨架并不复杂最核心的就是自注意力层、前馈网络和归一化层。本文不碰完整结构只把 7.1.1 小节的内容讲透自注意力是什么、数学上怎么算、为什么要用 Q/K/V、用 PyTorch 怎么写、因果掩码是怎么回事、多头注意力又是如何扩展的。文章会给出可以直接运行的 PyTorch 代码。建议你手动跑一遍先理解计算流程再看注意力矩阵长什么样最后把代码改成多头版本。这样你对自注意力的理解会比单纯背公式深刻得多。AI Study 深度学习笔记Transformer 7.1.1 自注意力机制1. 核心认知速览先给出一张速览表把自注意力机制的关键信息列清楚。后面所有分析都围绕这张表展开。项目说明所属范畴深度学习 / 自然语言处理 / Transformer 基础组件核心作用建模序列内部任意两个位置之间的依赖关系输入形式Token 嵌入序列形状为[batch_size, seq_len, embed_dim]输出形式与输入形状相同的增强表示序列核心计算Query、Key、Value 投影 点积注意力 Softmax 归一化 加权求和关键优势任意位置之间一步直达支持并行计算主要瓶颈计算复杂度为O(n²)序列越长开销越大推荐学习方式先用 PyTorch 从零实现单头自注意力再扩展到多头需要前置知识线性代数、Softmax、Word Embedding、深度学习基础常见变体因果自注意力、多头注意力、FlashAttention、线性注意力上面这张表里最需要长期记住的是核心计算流程和复杂度瓶颈。Transformer 之所以能成为大模型的基础架构本质原因是自注意力在长距离依赖建模上比 RNN 更直接在灵活性和并行度上比 CNN 更有优势。我们下面逐步展开。2. 为什么需要自注意力2.1 RNN 与 CNN 的局限性在 Transformer 出现之前序列建模主要靠 RNN 和 CNN。RNN 把输入从左到右逐个读入借助隐藏状态保存历史信息。这种顺序处理方式的缺点是当前时刻的输出必须等前面所有时刻处理完才能计算训练阶段难以并行长序列中信息经过多步传递后容易衰减或丢失。LSTM 和 GRU 通过门控机制缓解了梯度消失问题但长距离依赖仍然需要逐位置传递路径长度随距离增长。CNN 的做法则是通过卷积核提取局部特征用堆叠层数扩大感受野。卷积的特点是局部连接核大小为 3 时每个输出只能看到相邻 3 个位置的输入。要建立远距离依赖就必须叠加很多层或者使用空洞卷积、大卷积核等手段。也就是说CNN 的建模能力依赖层数路径变长了优化难度也随之上升。2.2 自注意力的思路自注意力换了一个角度每个输出位置不是只看局部而是直接与序列中所有位置计算相关性然后加权聚合信息。相关性权重是动态计算的不是固定参数。输入序列长度为 4 时第一个位置可以直接关注第二、第三、第四个位置不需要经过中间节点接力传递。用一句话概括自注意力让任意两个位置之间的交互路径长度变为 1。这个特性让长距离依赖建模变得非常直接同时每个位置的输出可以并行计算训练效率大幅提升。Google 在 2017 年发表的论文《Attention Is All You Need》正式提出 Transformer 架构把自注意力作为核心模块从此开启了大规模预训练模型时代。理解到这里你会明白一个关键点自注意力不是某种辅助机制它本身就是信息交换的主力。后面我们会把它的计算流程拆成四个步骤。3. 自注意力机制数学原理3.1 输入与投影假设输入是一个形状为[seq_len, embed_dim]的矩阵 X每一行代表一个位置的向量表示。自注意力首先通过三个可学习的线性层把 X 投影成三组向量Query查询向量Q X W_QKey键向量K X W_KValue值向量V X W_V其中 W_Q、W_K、W_V 都是[embed_dim, d_k]或[embed_dim, d_v]的可学习权重矩阵。实际项目中通常取d_k d_v embed_dim。Query 可以理解为我要找什么Key 是我是什么Value 是我能提供什么内容。3.2 注意力分数要判断位置 i 应该关注位置 j 多少天然做法是计算 Query_i 与 Key_j 的相似度。最常见的是点积score(i, j) Q_i · K_j将所有位置的分数堆叠起来得到矩阵QK^T形状为[seq_len, seq_len]。分数越大表示两个位置越相关。3.3 缩放因子直接使用点积存在一个问题当向量维度 d_k 较大时点积的方差会变大导致部分分数进入 Softmax 的饱和区梯度会变得很小。解决办法是除以√d_kAttention(Q, K, V) softmax(QK^T / √d_k) V这个缩放因子就是整个公式中最容易被忽略、但对训练稳定性至关重要的部分。3.4 Softmax 归一化与加权求和得到缩放后的分数矩阵后对每一行做 Softmax使每一行的权重和为 1。这样注意力权重矩阵就形成了概率分布位置 i 对位置 j 的注意力权重等于 softmax 的第 (i, j) 项。最后用权重矩阵乘 V得到加权求和后的输出。最终输出的第 i 行是 V 中所有行的加权平均权重由注意力分数决定。这就是自注意力的完整计算过程。4. 用 PyTorch 从零实现自注意力4.1 核心函数实现先写缩放点积注意力的核心函数import torch import torch.nn as nn import torch.nn.functional as F def scaled_dot_product_attention(query, key, value): 缩放点积注意力核心计算 参数: query: [batch_size, seq_len, d_k] key: [batch_size, seq_len, d_k] value: [batch_size, seq_len, d_v] 返回: output: [batch_size, seq_len, d_v] attn_weights: [batch_size, seq_len, seq_len] d_k query.size(-1) scores torch.matmul(query, key.transpose(-2, -1)) scores scores / (d_k ** 0.5) attn_weights F.softmax(scores, dim-1) output torch.matmul(attn_weights, value) return output, attn_weights4.2 单头自注意力模块接下来把 Q/K/V 投影封装成一个模块class SelfAttention(nn.Module): def __init__(self, embed_dim): super().__init__() self.embed_dim embed_dim self.query nn.Linear(embed_dim, embed_dim) self.key nn.Linear(embed_dim, embed_dim) self.value nn.Linear(embed_dim, embed_dim) def forward(self, x): x: [batch_size, seq_len, embed_dim] Q self.query(x) K self.key(x) V self.value(x) output, attn_weights scaled_dot_product_attention(Q, K, V) return output, attn_weights4.3 运行验证用一个小输入验证模块能否正常运行torch.manual_seed(42) batch_size, seq_len, embed_dim 2, 4, 8 x torch.randn(batch_size, seq_len, embed_dim) self_attn SelfAttention(embed_dim) output, attn_weights self_attn(x) print(输入形状:, x.shape) print(输出形状:, output.shape) print(注意力矩阵形状:, attn_weights.shape) print(每行注意力权重之和:, attn_weights.sum(dim-1))预期输出关键点输出形状与输入形状一致仍为[2, 4, 8]。注意力矩阵形状为[2, 4, 4]。每行注意力权重之和约为 1因为对每一行做了 Softmax。这个验证非常直观。你可以把seq_len改大一些比如 32 或 64观察运行时间的变化。你会发现序列长度对自注意力计算量的影响不是线性的这正是后面要讲的复杂度问题。5. 因果自注意力与掩码5.1 为什么需要因果掩码在 Transformer 的解码器Decoder中生成任务要求当前位置只能看到当前位置以及之前位置的信息不能提前看到未来。否则模型在预测下一个词时偷看答案训练就失去了意义。这种设计叫因果自注意力Causal Self-Attention也叫掩码自注意力。实现方式非常简单在 Softmax 之前把未来位置的分数替换为负无穷。负无穷经过 Softmax 后权重趋近于 0相当于完全被屏蔽。5.2 代码实现def causal_attention(query, key, value): 因果自注意力当前位置只能关注自身及之前的位置 batch_size, seq_len, d_k query.shape scores torch.matmul(query, key.transpose(-2, -1)) scores scores / (d_k ** 0.5) # 构造下三角掩码形状 [seq_len, seq_len] mask torch.tril(torch.ones(seq_len, seq_len)).bool() # 把未来位置的分数置为负无穷 scores scores.masked_fill(~mask, float(-inf)) attn_weights F.softmax(scores, dim-1) output torch.matmul(attn_weights, value) return output, attn_weights调用方式和前面一样。你可以在拿到注意力矩阵后打印第二行会发现第二个位置之后的所有权重都是 0只有第一个位置和第二个位置有权重。这就是因果自注意力和普通自注意力最直观的区别。5.3 常见实现细节实际的大模型实现中掩码矩阵往往提前算好并缓存避免每次前向都生成一次。同时为了方便批量处理不同长度的序列还会叠加一个 Padding Mask把无效位置也一起屏蔽。两套掩码组合使用一个是防止看到未来的三角形掩码一个是防止关注到无效填充位的掩码。理解这两者的区别对阅读大模型源码很有帮助。6. 多头注意力机制6.1 为什么需要多头单头自注意力虽然能建模全局依赖但所有信息都压缩在一套 Q/K/V 投影里。实际语义关系非常复杂比如一个词可能在语法关系上依赖前一个词在指代关系上又关联到很远的某个名词这两类关系需要不同的注意力模式来捕捉。多头注意力Multi-Head Attention的思路是把embed_dim拆分成h个子空间每个头在独立的低维空间计算自注意力最后把 h 个头的输出拼接起来再过一层线性投影。这样模型拥有了多套关注方式可以同时捕获不同类型的关系。6.2 公式与维度假设embed_dim 512num_heads 8那么每个头的维度是head_dim 512 / 8 64。输入 X 先通过一个线性层投影到3 * embed_dim的维度然后拆分成 Q、K、V各自再按head_dim切成 8 组。这个过程在代码里通常用 reshape 和 permute 配合完成。6.3 PyTorch 多头注意力实现下面给出一个完整的、带因果掩码的多头注意力实现。这个代码可以同时看到多头拆分、因果掩码、输出拼接和最终投影class MultiHeadAttention(nn.Module): def __init__(self, embed_dim, num_heads): super().__init__() assert embed_dim % num_heads 0, embed_dim 必须能被 num_heads 整除 self.embed_dim embed_dim self.num_heads num_heads self.head_dim embed_dim // num_heads self.qkv nn.Linear(embed_dim, 3 * embed_dim) self.proj nn.Linear(embed_dim, embed_dim) def forward(self, x): batch_size, seq_len, embed_dim x.shape # 1. 一步投影得到 QKV再拆分 qkv self.qkv(x) # [batch, seq_len, 3 * embed_dim] qkv qkv.reshape(batch_size, seq_len, 3, self.num_heads, self.head_dim) qkv qkv.permute(2, 0, 3, 1, 4) # [3, batch, heads, seq_len, head_dim] query, key, value qkv[0], qkv[1], qkv[2] # 2. 缩放点积注意力 scores torch.einsum(b h q d, b h k d - b h q k, query, key) scores scores / (self.head_dim ** 0.5) # 3. 因果掩码按需要选择是否开启 mask torch.tril(torch.ones(seq_len, seq_len)).bool() scores scores.masked_fill(~mask, float(-inf)) attn_weights F.softmax(scores, dim-1) # 4. 加权求和einsum 自动处理 batch 和 heads 维度 context torch.einsum(b h q k, b h k d - b h q d, attn_weights, value) # 5. 拼接所有头恢复 [batch, seq_len, embed_dim] context context.reshape(batch_size, seq_len, embed_dim) # 6. 最后的输出投影 output self.proj(context) return output, attn_weights测试方式与前面一致torch.manual_seed(0) batch_size, seq_len, embed_dim 2, 8, 16 num_heads 4 x torch.randn(batch_size, seq_len, embed_dim) mha MultiHeadAttention(embed_dim, num_heads) output, attn_weights mha(x) print(多头注意力输出形状:, output.shape) print(注意力权重形状:, attn_weights.shape) # [batch, heads, seq_len, seq_len]需要注意attn_weights的维度比单头时多了一维变成了[batch, heads, seq_len, seq_len]。这意味着你可以用attn_weights[0][2]查看第 0 个样本第 2 个头的注意力分布观察不同头关注的位置是否不同。7. 自注意力的性能观察与工程要点7.1 计算复杂度自注意力需要计算QK^T形状是[seq_len, seq_len]。序列长度为 n 时计算量约为 O(n²)。这里的 n 是序列长度不是模型宽度。所以当文本从 512 增长到 2048自注意力计算量会增长 16 倍这就是为什么长文本处理是大模型应用的难点。反观 RNN复杂度是 O(n)但无法并行。CNN 的复杂度与卷积核大小和层数相关。自注意力用更大的计算量换来了任意位置直接交互和可并行两个优势。工程上的优化方向主要是两点减少注意力计算的路径长文本稀疏化或者改变计算顺序FlashAttention 等融合算子。7.2 显存占用观察方法在本地实验时观察显存占用是一堂必修课。PyTorch 中可以用以下方法# GPU 环境下 import torch def print_gpu_memory(): print(当前显存占用:, torch.cuda.memory_allocated() / 1024**2, MB) print(显存缓存量:, torch.cuda.memory_reserved() / 1024**2, MB)在 CPU 上运行自注意力代码时显存占用可以忽略不计。真正需要关注显存的场景是把模型放到 GPU 上训练或推理。可以测试不同seq_len下的显存增长趋势512、1024、2048你会发现显存增长明显快于线性增长这就是平方复杂度在显存上的体现。7.3 数值精度的影响深度学习中浮点数格式的选择会影响显存、速度和精度。常见的格式有 FP32、FP16、BF16、TF32。FP32 精度高但占用显存多FP16 占用减半但动态范围有限容易出现梯度溢出BF16 与 FP32 动态范围一致适合训练TF32 是某些 GPU 上加速矩阵乘法的中间格式。在自注意力的QK^T和 Softmax 计算中FP16 下很容易出现数值不稳定所以很多框架在注意力计算中会混合精度处理。理解这一点能帮你排查训练时 loss 变成 NaN 的问题。7.4 降低计算开销的思路处理长序列时的常见优化方向包括局部注意力只关注附近的窗口、稀疏注意力按固定模式跳过部分位置、线性注意力用核函数近似点积、FlashAttention重排计算顺序并分块处理。这些方法都是以自注意力为基准做的工程优化。先掌握基础版本再看优化版本会容易得多。8. 常见问题与易混淆点8.1 自注意力与普通注意力的区别普通注意力比如编码器-解码器注意力中Query 来自解码器Key 和 Value 来自编码器是两个不同序列之间的信息交换。自注意力中Q、K、V 全部来自同一个序列所以叫自注意力。见到 Cross-Attention 时要能分清哪个序列提供 Query哪个序列提供 Key/Value。8.2 Q/K/V 到底是什么Q、K、V 是输入通过三个不同的线性层投影出来的向量。可以把它们想象成Query 负责提问Key 负责匹配Value 负责提供内容。相似度由 Query 和 Key 的点积决定内容的加权求和由 Value 承担。三者缺一不可。8.3 为什么需要缩放因子不缩放的话当 d_k 很大时点积的方差会变大。Softmax 对输入数值非常敏感过大的输入会让概率分布趋于极值导致梯度非常小训练几乎走不动。除以 √d_k 之后点积的方差被拉回合理范围训练稳定很多。8.4 多头注意力与单头注意力的区别单头注意力把整个embed_dim看作一个空间多头注意力把空间切成多个子空间每组独立计算。切分不是把输入切成几段而是用线性层投影到不同子空间后分别计算。这样做的好处是模型可以在多个表示子空间中同时捕获不同类型的依赖关系。8.5 注意力权重与全连接权重的区别全连接层的权重是训练后固定的参数与输入无关。注意力权重是根据当前输入动态计算出来的同一个模型处理不同输入时注意力矩阵不同。正因为这种动态特性自注意力可以建模复杂的上下文关系。8.6 常见问题排查表问题现象可能原因排查方式解决方案注意力矩阵全为均匀分布初始化导致分数太接近检查 QKV 投影的参数初始化调整初始化方式或增大训练步数训练时 Loss 变成 NaNFP16 溢出或学习率过大检查梯度数值和 Loss 曲线使用 BF16、降低学习率、加 Gradient Clipping多头注意力的 head_dim 设置出错embed_dim 无法被 num_heads 整除打印形状逐层检查确保 embed_dim 与 num_heads 匹配因果自注意力输出不满足因果性掩码没在 Softmax 之前应用打印注意力矩阵观察非零位置把负无穷填充放到 Softmax 之前序列变长后显存暴涨O(n²) 复杂度用不同 seq_len 测试显存增长使用稀疏注意力或 FlashAttention 等优化方案注意力权重行和不为 1拿到的是未归一化的分数确认是否经过 Softmax在 Softmax 之后取权重9. 最佳实践与学习建议9.1 从最小样例开始调试先设置batch_size2, seq_len4, embed_dim8这样的小参数打印每一步的张量形状。确认形状对齐后再调大参数。自注意力实现中最常见的错误是转置维度写错小规模调试能快速定位。9.2 可视化注意力权重拿到注意力矩阵后用matplotlib画一张热力图是直观理解自注意力的好办法。横轴为 Key 的位置纵轴为 Query 的位置颜色深浅代表权重大小。你可能会看到对角线附近的权重较高这表示模型倾向于关注自身位置也可能看到某些词会强烈关注远处的重要词。9.3 理解 einsum 写法很多开源项目使用torch.einsum表示注意力计算。einsum 用字符串描述张量运算维度比一堆 reshape 和 transpose 更可读。建议掌握b h q d, b h k d - b h q k这类写法的含义阅读源码时会快很多。9.4 递进式学习路径自注意力只是 Transformer 的其中一个组件。建议按下面顺序递进手写单头自注意力代码理解矩阵运算过程。实现因果掩码理解生成任务中的信息屏蔽。实现多头注意力理解子空间拆分。阅读原版 Transformer 论文的编码器和解码器结构。最后阅读 HuggingFace 中 GPT 或 BERT 的源码把代码与理论对应起来。9.5 注意版权与合规如果是基于开源模型做二次开发要仔细阅读模型的许可证如果自己训练数据来自网络文本、图片或语音需要确认数据来源是否合规如果涉及人脸、声音等敏感信息必须取得授权。这些在动手做实际项目时非常重要。10. 总结与下一步这次把自注意力的核心原理和代码实现全部串起来了。从 Q/K/V 投影、缩放点积、Softmax 归一化到因果掩码和多头注意力每一步都是 Transformer 模型的地基。最值得先验证的是那个最简单的单头自注意力代码把输入改成不同长度打印注意力矩阵观察每一行权重和是否为 1确认计算流程完全正确。最容易踩的坑有两个一是忘记缩放因子导致训练不稳定二是掩码位置放错把负无穷放到了 Softmax 之后导致权重计算错误。建议在代码里把这两个地方做成断言或打印检查形成习惯。下一步可以沿着两个方向继续深入一是把编码器Encoder中的自注意力、前馈网络、层归一化和残差连接组合起来搭建完整的 Transformer 编码器二是对照 GPT 源码看因果自注意力如何与大模型训练流程配合。到这一步你对深度学习 Transformer 的理解就已经超过大部分只看论文的人了。如果正在用 PyTorch 做 NLP 项目建议把这篇文章里的多头注意力模块保存下来作为工具函数后面写模型时可以直接复用。