Transformer核心:多头注意力机制原理与PyTorch实现
2026/8/30 6:48:04 网站建设 项目流程

这次我们来看 Transformers 里最核心,也最容易在阅读源码时绕晕的一个模块:多头注意力。它属于 Transformer 章节 7.1.2 的内容,前接 7.1 注意力机制基础,后面直接通向 BERT、GPT、ViT 这些实际模型。很多同学在看结构图时觉得 Q、K、V 三条分支懂了,但一旦落到 PyTorch 代码里,就不清楚每个张量到底是几维、为什么要算点积、为什么拆成多头后又要拼回去。这篇文章要做的,就是把这条链路彻底打通:先拆注意力机制的公式和直觉,再讲多头注意力到底在做什么,然后给出一版可运行的 PyTorch 从零实现,并和官方nn.MultiheadAttention做对比验证。

先说结论。多头注意力的核心价值,不是简单增加参数量,而是把注意力计算切到多个子空间并行执行,让模型能够同时关注不同位置的多种依赖关系。从工程角度看,它只是在缩放点积注意力前后做几次 reshape 和线性变换,计算量没有本质增加,但表达能力更强。全程不需要高端显卡,CPU 就能跑通。

文章按这个顺序展开:第一部分给出多头注意力的核心认知速览;第二部分从缩放点积注意力讲起,这是所有后续内容的基础;第三部分拆解多头注意力的完整流程;第四部分给环境准备和 PyTorch 实现;第五部分用代码验证维度、掩码和注意力可视化;第六部分看它在真实模型中的应用;最后是常见问题排查和最佳实践。

1. 多头注意力核心认知速览

先给一张速览表,把多头注意力这个模块的关键信息放在最前面,方便你判断这篇文章是否值得往下读。

维度说明
所属模块Transformer 编码器和解码器的核心子层
核心公式Attention(Q,K,V) = softmax(QKᵀ / √d_k)V
多头计算形式MultiHead(Q,K,V) = Concat(head₁, …, head_h)W^O
核心思想将 d_model 维空间切成 h 个 d_k 维子空间,并行计算注意力
可训练参数Q/K/V 三个线性层 + 输出线性层,共 4 组权重
典型 head 数Transformer 论文默认 8;BERT-base 为 12;GPT-2 为 12
硬件门槛理解原理无需 GPU;小规模验证 CPU 可运行
与单头区别多头能同时捕获多种关系,单头只能做一种加权聚合

这里先解释几个常见符号,后面代码和公式都会用到:

  • d_model:输入向量的维度,也是每个 token 的表示维度。
  • n_head:多头数量。
  • d_k:每个 head 的查询/键向量维度,通常d_k = d_model / n_head
  • d_v:每个 head 的值向量维度,通常在标准实现中d_v = d_k

理解多头注意力不需要先读完整篇 Transformer,只需要知道它接受一个形状为(batch_size, seq_len, d_model)的张量,经过 Q、K、V 三个线性映射和若干次 reshape 后,输出和输入同形状的张量,同时内部还产出一组注意力权重。

2. 注意力机制基础:从加权求和说起

2.1 Query、Key、Value 的直觉

注意力机制可以这样理解:把一份文本切成长度为 n 的 token 序列后,每个 token 被表示成一个向量。当我们处理第 i 个 token 时,希望模型能动态地决定“应该重点关注序列里的哪些位置”。

这里的 Query 可以理解为当前查询向量,Key 是其他所有位置的索引向量,Value 是其他位置的内容向量。模型先用当前 Query 和所有 Key 做相似度计算,得到一组权重,再把这些权重应用到 Value 上,最终得到当前 token 的上下文表示。这就是一次“加权求和”的过程。

Query 来自“我要查什么”,Key 是“别人能提供什么索引”,Value 是“别人实际提供的内容”。注意力机制就是在给定 Query 的情况下,从所有 Key 中找出相关度,然后按相关度聚合 Value。

2.2 缩放点积注意力的数学形式

