资讯动态

多头注意力机制详解:从原理到 PyTorch 实现

发布时间:2026/8/29 17:59:29 来源:尧图企业网站定制
任何接触过 Transformer 或 BERT 的开发者一定都听过“多头注意力”这个词。初学阶段最让人困惑的往往是既然自注意力已经能够计算每个 token 和其他 token 的关系为什么还要把注意力分成多个“头”分头之后计算过程发生了什么变化多出来的这些头到底带来了什么收益本文将从基础推导到 PyTorch 实现完整拆解“多头”注意力机制的原理并给出可直接运行的代码帮助你彻底理解这一层。1. 从自注意力到多头注意力1.1 自注意力的本质回顾在自然语言处理中每个输入 token词或子词都需要被编码成一个包含上下文语义的向量。自注意力机制做的事情可以概括为一句话让序列中的每个 token 都关注到序列中的其他 token并根据关注程度聚合信息。以句子“我喜欢猫因为它很可爱”为例当模型处理“它”这个词时自注意力机制会计算“它”与“我”“喜欢”“猫”“因为”“很”“可爱”之间的相关度发现“猫”和“可爱”的权重更高于是“它”的最终表示就会包含更多来自“猫”和“可爱”的信息。这就是自注意力的直觉解释。自注意力的计算过程可以用三个向量来描述Query查询向量表示当前 token“想找什么”。Key键向量表示当前 token“我能提供什么”。Value值向量表示当前 token“实际携带的信息”。计算流程是当前 token 的 Query 与所有 token 的 Key 做点积得到注意力分数。将分数除以缩放因子再经过 Softmax 归一化得到注意力权重。用注意力权重对所有 Value 向量加权求和得到当前 token 的输出。用公式表达就是Attention(Q, K, V) Softmax(QK^T / sqrt(d_k)) V其中 d_k 是 Key 向量的维度除以 sqrt(d_k) 是为了防止点积结果过大导致梯度消失。1.2 自注意力层的瓶颈自注意力单独使用是有效的但在实际处理复杂语言现象时它存在一个显著问题一种注意力模式只能捕获一种类型的关系。语言中的关系是多维度的。还是看“我喜欢猫因为它很可爱”这句话语法关系“它”指代“猫”这是一种指代关系。语义关系“可爱”在修饰“猫”这是一种修饰关系。甚至位置关系某些情况下模型需要关注相邻词某些情况下需要关注远处的词。如果只有一个注意力头模型只能学习到一个综合的 Q、K、V 映射。这个映射必须在“指代关系”“修饰关系”“位置关系”等所有任务之间做一个折中。最终的结果可能是每种关系都只能学到一部分精度受限。这就产生了一个需求能不能让模型同时从多个角度去关注序列中的信息多头注意力Multi-Head Attention正是为了解决这个问题而提出的。2. 多头注意力的核心原理2.1 多头注意力做了什么多头注意力的思想非常直接与其只用一个注意力函数不如把 Query、Key、Value 分别映射到多个不同的低维子空间在每一个子空间独立计算注意力最后把所有子空间的结果拼接起来再做一次线性变换。这样做的好处是每个头可以学习到不同的注意力模式。有的头倾向于关注相邻词汇有的头倾向于关注句法依存关系有的头倾向于关注指代关系。多个头并行工作模型就能从不同角度捕获更丰富的语义信息。下面我们从数学角度拆解多头注意力。2.2 多头注意力的数学定义假设输入序列的向量表示为 X维度为 d_model。多头注意力首先将 X 通过三组权重矩阵映射为 Q、K、VQ XW_Q K XW_K V XW_V但多头注意力不是用一个完整的 W_Q、W_K、W_V而是把头数 h 考虑进去。对于第 i 个头Q_i XW_{Q_i} K_i XW_{K_i} V_i XW_{V_i}这里的 W_{Q_i} 的维度是 d_model × d_k其中 d_k d_model / h。也就是说每个头在维度为 d_k 的子空间中计算注意力head_i Attention(Q_i, K_i, V_i) Softmax(Q_i K_i^T / sqrt(d_k)) V_i最后把所有头的输出拼接起来MultiHead(Q, K, V) Concat(head_1, head_2, ..., head_h) W_O其中 W_O 的维度是 d_model × d_model作用是融合不同头的信息并恢复原始维度。2.3 一个具体的维度示例为了更直观地理解维度变化我们看一个具体例子。假设d_model 512头数 h 8每个头的维度 d_k d_model / h 64计算过程中的维度变化如下阶段张量维度输入X(batch_size, seq_len, 512)每个头的 Q/K/VQ_i/K_i/V_i(batch_size, seq_len, 64)缩放点积注意力head_i(batch_size, seq_len, 64)拼接Concat(head_1...head_8)(batch_size, seq_len, 512)输出投影MultiHead(batch_size, seq_len, 512)注意整个多头注意力层的输入和输出维度都是 d_model因此它可以直接作为 Transformer Encoder 或 Decoder 中的一个子层也方便残差连接的实现。2.4 多头注意力与单头注意力的计算量对比这里有一个常被问到的问题多头注意力的计算量是不是比单头大 h 倍答案是否定的。虽然多头注意力的“头数”增多了但每个头的维度被压缩到了 d_model / h总的计算量并没有显著增加。从矩阵乘法的角度看单头注意力是 (seq_len, d_model) 乘以 (d_model, d_model)多头注意力可以理解为把一个大矩阵乘法拆成 h 个小矩阵乘法。在理想情况下参数量与计算量基本持平。这也是多头注意力设计巧妙的地方在不增加整体参数量级的前提下通过分解子空间实现了多角度特征提取。3. 为什么需要“多”个头3.1 单头注意力的局限性为了理解多头注意力的优势我们可以做一个实验用单头注意力训练一个简单的机器翻译模型然后可视化它的注意力权重。你会发现模型被迫在一个头里同时承担多种关系建模任务注意力权重图通常“糊成一团”很难清晰地区分不同语言现象。问题在于单头注意力学到的 Q、K、V 变换是一个整体映射。当语料复杂、任务多样时这种映射的容量不够模型会在不同注意力模式之间互相妥协。3.2 多头的可解释性Google 在《Attention Is All You Need》论文中展示了多头注意力的可视化结果。他们发现有的头在编码英文句子时会学习到语法依存关系比如动词与它的宾语之间有较高的注意力权重。有的头会学习到指代关系比如代词与它指代的名词之间有较高权重。有的头会关注到相邻词捕获局部语境。这些信息在单头注意力中是混叠在一起的不容易被有效分离。多头机制给模型提供了“并行多个子空间”的能力让不同的头可以分化出不同的职责。3.3 多头注意力与感受野的关系从信息聚合的角度看多头注意力还扩大了模型的“表达能力”。每个头可以关注序列中不同范围的信息某些头只关注附近几个 token相当于局部感受野。某些头关注整个序列中距离较远的 token相当于全局感受野。这种混合方式使得模型既能捕获细粒度的局部特征又能捕获长距离依赖这对机器翻译、文本分类、命名实体识别等 NLP 任务都非常重要。4. 使用 PyTorch 实现多头注意力4.1 从零手写多头注意力理论讲完了下面我们来看代码。使用 PyTorch 从零实现一个多头注意力层可以加深对计算过程的理解。首先是核心实现import torch import torch.nn as nn import torch.nn.functional as F class MultiHeadAttention(nn.Module): def __init__(self, d_model512, num_heads8, dropout0.1): super(MultiHeadAttention, self).__init__() assert d_model % num_heads 0, d_model 必须能被 num_heads 整除 self.d_model d_model self.num_heads num_heads self.d_k d_model // num_heads self.W_Q nn.Linear(d_model, d_model) self.W_K nn.Linear(d_model, d_model) self.W_V nn.Linear(d_model, d_model) self.W_O nn.Linear(d_model, d_model) self.dropout nn.Dropout(dropout) def forward(self, query, key, value, maskNone): 参数说明 query/ key/ value: (batch_size, seq_len, d_model) mask: 形状可以是 (batch_size, 1, seq_len, seq_len) 或 (seq_len, seq_len) 返回 output: (batch_size, seq_len, d_model) batch_size query.size(0) # 1. 线性投影并拆分为多头 Q self.W_Q(query).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) K self.W_K(key).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) V self.W_V(value).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) # 2. 缩放点积注意力 scores torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(self.d_k, dtypetorch.float32)) if mask is not None: scores scores.masked_fill(mask 0, -1e9) attention_weights F.softmax(scores, dim-1) attention_weights self.dropout(attention_weights) context torch.matmul(attention_weights, V) # 3. 拼接所有头 context context.transpose(1, 2).contiguous().view( batch_size, -1, self.d_model ) # 4. 输出投影 output self.W_O(context) return output这段代码的实现要点如下d_model是模型的隐藏维度num_heads是注意力头数。通过view方法在d_model维度上切分成num_heads个维度为d_k的子向量。使用transpose将维度交换为 (batch_size, num_heads, seq_len, d_k)这样每个头可以并行计算。mask用于屏蔽某些位置Decoder 的因果注意力会用到。最后通过contiguous()确保内存连续再view拼接头部。4.2 使用 PyTorch 内置实现实际工程中我们不需要手写。PyTorch 提供了现成的nn.MultiheadAttention但在使用时需要注意它与手写版本的 API 差异。下面是一个使用内置模块的例子import torch import torch.nn as nn d_model 512 num_heads 8 seq_len 10 batch_size 2 mha nn.MultiheadAttention(embed_dimd_model, num_headsnum_heads, dropout0.1) # 注意内置模块的输入形状是 (seq_len, batch_size, d_model) query torch.randn(seq_len, batch_size, d_model) key torch.randn(seq_len, batch_size, d_model) value torch.randn(seq_len, batch_size, d_model) attn_output, attn_weights mha(query, key, value) print(输出形状:, attn_output.shape) # (seq_len, batch_size, d_model) print(注意力权重形状:, attn_weights.shape) # (batch_size, seq_len, seq_len)这里容易踩坑nn.MultiheadAttention默认的输入布局是(seq_len, batch, embed_dim)而很多新手按照直觉输入(batch, seq_len, embed_dim)导致维度错误。如果需要使用(batch, seq_len, embed_dim)布局可以设置batch_firstTruemha nn.MultiheadAttention(embed_dimd_model, num_headsnum_heads, batch_firstTrue) query torch.randn(batch_size, seq_len, d_model) attn_output, attn_weights mha(query, query, query) print(batch_first 输出形状:, attn_output.shape) # (batch_size, seq_len, d_model)内置实现还集成了带有 mask、bias、dropout 等功能工程中使用更方便。4.3 多头注意力各头权重可视化为了验证多个头确实学到了不同的注意力模式我们可以提取每个头的注意力权重并可视化。以下代码可以输出每个头的注意力权重分布import matplotlib.pyplot as plt import numpy as np # 假设已经通过 mha 计算得到 attn_weights # attn_weights 形状为 (batch_size, num_heads, seq_len, seq_len) def plot_attention_heads(attn_weights, head_indicesNone): attn attn_weights[0].detach().numpy() # 取第一个 batch num_heads attn.shape[0] if head_indices is None: head_indices range(num_heads) fig, axes plt.subplots(2, 4, figsize(16, 8)) for idx, head in enumerate(head_indices): ax axes[idx // 4][idx % 4] im ax.imshow(attn[head], cmapBlues) ax.set_title(fHead {head}) ax.set_xlabel(Key) ax.set_ylabel(Query) plt.colorbar(im, axax) plt.tight_layout() plt.show() # 先运行一个前向传播 query torch.randn(seq_len, batch_size, d_model) attn_output, attn_weights mha(query, query, query) plot_attention_heads(attn_weights)在实际模型中你通常会看到不同头的注意力图有明显的结构差异这说明多头机制确实在捕获不同类型的关系。5. 多头注意力在 Transformer 中的完整应用5.1 Transformer Encoder 标准结构在 Transformer 中多头注意力并不是孤立使用的。标准的 Encoder 层结构如下多头自注意力子层。残差连接与层归一化。前馈神经网络子层MLP。残差连接与层归一化。伪代码如下class TransformerEncoderLayer(nn.Module): def __init__(self, d_model512, num_heads8, d_ff2048, dropout0.1): super().__init__() self.self_attn nn.MultiheadAttention(d_model, num_heads, dropoutdropout) self.linear1 nn.Linear(d_model, d_ff) self.linear2 nn.Linear(d_ff, d_model) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.dropout1 nn.Dropout(dropout) self.dropout2 nn.Dropout(dropout) def forward(self, src, src_maskNone): # 自注意力 残差 归一化 attn_out, _ self.self_attn(src, src, src) src self.norm1(src self.dropout1(attn_out)) # 前馈网络 残差 归一化 ff_out self.linear2(F.relu(self.linear1(src))) src self.norm2(src self.dropout2(ff_out)) return src这里的关键是残差连接。多头注意力层的输入输出维度相同因此可以直接相加。如果不使用残差连接深层网络会面临梯度消失和退化问题。5.2 Decoder 中的掩码多头注意力在 Decoder 中多头注意力有两种类型自注意力需要加因果掩码确保预测当前位置时只能看到当前位置及之前的位置不能看到未来信息。交叉注意力Query 来自 DecoderKey 和 Value 来自 Encoder 的输出让 Decoder 每个位置都能关注到 Encoder 的所有位置。其中因果掩码的实现方式如下def create_causal_mask(seq_len): 生成因果掩码上三角为 0下三角为 1 mask torch.tril(torch.ones(seq_len, seq_len)).bool() return mask mask create_causal_mask(10) print(mask)输出结果tensor([[ True, False, False, ...], [ True, True, False, ...], [ True, True, True, ...], ...])将这个 mask 传入注意力计算的masked_fill阶段未来位置就会被置为极小的负数经过 Softmax 后权重趋近于 0从而实现因果约束。6. 多头注意力的变体与应用场景6.1 主流模型中的多头注意力多头注意力已经成为深度学习中最重要的基本模块之一。不同模型对头数的选择略有不同模型d_model头数每头维度Transformer base512864Transformer big10241664BERT base7681264BERT large10241664GPT-27681264GPT-31228896128可以看到每头维度 d_k 通常保持在 64 左右头数随着 d_model 的增大而增加。这个设计并非绝对但 64 的维度在多数任务中表现良好是经过大量实验验证的常用配置。6.2 不同类型的注意力机制对比在网络资料中经常会看到很多注意力变体这里简单区分一下自注意力Self-AttentionQuery、Key、Value 都来自同一个序列用于建模序列内部的关系。多头自注意力Multi-Head Self-Attention将自注意力分成多个头并行计算是 Transformer 的核心组件。因果自注意力Causal Self-Attention在自注意力基础上加因果掩码只能看到当前位置及之前的位置用于自回归生成模型。交叉注意力Cross-AttentionQuery 来自一个序列Key、Value 来自另一个序列用于 Encoder-Decoder 模型。通道注意力Squeeze-and-Excitation Attention常用于图像分类通过全局池化建模通道之间的依赖关系是 SENet 的核心模块。时序注意力可以理解为一类关注时间维度重要性的注意力机制在时序预测任务中比较常见。需要说明的是本文讨论的多头注意力属于“自注意力”家族它解决的是序列内部建模问题。理解了它再去看其他注意力变体会容易很多。7. 常见问题与排查思路7.1 Query、Key、Value 维度不一致怎么办nn.MultiheadAttention的kdim和vdim参数允许 Key、Value 的维度与 Query 不同。例如在交叉注意力中Encoder 输出维度是 512但 Decoder 的 Query 维度是 768此时可以设置mha nn.MultiheadAttention( embed_dim768, num_heads8, kdim512, vdim512 )这种情况在很多跨模态模型中很常见。7.2 d_model 无法被 num_heads 整除手写实现中如果d_model % num_heads ! 0会直接抛异常。解决办法是调整num_heads保证d_model是num_heads的整数倍。例如 d_model512 可以选择 8、16 头d_model768 可以选择 12 头。也可以使用nn.MultiheadAttention它内部会对维度做兼容处理但最稳妥的方式仍然是保证整除。7.3 训练时损失不下降注意力权重过于均匀这种情况通常是以下原因导致的问题现象常见原因解决思路注意力权重接近均匀分布模型尚未充分训练或学习率过大降低学习率增加训练步数多头中的某些头一直学到相同模式初始化问题或参数共享问题使用不同的随机种子检查权重初始化输出出现 NaN注意力分数过大导致 Softmax 溢出检查是否忘记除以 sqrt(d_k)因果掩码未生效生成结果泄露未来信息mask 传递或形状错误打印 mask检查 mask 是否被正确传入7.4 内置模块与手写实现结果不一致nn.MultiheadAttention在计算时默认会对输入做一种特殊处理它将嵌入维度在内部拆分时保持顺序并且部分版本会加上一个可选的 bias 项。如果你用手写实现做对照实验需要确认是否包含 bias、dropout 等细节。一般来说两者数学上等价但存在浮点误差差异在 1e-5 以下属于正常现象。8. 最佳实践与工程建议8.1 头数的选择策略在实际项目中头数并不是越大越好。头数太多会导致每个头的维度太小单个头的表达能力不足头数太少又无法充分捕获多角度关系。经验值如下d_model 512 时选择 8 个头。d_model 768 时选择 12 个头。每头维度 d_k 保持在 64 左右是常见选择。如果你需要调整建议优先保持 d_k 在 32 到 128 之间同时关注训练曲线。8.2 先检查 mask再检查模型很多 Transformer 模型的 bug 都出在 mask 上。建议在训练前单独打印 mask 并进行单元测试mask create_causal_mask(10) assert torch.all(mask.diagonal()), 对角线位置必须为 True assert not torch.any(torch.triu(mask, diagonal1)), 上三角必须全为 False print(mask 检查通过)这种简单的断言检查可以避免大量隐蔽的训练异常。8.3 性能优化与显存控制多头注意力的时间复杂度是 O(n^2 d)其中 n 是序列长度。当序列很长时显存占用会急剧上升。工程上有几种常见优化手段使用 FlashAttention 等融合实现减少中间矩阵的显存占用。在超长序列场景下使用稀疏注意力只计算局部窗口内的注意力分数。如果使用 PyTorch 内置多头注意力保持输入为batch_firstTrue可以减少显式转置操作。推理阶段可以缓存历史 Key、Value避免重复计算。8.4 实现建议手写实现适合学习工程中使用内置模块更适合。推荐路线是先手写一遍理解原理再切换到nn.MultiheadAttention或大模型框架中的成熟实现。这样既能深入底层又能保证工程代码的稳定性和性能。8.5 可解释性分析在调试文本生成、翻译等模型时可视化注意力权重是一个非常有效的诊断手段。可以观察注意力权重是否集中在合理的 token 上。是否存在注意力“塌缩”所有 token 关注同一个位置。不同头是否显示出不同的关注模式。这些观察能帮助你判断模型是否学习到有效的语义关系而不是“糊弄”着过拟合训练数据。9. 总结与下一步学习方向多头注意力机制是 Transformer 架构中最核心的组件之一。本文从自注意力的基础出发解释了为什么要引入多“头”给出了数学定义、完整维度变化和 PyTorch 手写代码并结合 Transformer Encoder/Decoder 展示了实际应用最后整理了常见问题与工程建议。理解多头注意力之后你可以继续探索以下方向多头注意力在 BERT、GPT 系列模型中的具体配置差异。FlashAttention 如何在不改变数学结果的前提下大幅提升训练速度。各种注意力变体如分组查询注意力、滑动窗口注意力等。注意力可视化工具的使用如 BertViz直观观察不同头的关注模式。如果本文对你有帮助可以收藏备用。动手实践永远是理解深度学习模块的最好方式建议你亲自运行一遍文中的代码修改头数和 d_model观察输出的变化。写得多了、跑得多了你对多头注意力的理解会越来越深。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价