今天我们不聊库怎么装,也不聊某个模型怎么跑,而是把 Transformer 里面最容易被跳过、又最难啃的一块骨头单独拿出来拆干净:多头注意力机制。很多同学看 Transformer 源码时会发现,代码里有一堆view、transpose、matmul,如果你没有把多头注意力的数据结构彻底搞清楚,这几行代码看一天也看不懂。这篇文章会把多头注意力原理、维度变换、PyTorch 实现、因果掩码、显存占用估算一次讲完,并且给出可以直接运行的代码片段。适合正在看 Transformer 论文、读 PyTorch 源码、准备自己写注意力模块,或者想搞清楚 BERT/GPT 内部结构的人。
先把结论放在前面:多头注意力不是让模型“多算几次注意力”,而是把输入向量切到多个子空间,在每个子空间里独立计算注意力,再拼回来做一次线性变换。它解决的核心问题,是单个注意力头只能从一种关系或一种模式去计算依赖,而真实语言的依赖非常复杂,比如指代关系、语法关系、语义相似性往往需要同时捕捉。多头注意力用更低的维度并行处理多组关系,计算成本接近原来的单头注意力,但表达能力明显更强。
本文会覆盖五个方面:多头注意力的数学原理、Q/K/V 和维度变换、因果自注意力掩码、PyTorch 可运行实现、以及训练和推理中的性能边界。看完之后,你能独立写出一个多头注意力模块,也能理解 BERT、GPT 代码中scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)这一行到底在做什么。
1. 多头注意力核心概念速览
| 项目 | 说明 |
|---|---|
| 所属领域 | 深度学习基础模块,Transformer 的核心组件 |
| 核心作用 | 让模型同时从多个子空间捕获输入序列中的依赖关系 |
| 基本组成 | 线性投影层、多头切分、缩放点积注意力、拼接、输出投影 |
| 典型模型配置 | BERT base:12 层、12 头、hidden_size=768,每头维度 64 |
| 典型模型配置 | BERT large:24 层、16 头、hidden_size=1024,每头维度 64 |
| 典型模型配置 | GPT-2:12 层、12 头、hidden_size=768,因果自注意力 |
| 适用任务 | 文本分类、机器翻译、文本生成、多模态、图网络等 |
| 主要计算瓶颈 | attention score 矩阵的大小为 batch × heads × seq_len × seq_len |
| 是否独立训练 | 通常不作为单独模型训练,嵌入在 Transformer Block 中 |
这里要强调一点:在标准设置下,多头注意力的总参数量和单头注意力的参数量完全一样。因为多头会把d_model切成h个头,每个头维度是d_k = d_model / h,所有头的参数量之和还是原来的量级。真正让多头变强的原因,不是参数量变大,而是计算结构变化了。
2. 适用场景与理解边界
多头注意力是 Transformer、BERT、GPT、T5 等一系列模型的基础构件,几乎所有现代大模型结构里都有它。它适合这些场景:
- 文本序列建模:捕捉位置之间长短不一的依赖关系。
- 机器翻译:对齐源语言和目标语言,同时处理多个语义层面的对应关系。
- 文本生成:GPT 系列用因果自注意力限制每个位置只能看左侧内容。
- 多模态:把图像 patch 和文本 token 放在同一序列里做交叉注意力。
- 代码模型:捕捉变量定义与引用之间的跨行长距离依赖。
但多头注意力并不是万能的。第一,它的计算量随序列长度平方增长,直接用在超长序列上会非常吃力。第二,对于短序列或者非常简单的任务,多头带来的提升可能不明显,反而会引入更多的超参数调优成本。第三,多头注意力内部的可解释性有限,所谓“不同的头学到不同模式”并不总是成立,很多头训练后可能功能高度重合,甚至基本退化。理解这一点很重要:多头只是提供了一种更丰富的建模方式,不是保证模型变聪明的魔法。
另外,在使用包括多头注意力在内的深度学习技术时,要注意数据合规问题。如果训练数据涉及人脸、声音、隐私文本或版权素材,必须确认已经获得合法授权;在生产环境部署相关模型时,也要遵守平台的隐私和数据安全规定。
3. 多头注意力机制原理详解
3.1 从自注意力到缩放点积注意力
自注意力(Self-Attention)的输入是一个序列向量矩阵X,形状通常为(batch_size, seq_len, d_model)。它通过三个可学习的投影矩阵,把X映射成一组查询、键、值:
- Q(Query):代表当前词“想找什么”,可以理解为提问。
- K(Key):代表当前词“能提供什么”,可以理解为索引标签。
- V(Value):代表当前词“真正携带的信息”,可以理解为内容。
在缩放点积注意力中,模型先计算 Q 和 K 的点积,得到两两之间的相关分数,再除以缩放因子sqrt(d_k)以避免点积结果过大导致 softmax 梯度消失,最后过 softmax 得到注意力权重,并与 V 做加权求和。
缩放点积注意力的公式如下:
Attention(Q, K, V) = softmax(Q K^T / sqrt(d_k)) V其中d_k是每个头的键向量维度。除以sqrt(d_k)是 Attention Is All You Need 论文里的关键设计。如果不做缩放,当d_k很大时,Q 和 K 点积的方差会变大,softmax 的输入分布会进入饱和区,梯度很小,训练容易不稳定。
3.2 为什么要“多头”
单头注意力只能计算一组 Q/K/V 拟合一种关系。但真实文本里,一个词往往同时和多个词存在不同类型的关系。比如“小明把书递给小红”这句话,“传递”这个动作可能和“小明”“书”“小红”同时相关,但这种相关性在不同抽象层次上表现不一样。单头注意力会把所有关系混在一个平均后的权重里,表达能力受限。
多头注意力的思路是:把d_model维的 Q、K、V 全部切分成h份,每一份代表一个子空间,每组子空间独立计算注意力。这样模型可以并行学习多套不同的注意力模式。比如一头可能偏向距离较近的词,另一头偏向句法关系,还有一头可能负责指代关系。虽然这种“分工”不是显式监督出来的,但在一定语义任务上确实能观察到不同头关注不同位置的倾向。
3.3 多头注意力的完整计算过程
以一个输入向量维度d_model = 768、头数h = 12的配置为例,每个头的维度是d_k = d_model / h = 64。
计算过程分为四步:
- 对输入
X做三次线性投影,得到 Q、K、V,三者形状都是(batch_size, seq_len, d_model)。 - 把 Q、K、V 按最后一维切分成
h块,得到形状(batch_size, seq_len, h, d_k),再变换成(batch_size, h, seq_len, d_k)。这里注意,transpose(1, 2)是为了让每个头独立完成批量矩阵乘法。 - 对每个头分别计算缩放点积注意力,得到形状为
(batch_size, h, seq_len, d_k)的输出。 - 把全部头的输出拼接回
(batch_size, seq_len, d_model),再通过输出投影矩阵W_o融合起来。
形式化写法是:
MultiHead(Q, K, V) = Concat(head_1, head_2, ..., head_h) W^O 其中 head_i = Attention(Q W_i^Q, K W_i^K, V W_i^V)这里W_i^Q、W_i^K、W_i^V是每个头的独立投影,实际工程实现中一般不用单独定义h个矩阵,而是先做一次大的nn.Linear(d_model, d_model),再通过 reshape 切头。
3.4 多头注意力在 Transformer Block 中的完整链条
多头注意力不会孤零零地工作。在标准的 Transformer Encoder Block 中,它处在这样的链条里:
输入 X -> 多头自注意力(Multi-Head Self-Attention) -> 残差连接(Residual Connection,将输入和注意力输出相加) -> 层归一化(Layer Normalization) -> 前馈网络(Feed-Forward Network / MLP) -> 残差连接 -> 层归一化 -> 输出层归一化在这里非常关键。因为注意力输出和 MLP 输出的数值分布会随着层数加深不断变化,LayerNorm 可以把每一层的输入拉到比较稳定的范围,让训练更稳。而多层感知机(MLP)部分通常是两个全连接层加激活函数,比如d_model -> 4*d_model -> d_model,给了模型在注意力聚合之后做进一步非线性变换的能力。所以在学习多头注意力时,最好连带着理解残差连接、层归一化和 MLP,它们共同构成了完整的 Transformer 基本单元。
3.5 因果自注意力:生成模型的掩码机制
GPT 这类自回归生成模型使用因果自注意力(Causal Self-Attention)。它和多头注意力不是对立关系,而是多头注意力在自回归场景下的一种约束。生成第t个 token 时,模型不能看到第t+1及之后的 token,因此需要在计算 attention score 时把未来位置遮住。
具体做法是构造一个上三角掩码矩阵,形状为(seq_len, seq_len),将当前位置右侧的元素置为-inf或False。在 PyTorch 中常见写法是:
import torch seq_len = 8 causal_mask = torch.tril(torch.ones(seq_len, seq_len, dtype=torch.bool))# 显示掩码 causal_mask输出是一个下三角为 True、上三角为 False 的矩阵。在多头注意力实现中,对scores执行masked_fill,把上三角位置填充成负无穷,softmax 之后这些位置的概率就会变成 0。这样每个位置只能看到自己和之前的位置,保证了自回归的因果性。
4. PyTorch 实现多头注意力
下面给出一份可以直接运行的多头自注意力实现。代码不依赖 HuggingFace,只用 PyTorch,把维度变换和掩码逻辑摆出来,方便你对照公式理解。
import torch import torch.nn as nn import math class MultiHeadSelfAttention(nn.Module): def __init__(self, d_model, n_heads, dropout=0.1): super().__init__() assert d_model % n_heads == 0, "d_model must be divisible by n_heads" self.d_model = d_model self.n_heads = n_heads self.d_k = d_model // n_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, x, mask=None): batch_size, seq_len, _ = x.size() # 1. 线性投影并切分到多头 Q = self.W_q(x).view(batch_size, seq_len, self.n_heads, self.d_k).transpose(1, 2) K = self.W_k(x).view(batch_size, seq_len, self.n_heads, self.d_k).transpose(1, 2) V = self.W_v(x).view(batch_size, seq_len, self.n_heads, self.d_k).transpose(1, 2) # 2. 计算缩放点积注意力分数 scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) # 3. 可选掩码 if mask is not None: scores = scores.masked_fill(mask == 0, float("-inf")) # 4. softmax 和 dropout attn_weights = torch.softmax(scores, dim=-1) attn_weights = self.dropout(attn_weights) # 5. 加权求和,然后拼接多头 context = torch.matmul(attn_weights, V) context = context.transpose(1, 2).contiguous().view(batch_size, seq_len, self.d_model) # 6. 输出投影 output = self.W_o(context) return output验证一下输出形状:
d_model = 768 n_heads = 12 seq_len = 128 batch_size = 2 x = torch.randn(batch_size, seq_len, d_model) mha = MultiHeadSelfAttention(d_model, n_heads, dropout=0.1) out = mha(x) print("输入形状:", x.shape) print("输出形状:", out.shape)预期输出:
输入形状: torch.Size([2, 128, 768]) 输出形状: torch.Size([2, 128, 768])再看因果掩码的用法。构造一个(batch_size, n_heads, seq_len, seq_len)的掩码,或者只用二维掩码广播也行,代码里masked_fill(mask == 0, float("-inf"))会按广播规则处理:
# 构造因果掩码,形状为 (seq_len, seq_len) causal_mask = torch.tril(torch.ones(seq_len, seq_len, dtype=torch.bool)) # 扩展到 (1, 1, seq_len, seq_len),方便广播 causal_mask = causal_mask.unsqueeze(0).unsqueeze(0) out_causal = mha(x, mask=causal_mask) print("因果注意力输出形状:", out_causal.shape)这段代码里最关键的两步是transpose(1, 2)和contiguous().view(...)。前者把维度从(batch, seq, head, d_k)变成(batch, head, seq, d_k),让每一头的序列元素可以独立做矩阵乘法;后者在拼接前把内存布局理顺,避免view报错或得到错误结果。
如果不想手写,PyTorch 也提供了现成的nn.MultiheadAttention模块。它的参数默认顺序是(query, key, value),并且默认使用第 0 维作为序列维度,需要在传入前转置。日常做实验直接用官方模块即可,但手写一遍能更清楚地理解内部结构。
5. 关键参数设计与影响
5.1 头数h和每头维度d_k的取值逻辑
在标准 Transformer 中,d_model和h通常是固定搭配,常见做法是让d_k = d_model / h = 64。从 BERT base 到 LLaMA 的大多数模型,都沿用了“每头维度 64 或 128”的经验值。比如d_model=768时用 12 头,d_model=1024时用 16 头,d_model=4096时可能用 32 头。这样做的好处是保持了每个头内部矩阵乘法的计算特性和初始化尺度相对稳定。
如果h设置过大,每个头能看到的维度太细,容易让单个头退化,噪声增加;如果h设置过小,多子空间表达的优势就不明显。一般不建议单独把h拉到很大,比如说在d_model=768时设 64 头,每头只有 12 维,实践中往往不稳定。
5.2 多头数量与模型容量的关系
虽然多头参数总量和单头一致,但多头实际上增加了类似“并行表达能力”的效果。大量实验观察表明,头数越多,模型能在更多位置上同时捕捉不同模式,但也更容易出现过拟合,特别是在小数据集上。因此头数属于需要根据任务和模型规模调节的超参数,没有绝对最优。
5.3 残差连接和层归一化对多头注意力的保护
深层 Transformer 中,多头注意力的输出和输入之间一定会加残差连接和 LayerNorm。原因是多头注意力内部包含若干矩阵乘法和 softmax,输出分布不稳定,直接堆叠会放大方差,训练容易发散。加入 LayerNorm 后,输入到下一层的特征尺度被约束在一个相对稳定的范围。这里可以顺便看到网络热词里的“层归一化”和“多层感知机”在 Transformer 里的实际位置:它们和多头注意力配合,组成完整的 Transformer Block。
5.4 不同序列长度下的行为差异
多头注意力对序列长度非常敏感。序列越长,注意力矩阵越大,每个位置需要聚合的信息越多,最终输出的语义会更偏向全局平均,而短序列下每个位置能关注的邻近信息相对有限。因此,在处理超长文本时,很多人会改用稀疏注意力、滑动窗口注意力或 FlashAttention 等优化方法,而不是无脑加大seq_len。
6. 资源占用与性能观察方法
这一节主要从公式层面说明如何估算多头注意力带来的显存或内存占用,具体数值要结合你本机的 PyTorch、CUDA 版本和模型配置实测。
6.1 参数占用量估算
多头注意力模块中,Q/K/V 和输出投影各自是一个nn.Linear(d_model, d_model),四个矩阵的参数量约为:
4 * d_model * d_model以d_model=768为例,参数约 236 万。这部分只占 Transformer 总参数量的一小部分,因为 MLP 部分的参数量通常更大。
6.2 中间张量显存占用量估算
训练过程中,真正的显存大头来自 attention score 矩阵。它的形状是:
(batch_size, n_heads, seq_len, seq_len)假设batch_size=2、n_heads=12、seq_len=512,那么 score 矩阵有约 629 万个元素,每个元素如果使用 FP16 存储,则约 12.6 MB。如果seq_len变成 4096,则占用量会变成约 805 MB。这就是为什么长序列下显存爆炸非常快。实际训练时还会保存梯度以做反向传播,占用量会进一步翻倍。
6.3 如何观察显存占用
在 PyTorch 里,可以这样观察当前 GPU 显存占用:
import torch if torch.cuda.is_available(): allocated = torch.cuda.memory_allocated() / 1024**2 reserved = torch.cuda.memory_reserved() / 1024**2 print(f"当前分配显存: {allocated:.2f} MB") print(f"当前保留显存: {reserved:.2f} MB")如果想看某个算子的精确显存消耗,可以在推理阶段用torch.profiler或torch.cuda.memory._record_memory_history,不过这两者在不同的 PyTorch 版本中 API 会有差异。更简单的做法是控制变量:固定 batch 和序列长度,逐步增加头数,观察显存变化曲线,判断当前模型的瓶颈。
6.4 如何降低资源占用
在多头注意力中降低显存占用,常见的思路有:
- 降低 batch size 或 seq_len。
- 使用 FlashAttention 等 IO 感知注意力实现,避免显式生成完整的
seq_len × seq_len矩阵。 - 推理阶段使用 KV Cache,避免重复计算已经算过的 K/V。
- 使用 GQA/MQA,让多个头共享 K/V,减少 KV Cache 占用。
- 混合精度训练,用 FP16/BF16 减少中间张量大小。
但要注意,这些优化方式大多数需要配合特定硬件和库版本进行测试,不是所有环境都能直接获得收益。
7. 常见问题排查表
| 问题现象 | 可能原因 | 排查方式 | 解决方案 |
|---|---|---|---|
| 训练时 loss 不下降 | attention score 未做缩放,softmax 饱和 | 检查是否除以 sqrt(d_k) | 在 scores 计算中加入缩放因子 |
| 训练时输出 NaN | softmax 输入包含 NaN 或极大值 | 检查输入张量和梯度 | 缩小 learning rate,检查初始化 |
| 因果注意力结果异常 | mask 形状或位置不正确 | 打印 mask 的前几行,检查是否覆盖未来位置 | 使用torch.tril构造 mask,并确认 broadcast 维度 |
| 多头拼接后维度错误 | view前没有contiguous | 检查报错信息和张量形状 | 使用transpose().contiguous().view() |
| 显存不足 | seq_len 过大导致 attention 矩阵爆炸 | 打印 score 张量形状 | 降低 seq_len/batch_size,或用 FlashAttention |
| 多头效果和单头差不多 | 头数过多或数据集太小 | 观察多头注意力权重是否分散 | 减少头数,或增强模型容量和训练数据 |
| 模型推理速度慢 | 未使用 KV Cache 或注意力实现低效 | 检查推理时重复计算的 K/V | 使用 KV Cache 或优化注意力实现 |
| 输出权重分布太平均 | 注意力计算没有学到有效依赖 | 检查输入编码和位置编码 | 增加训练步数,或调整头数 |
8. 多头注意力的最佳实践与使用建议
8.1 先搞清需求再选参数
如果是复现 BERT 或 GPT,不要自行魔改头数,直接沿用公开配置:d_model / n_heads尽量保持 64 或 128。如果是自定义小模型,建议从 8 头或 12 头起步,通过验证集调参,不要一开始就把头数调到 64。
8.2 用库优先,手写仅用于学习
日常开发建议使用 PyTorch 的nn.MultiheadAttention或 HuggingFace Transformers 里的注意力实现,这些模块经过大量测试,效率和稳定性更高。手写多头注意力有助于理解原理,但不建议直接上生产环境。
8.3 训练时监控梯度与注意力分布
可以打印以下内容进行诊断:
- attention 权重的均值、方差、稀疏程度。
- Q/K 的梯度范数。
- LayerNorm 前后输出的均值和标准差。
如果注意力权重一直非常接近均匀分布,说明模型没有学到有效的信息选择;如果 attention 分数出现过大的正负值,则要确认缩放因子和初始化是否合理。
8.4 长序列场景优先考虑优化方案
在长文本、长视频序列或高分辨率图片建模任务里,标准多头注意力的计算成本会迅速超过模型本身的计算能力。这时候可以优先考虑稀疏注意力、滑动窗口注意力、线性注意力,或者 FlashAttention。注意这些优化方案往往带有一定的近似性,准确率、显存、速度三者需要实际测量后选择。
8.5 数据合规和部署安全
如果多头注意力模块被用来处理真实人物的人脸、声音、用户隐私文本或受版权保护的素材,必须事先获得合法授权。在部署模型服务时,建议限制接口访问范围,避免生成内容被滥用;涉及商用场景要复核模型的稳定性和安全性。
9. 总结与下一步
多头注意力是目前最值得反复理解的一个深度学习基础模块。它解决了单头注意力表达单一的问题,用“切分、并行、拼接、投影”四个步骤,让模型在计算量几乎不变的情况下获取多子空间建模能力。理解它的关键在于 Q/K/V 投影、维度变换和掩码机制,尤其是(batch, heads, seq_len, d_k)这种四维张量的流转过程。只要把多头注意力代码手写一遍,再回头看 BERT 或 GPT 的源码,你会发现很多困惑会自然消失。
下一步建议你在自己熟悉的框架里完成三个练习:第一,用随机输入跑通本文代码;第二,给模块加上因果掩码,测试自回归生成场景;第三,把多头注意力嵌入一个两层的 Transformer Block,训练一个小的文本分类或语言模型任务,观察不同头数对收敛速度和最终指标的影响。最容易踩的坑是维度拼接和掩码广播,一定要多打印中间张量形状。上面这些内容如果对你有帮助,建议收藏备用,后面写 Transformer 相关代码时可以随时对照。