论文《Attention Is All You Need》中给出的注意力函数是缩放点积注意力,公式如下:

Attention(Q, K, V) = softmax(QKᵀ / √d_k)V

这个公式看起来简单,但每个矩阵的维度必须清楚:

  • Q 的形状是(batch_size, seq_len_q, d_k)
  • K 的形状是(batch_size, seq_len_k, d_k)
  • V 的形状是(batch_size, seq_len_k, d_v)
  • QKᵀ 计算后得到(batch_size, seq_len_q, seq_len_k),表示每个 Query 和每个 Key 之间的相似度。
  • softmax 作用在最后一个维度上,让每一行的注意力权重之和为 1。
  • 最后乘 V,得到(batch_size, seq_len_q, d_v)

在自注意力场景里,seq_len_q 和 seq_len_k 相等,Q、K、V 都来自同一个输入序列。在编码器-解码器注意力场景里,Q 来自解码器,K 和 V 来自编码器输出。

2.3 为什么要除以根号 d_k

缩放因子 √d_k 是公式里最容易忽略但最关键的部分。如果不做缩放,两个 d_k 维向量的点积结果会随着维度增加而变大。当 d_k 较大时,点积结果的数值会很大,softmax 的输入进入梯度饱和区,表现为梯度非常小,训练不稳定。除以 √d_k 后,点积的方差被拉回 1 附近,softmax 的梯度能保持在合理范围。

从实现角度看,这个缩放不需要额外学习参数,只是一个常数操作。但它在训练稳定性上非常重要。手写多头注意力时,容易把这个缩放漏掉,导致训练时 loss 不下降,这是常见问题之一。

2.4 单头注意力的局限

如果只做一次注意力计算,模型只能学到一组加权模式。但一句话里往往同时存在多种关系:相邻词之间的语法关系、远距离的指代关系、否定词的作用范围、句法结构中的父子节点关系。单头注意力只能把这些关系混合在同一个加权平均里,无法分别建模。

多头注意力要解决的,正是这个问题:让不同的头去学习不同类型的依赖关系,最终把多个子空间的信息拼接起来,交给输出投影层融合。这也是为什么不是“多一份参数”这么简单,而是“多一份表达能力”。

3. 多头注意力的原理拆解

3.1 整体思路

多头注意力可以拆成五个步骤:

  1. 对输入做 Q、K、V 三个线性投影。
  2. 把 Q、K、V 按头数 reshap 并转置,拆成 n_head 个子空间。
  3. 在每个子空间里独立执行缩放点积注意力。
  4. 把所有头的输出拼接回d_model维。
  5. 过输出线性投影 W^O。

整体公式为:

MultiHead(Q,K,V) = Concat(head₁, …, head_h)W^O

其中每个头为:

head_i = Attention(QW_i^Q, KW_i^K, VW_i^V)

注意 Q、K、V 并不是一开始就拆好的,而是先通过三组线性层映射到同样维度,再在计算时切成多段。这个“先映射再切分”的做法,在工程实现上非常高效。

3.2 张量形状变化

为了不抽象,这里把张量形状变化列成一张表,以batch_size=2seq_len=10d_model=512n_head=8为例,每个头的维度d_k=d_v=64

步骤输入形状输出形状
Q/K/V 线性投影(2, 10, 512)(2, 10, 512)
拆分为多头(2, 10, 512)(2, 8, 10, 64)
每个头单独注意力(2, 8, 10, 64)(2, 8, 10, 64)
拼接所有头(2, 8, 10, 64)(2, 10, 512)
输出投影(2, 10, 512)(2, 10, 512)

reshape 的细节要特别注意。一个(2, 10, 512)的张量先变成(2, 10, 8, 64),然后用transpose(1, 2)变成(2, 8, 10, 64)。这样才能保证每个头都看到完整的序列,而不是把序列拆成多段。

3.3 为什么有效:不同头关注不同关系

论文中提到,训练完毕后观察不同头的注意力权重,会发现不同头分布在不同区域:有些头主要关注相邻词,有些头关注远距离依赖,还有些头关注特定的语法关系。这说明多头并不是简单的重复,而是模型在训练中自动把不同的“关系查找任务”分配给了不同的子空间。

不过也要说明,并非每个头都一定学到可解释的模式,有部分头可能互相冗余。这也催生了后续的 MQA(Multi-Query Attention)和 GQA(Grouped Query Attention)等优化,它们核心思想都是减少 Key 和 Value 的冗余头数,降低推理显存和带宽开销。理解标准多头注意力,是理解这些优化方案的前提。

4. 环境准备与 PyTorch 从零实现

4.1 环境准备

多头注意力是理论模块,不需要特殊硬件。建议用 Python 3.8 以上版本和 PyTorch 1.10 或 2.x。安装命令如下,具体 PyTorch 版本请以官方安装页为准:

pip install torch numpy matplotlib

安装完成后,可以用下面的命令确认 PyTorch 是否可用:

python -c "import torch; print(torch.__version__)"

如果你的环境里已经有 PyTorch,可以直接跳过安装步骤。下面所有实验在 CPU 上就能运行。

4.2 手写多头注意力模块

从零实现一版多头注意力,核心代码并不长。这里给出一个适合学习的最小实现,没有封装太多复杂细节,方便对照公式。

import math import torch import torch.nn as nn import torch.nn.functional as F class MultiHeadAttention(nn.Module): """ 从零实现的多头注意力模块。 d_model: Transformer 模型的宽度 n_head: 头数 """ def __init__(self, d_model, n_head, dropout=0.1): super().__init__() assert d_model % n_head == 0, "d_model 必须能被 n_head 整除" self.d_model = d_model self.n_head = n_head self.d_k = d_model // n_head self.d_v = d_model // n_head 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, mask=None): batch_size = query.size(0) # 1. 线性投影 Q = self.w_q(query) # (batch, seq_len, d_model) K = self.w_k(key) # (batch, seq_len, d_model) V = self.w_v(value) # (batch, seq_len, d_model) # 2. 拆分多头 # 先 view 成 (batch, seq_len, n_head, d_k) # 再 transpose 成 (batch, n_head, seq_len, d_k) Q = Q.view(batch_size, -1, self.n_head, self.d_k).transpose(1, 2) K = K.view(batch_size, -1, self.n_head, self.d_k).transpose(1, 2) V = V.view(batch_size, -1, self.n_head, self.d_k).transpose(1, 2) # 3. 缩放点积注意力 # scores: (batch, n_head, seq_len, seq_len) scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) if mask is not None: scores = scores.masked_fill(mask == 0, float("-inf")) attn_weights = F.softmax(scores, dim=-1) attn_weights = self.dropout(attn_weights) # context: (batch, n_head, seq_len, d_k) context = torch.matmul(attn_weights, V) # 4. 拼接多头 # 先转置回 (batch, seq_len, n_head, d_k) # 再 contiguous + view 回 (batch, seq_len, d_model) context = context.transpose(1, 2).contiguous().view( batch_size, -1, self.d_model ) # 5. 输出投影 output = self.w_o(context) return output, attn_weights

代码中几个容易出错的地方:

  • view之后必须注意维度的排列顺序。先view(batch_size, -1, n_head, d_k)transpose(1, 2),得到的才是(batch, n_head, seq_len, d_k)
  • transpose之后,张量内存可能不连续,拼接前需要调用contiguous(),否则view会报错。
  • masked_fill(mask == 0, float("-inf"))会把被 mask 掉的位置变成负无穷,softmax 后这些位置的权重接近 0。

4.3 与 nn.MultiheadAttention 对比测试

PyTorch 官方已经提供了nn.MultiheadAttention,我们可以并行跑一遍,验证形状是否一致。

import torch import torch.nn as nn d_model = 512 n_head = 8 batch_size = 2 seq_len = 10 dropout = 0.1 custom_mha = MultiHeadAttention(d_model, n_head, dropout) builtin_mha = nn.MultiheadAttention(d_model, n_head, dropout=dropout, batch_first=True) x = torch.randn(batch_size, seq_len, d_model) out_custom, attn_custom = custom_mha(x, x, x) out_builtin, attn_builtin = builtin_mha(x, x, x) print("自定义多头注意力输出形状:", out_custom.shape) print("内置多头注意力输出形状:", out_builtin.shape) print("自定义注意力权重形状:", attn_custom.shape) print("内置注意力权重形状:", attn_builtin.shape)

预期输出是:

自定义多头注意力输出形状: torch.Size([2, 10, 512]) 内置多头注意力输出形状: torch.Size([2, 10, 512]) 自定义注意力权重形状: torch.Size([2, 8, 10, 10]) 内置注意力权重形状: torch.Size([2, 10, 8, 10])

两边的输出张量形状完全一致,只是注意力权重的维度排布不同。自定义实现里注意力权重是(batch, n_head, seq_len, seq_len),内置模块默认返回(batch, seq_len, n_head, seq_len)。这是因为官方接口里把 seq_len 放在了前面,语义上没有差别。

数值上两边不会完全一致,因为线性层初始化参数不同。需要验证的是变换逻辑,而不是数值相等。如果你希望严格对齐,可以手动把内置模块的in_proj_weightin_proj_bias拷贝到自定义实现中,再对比输出,但一般学习阶段不需要做这一步。

5. 功能验证:维度、掩码与注意力可视化

5.1 维度变化验证

用一个小例子,逐步打印每个阶段的张量形状,是最快理解多头注意力的方式。下面脚本把上一节的 Q、K、V 中间变量分别打印出来:

import math import torch import torch.nn as nn d_model = 64 n_head = 4 batch_size = 1 seq_len = 6 class DebugMHA(nn.Module): def __init__(self): super().__init__() self.n_head = n_head self.d_k = d_model // n_head 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) def forward(self, x): B = x.size(0) Q = self.w_q(x) K = self.w_k(x) V = self.w_v(x) print("Q shape:", Q.shape) Q = Q.view(B, -1, self.n_head, self.d_k).transpose(1, 2) K = K.view(B, -1, self.n_head, self.d_k).transpose(1, 2) V = V.view(B, -1, self.n_head, self.d_k).transpose(1, 2) print("Q after split:", Q.shape) scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) print("scores shape:", scores.shape) attn = torch.softmax(scores, dim=-1) context = torch.matmul(attn, V) print("context shape:", context.shape) context = context.transpose(1, 2).contiguous().view(B, -1, d_model) print("context after concat:", context.shape) out = self.w_o(context) print("output shape:", out.shape) return out x = torch.randn(batch_size, seq_len, d_model) model = DebugMHA() model(x)

运行这段代码,你会清楚看到每一步的维度变化。判断实现是否正确,标准就是最终输出形状和输入形状一致,以及scores的形状符合(batch, n_head, seq_len, seq_len)

5.2 mask 掩码对注意力的影响

在 Transformer 中,mask 主要有两种用途:

  • padding mask:把无效位置填充为很小的数,避免模型关注填充 token。
  • 因果 mask:解码器中避免模型看到未来 token。

以 padding mask 为例,假设序列长度为 4,其中最后一个 token 是填充项,mask 向量为[1, 1, 1, 0]。在注意力计算里,这个 mask 会被广播到所有 head。

用上一节的手写模块测试:

import torch from your_implementation import MultiHeadAttention model = MultiHeadAttention(d_model=64, n_head=4) x = torch.randn(2, 4, 64) # 假设第二批次的最后一个 token 是 padding mask = torch.tensor([ [1, 1, 1, 1], [1, 1, 1, 0] ]).unsqueeze(1).unsqueeze(2) # (2, 1, 1, 4) out, attn = model(x, x, x, mask=mask) print(attn)

mask 形状需要能广播到(2, 4, 4, 4)的 scores 上,所以这里是(batch, 1, 1, seq_len)。可以看到被 mask 位置对应的注意力权重几乎为 0,因为masked_fill将分数设置为-inf后,softmax 会把它们压到 0。

常见错误是把 mask 的形状搞错。如果 mask 少了维度,masked_fill会广播失败或者没有按预期遮挡。建议在实现里写成scores.masked_fill(mask == 0, float("-inf")),并通过打印 shape 确认广播后形状。

5.3 注意力权重可视化

注意力权重是理解模型行为的直接入口。以一句话为例子,可以画出每个 token 到其他 token 的热力图。

import matplotlib.pyplot as plt import torch from your_implementation import MultiHeadAttention model = MultiHeadAttention(d_model=64, n_head=4) tokens = ["我", "爱", "深度学习", "和", "自然语言处理"] x = torch.randn(1, len(tokens), 64) _, attn = model(x, x, x) head_idx = 0 plt.figure(figsize=(6, 5)) plt.imshow(attn[0, head_idx].detach().numpy(), cmap="Blues") plt.xticks(range(len(tokens)), tokens, rotation=45) plt.yticks(range(len(tokens)), tokens) plt.colorbar() plt.title("Head 0 Attention Weights") plt.tight_layout() plt.show()

如果使用随机初始化的模型,热力图通常比较平滑,没有明显规律。如果模型已经训练好,你会看到不同 head 的关注点有明显差异。这也是验证“多头有效”最直观的方式。

5.4 资源与性能观察思路

多头注意力的主要计算开销来自注意力矩阵:

(batch_size, n_head, seq_len, seq_len)

序列长度增大时,注意力矩阵按平方增长,这是 Transformer 被称为“二次复杂度”模型的原因。观察资源占用可以分两部分:

  • CPU 或 GPU 推理耗时可以用time模块粗略统计。
  • 如果使用 GPU,可以用torch.cuda.max_memory_allocated()查看峰值显存。

实际数值取决于序列长度、batch size 和 head 数,不能一概而论。重点要记住:head 数增加会加大注意力矩阵的数量,但每个 head 的维度变小,总的参数量和计算量并不会成倍增长,因为每个头只负责d_model / n_head维子空间。

6. 多头注意力在真实模型中的应用

6.1 Transformer 编码器与解码器

在 Transformer 编码器中,每层包含两个子层:多头注意力和前馈网络。输入经过多头注意力之后,会经过残差连接和 LayerNorm,再进入前馈网络。

在解码器中,多头注意力出现两次:

  • 第一次是自注意力,用因果 mask 屏蔽未来 token。
  • 第二次是编码器-解码器注意力,Query 来自解码器,Key 和 Value 来自编码器输出,帮助解码器获取输入序列的信息。

这两种场景下,多头注意力的计算逻辑完全相同,区别只在于 mask 和输入来源。

6.2 BERT、GPT 等预训练模型

BERT-base 使用 12 层 Transformer,每层 12 个 head,d_model=768,每个头的维度是 64。GPT-2 同样使用 12 个 head,更大规模版本会增加层数和 head 数。

阅读 BERT 和 GPT 源码时,你会发现它们对多头注意力的实现有两种风格:一种是显式做 QKV 线性投影后再 reshape 拆头;另一种是使用nn.MultiheadAttention封装。前者的好处是可控性强,后者更简洁。理解了手写版本后,再去读这两类源码都会轻松很多。

6.3 视觉 Transformer(ViT)

ViT 把图片切分成固定大小的 patch,每个 patch 展开成 token 后送入标准 Transformer。图像里的多头注意力同样按 patch 位置计算相关性,因此不同 head 可能学到不同尺度或方向的图像特征。这是近两年视觉领域大量使用 Transformer 结构的直接原因之一。

6.4 head 数量怎么选

论文中默认 8 个 head,之后大量模型沿用 12 或 16 个 head。head 数增加能提升模型表达能力,但显存和训练时间也会增加。更关键的是,很多研究发现并非所有 head 都对最终效果有贡献,剪掉部分 head 性能下降有限。这也是多查询注意力(MQA)和分组查询注意力(GQA)的出发点:在推理阶段减少 K、V 头数,降低显存带宽开销,同时保持模型输出质量。理解这些优化的前提,还是把标准多头注意力吃透。

7. 常见问题与排查方法

问题现象可能原因排查方式解决方案
代码报错,无法viewtranspose后张量内存不连续检查报错信息,确认是否提示contiguous拼接前调用contiguous()
d_model无法整除n_head维度配置不合理打印d_modeln_head调整d_modeln_head,让两者整除
训练 loss 不下降忘记除以√d_k,softmax 梯度饱和检查注意力代码中是否有math.sqrt(self.d_k)补上缩放因子
注意力权重几乎均匀分布模型未训练,或初始化不合理打印注意力矩阵,观察是否集中在少数 token先跑小数据集验证,再调学习率
mask 不生效mask 形状不对,广播后无法匹配 scores打印 mask 和 scores 的形状把 mask 扩展到(batch, 1, 1, seq_len)形式
解码器看到未来信息因果 mask 设置错误检查 mask 是否为上三角矩阵使用torch.triu(..., diagonal=1)生成上三角掩码
序列稍长就显存不足注意力矩阵 O(n²) 占用过大查看峰值显存缩短序列、使用 FlashAttention 或分块注意力
自定义实现和官方输出差异大初始化参数不同对比两者的逻辑和形状,不对比具体数值如需严格对齐,复制官方权重初始化方式

8. 最佳实践与学习建议

8.1 实现与调试建议

第一次写多头注意力时,不要上来就跑大模型。建议先用d_model=64n_head=4seq_len=8这样的小参数把链路跑通,确认每一步的维度都如预期。

调试时可以沿用三件套:

  1. 在每个阶段打印张量形状。
  2. 写一个和nn.MultiheadAttention的对比脚本,校验形状。
  3. 用一个你熟悉的小任务,比如简单的序列复制或情感分类,让模型训练 20 步左右,观察 loss 是否能下降。

还有一个容易被忽略的点是初始化。PyTorch 的nn.Linear默认初始化通常能直接使用,但如果你复现论文时要严格对齐,需要关注权重初始化的细节。学习阶段不需要过度纠结这一点。

8.2 从理解到工程

理解标准多头注意力后,就可以继续往下读这些方向:

  • 位置编码:注意力本身不包含位置信息,需要靠位置编码补充。
  • 完整 Transformer Encoder/Decoder:多头注意力只是其中一个子层。
  • FlashAttention:通过分块计算和 IO 优化,把注意力计算变得更快更省显存。
  • MQA / GQA:在推理阶段减少 K、V 头数的优化方案。
  • 不同变体的源码实现:读 transformers 库或 Fairseq 的注意力实现,看工程化封装思路。

建议动手做一个小实验:用标准多头注意力实现一个最小字符级语言模型,在几十万字符的语料上训练几轮,然后观察生成结果。这个实验能让你把公式、代码和实际效果连通。

9. 总结与下一步

多头注意力是整个 Transformer 架构中最需要熟练掌握的模块之一。它建立在缩放点积注意力之上,通过拆接多个子空间,让模型可以同时表达多种依赖关系。这篇文章给出了完整公式、形状变化、PyTorch 从零实现和与官方模块的对比验证,重点在于你能否跟着代码跑一遍。

如果你已经看懂了这版实现,下一步建议做三件事:

  1. 把代码里的 mask 改成因果 mask,实现一个最小解码器。
  2. 找一个已经训练好的小型 BERT 或 GPT 模型,提取某层注意力权重并可视化。
  3. 对比 FlashAttention 的原理,理解标准实现里哪些地方可以优化。

在跑通上述内容之前,不需要急着读大模型源码,更不用一开始就在大批量数据上训练。先把多头注意力的每个张量形状烂熟于心,再去碰完整 Transformer 会顺畅得多。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询