PyTorch 2.x 深度学习专题【左扬精讲】—— 多头注意力机制:为什么一个头不够用?
PyTorch 2.x 深度学习专题【左扬精讲】—— 多头注意力机制:为什么一个头不够用?
上一篇文章我们讲了 Self-Attention(自注意力机制),它能让序列中任意两个位置直接"对话",不受距离限制。这已经很厉害了,但还有一个问题:如果只用单一注意力头,模型只能学习到一种类型的相关性。
举个例子。我们在分析一条运维日志的时候,脑子里其实会同时关注很多不同的东西:
-
-
- 错误关键词是什么?(ERROR、timeout、exception)
- 时间间隔正常吗?(这次报错和上次隔了多久?)
- 和哪些服务有关?(order-api、database、payment)
- 上下文有没有关联?(前面是不是刚刚重启过?)
-
如果只有一个注意力头,它可能只学会了关注"错误关键词",而忽略了其他同样重要的信息。这就是多头注意力(Multi-Head Attention)要解决的问题。
torch.nn.MultiheadAttention ← PyTorch 多头注意力实现
torch.matmul ← 矩阵乘法,用于计算注意力分数
torch.nn.functional.softmax ← softmax 函数,用于归一化注意力权重
torch.nn.Linear ← 线性变换层,用于 QKV 投影
Multi-Head Attention 注意力机制 Transformer QKV投影
学习重点提示
- 必须掌握:多头注意力的核心思想、QKV 投影和分头计算
- 需要理解:为什么多个头比一个头好、每个头学到了什么
- 建议了解:不同头的可视化解释、头数如何选择
一、从一个问题开始:单一注意力头有什么局限?
What — 什么是单一注意力头的问题?
想象一下,你是一个 SRE 工程师,正在分析一条复杂的日志:
2024-01-15 10:30:45 ERROR [order-api] connection_timeout
retry=3 duration=5000ms service=database
你一眼就能看出很多信息:这是个 ERROR 级别的问题,order-api 服务连不上 database,超时了 5 秒,重试了 3 次才失败。
但如果只有一个注意力头呢?它可能会这样看这个问题:
- 计算"ERROR"和"connection_timeout"的相关性 → 高
- 计算"order-api"和"database"的相关性 → 高
- 计算"retry=3"和"timeout"的相关性 → 中
看起来不错,但它只能学习到一种"看问题的方式"。它不知道"时间间隔"可能很重要,不知道"服务间调用关系"和"错误类型"可能是两件完全不同的事。
Why — 为什么需要多个头?
问题一:不同类型的相关性需要不同的"视角"
在日志分析中,我们需要同时关注:
- 词汇相关性:ERROR 和 timeout 通常一起出现
- 服务调用关系:order-api 调用了 database
- 时序模式:这次错误和 5 分钟前的重启有关
- 因果关系:因为 database 慢,所以 order-api 超时
每种相关性都需要不同的"注意力模式"。一个头很难同时学会所有这些。
问题二:单一头的表示空间有限
假设模型维度是 512,每个头分到的维度是 512/8=64(8个头的情况下)。64 维的空间能表达的信息是有限的。如果把词汇关系、服务关系、时序关系都塞进 64 维,就像把百科全书塞进一个小文件夹,肯定会有信息丢失。
没有多头注意力会发生什么?
- 模型只能学到一种类型的相关性,漏检复杂故障
- 不同类型的信息相互干扰,效果打折扣
- 模型表达能力受限,无法处理复杂场景
本节要点
- 单一头的局限:只能学习一种"视角",表示空间有限
- 多头解决思路:把一个复杂的任务拆成多个简单的子任务
- 类比:就像一个全科医生 vs 多个专科医生联合会诊
二、多头的核心思想:分而治之
What — 多头注意力到底是什么?
多头注意力的核心思想很简单:不要让一个头做所有事,让多个头分工合作。
具体来说:
- 把 d_model 维的输入分成 h 个部分(h = 头数)
- 每个头只处理自己的那部分维度
- 每个头独立计算注意力权重
- 最后把所有头的输出拼接起来,再做一个线性变换
这样,每个头可以专注于学习一种类型的相关性,而不必和别的类型竞争表示空间。
Why — 多头设计为什么有效?
设计意图一:专业的分工
想象一个破案团队:
-
-
- 侦探A专门负责分析人际关系
- 侦探B专门负责分析时间线
- 侦探C专门负责分析动机
- 侦探D专门负责分析物证
-
最后大家汇总信息,比一个人单打独斗强多了。多头注意力就是这个道理。
设计意图二:增加模型的表达能力
从数学上看,8 个头 × 64 维 = 512 维,拼接后还是 512 维。但为什么表达能力增强了?
关键在于:每个头有自己独立的 QKV 投影矩阵。它们看到的是同一个输入的不同"投影",就像从不同的角度看同一个物体,每个角度都能发现一些新细节。拼接后,模型可以用 512 维的表示同时包含 8 种不同类型的相关性信息。
设计意图三:提高泛化能力
如果只有单头,模型必须在一个小空间里同时表示所有关系,很容易过拟合。多头让每个子空间只关注一种关系,降低了过拟合风险。
让我们用大白话拆解多头注意力的计算流程:
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
class MultiHeadAttentionDemo(nn.Module):
def __init__(self, d_model=512, n_heads=8, dropout=0.1):
super().__init__()
assert d_model % n_heads == 0 # 确保能整除
self.d_model = d_model # 模型总维度
self.n_heads = n_heads # 头的数量
self.d_k = d_model // n_heads # 每个头的维度
# 核心:每个头有自己独立的 QKV 投影
# 8个头 -> 8组不同的 W_q, W_k, W_v
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.fc = nn.Linear(d_model, d_model) # 最终输出投影
def split_heads(self, x, batch_size):
# 关键步骤:把 d_model 维度拆成 n_heads 个 d_k
# 输入: (batch, seq_len, d_model) = (batch, seq_len, 512)
# 输出: (batch, n_heads, seq_len, d_k) = (batch, 8, seq_len, 64)
x = x.view(batch_size, -1, self.n_heads, self.d_k) # 拆分维度
return x.transpose(1, 2) # 交换seq_len和n_heads位置
def forward(self, query, key, value, mask=None):
batch_size = query.size(0)
# ========== 步骤1:QKV 投影 ==========
# 所有头共享一套投影矩阵,但分开后各自独立计算
Q = self.W_q(query) # (batch, seq_len, d_model)
K = self.W_k(key)
V = self.W_v(value)
# ========== 步骤2:分头 ==========
# 每个头只处理 d_k = 64 维
Q = self.split_heads(Q, batch_size) # (batch, 8, seq_len, 64)
K = self.split_heads(K, batch_size)
V = self.split_heads(V, batch_size)
# ========== 步骤3:每个头独立计算注意力 ==========
# 头1算词汇相关性,头2算服务关系,头3算时序...
# scores shape: (batch, 8, 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, -1e9)
attn_weights = F.softmax(scores, dim=-1) # 每个头独立归一化
# 加权求和
context = torch.matmul(attn_weights, V) # (batch, 8, seq_len, 64)
# ========== 步骤4:合并多头 ==========
# 把 8 个头的输出拼接起来
context = context.transpose(1, 2).contiguous() # (batch, seq_len, 8, 64)
context = context.view(batch_size, -1, self.d_model) # (batch, seq_len, 512)
# 最终线性变换
output = self.fc(context)
return output, attn_weights
# ========== 可视化:每个头学到了什么 ==========
def visualize_attention_heads(model, input_tensor):
"""演示不同头关注不同类型的相关性"""
output, attn_weights = model(input_tensor, input_tensor, input_tensor)
# attn_weights shape: (batch, n_heads=8, seq_len, seq_len)
# attn_weights[0] 是第1个样本的注意力权重
# attn_weights[0, 0] 是头0的注意力矩阵
# attn_weights[0, 1] 是头1的注意力矩阵
# ...
print(f"注意力权重形状: {attn_weights.shape}")
print(f"如果有8个词,注意力矩阵就是8x8")
print(f"头0关注词汇相关性,行0列2值最大=0.9 表示第1个词最关注第3个词")
print(f"头3关注服务关系,行1列5值最大=0.85 表示第2个词最关注第6个词")
return attn_weights
本节要点
- 分头:把 d_model 拆成 h 个 d_k,每个头独立计算
- 独立:每个头有自己的 QKV 投影,互不干扰
- 拼接:把 h 个头的输出拼接,还原到 d_model 维度
- 分工:不同头可以专注于学习不同类型的相关性
三、How — PyTorch 代码实现详解
What — torch.nn.MultiheadAttention 是什么?
PyTorch 已经帮我们实现好了 torch.nn.MultiheadAttention,可以直接用。但理解它的内部原理很重要,这样你才知道:
- 参数该怎么调
- 为什么会慢、要怎么优化
- 怎么调试注意力权重
让我们对比一下:官方实现(简洁)和手动实现(理解原理):
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
# ========== 方式一:直接用 PyTorch 官方实现(生产环境推荐) ==========
def use_official_multihead_attention():
d_model = 512
n_heads = 8
dropout = 0.1
# 一行代码创建多头注意力层
mha = nn.MultiheadAttention(
embed_dim=d_model, # 输入维度
num_heads=n_heads, # 头数量
dropout=dropout, # dropout概率
batch_first=True # True: (batch, seq, dim),False: (seq, batch, dim)
)
# 输入: (batch, seq_len, d_model)
query = torch.randn(2, 10, 512) # 2个样本,每个10个token,512维
key = torch.randn(2, 10, 512)
value = torch.randn(2, 10, 512)
# 前向传播
# 返回: (output, attn_weights)
# output: (batch, seq_len, d_model)
# attn_weights: (batch, n_heads, seq_len, seq_len)
output, weights = mha(query, key, value)
print(f"输出形状: {output.shape}") # (2, 10, 512)
print(f"注意力权重形状: {weights.shape}") # (2, 8, 10, 10)
return output, weights
# ========== 方式二:手动实现(理解原理) ==========
def manual_multihead_attention():
batch_size = 2
seq_len = 10
d_model = 512
n_heads = 8
d_k = d_model // n_heads # 64
# 输入
x = torch.randn(batch_size, seq_len, d_model)
# QKV 投影(所有头共享一套投影)
W_q = nn.Linear(d_model, d_model)
W_k = nn.Linear(d_model, d_model)
W_v = nn.Linear(d_model, d_model)
Q = W_q(x) # (2, 10, 512)
K = W_k(x)
V = W_v(x)
# 关键:重塑张量,每个头分 d_k=64 维
# (2, 10, 512) -> (2, 8, 10, 64)
Q = Q.view(batch_size, seq_len, n_heads, d_k).transpose(1, 2)
K = K.view(batch_size, seq_len, n_heads, d_k).transpose(1, 2)
V = V.view(batch_size, seq_len, n_heads, d_k).transpose(1, 2)
# 计算注意力分数
# Q: (2, 8, 10, 64), K^T: (2, 8, 64, 10)
# scores: (2, 8, 10, 10)
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)
# softmax 归一化
attn_weights = F.softmax(scores, dim=-1)
# 加权求和
# attn_weights: (2, 8, 10, 10), V: (2, 8, 10, 64)
# context: (2, 8, 10, 64)
context = torch.matmul(attn_weights, V)
# 合并多头:把8个头的输出拼接
# (2, 8, 10, 64) -> (2, 10, 8, 64) -> (2, 10, 512)
context = context.transpose(1, 2).contiguous()
context = context.view(batch_size, seq_len, d_model)
# 最终输出投影
W_o = nn.Linear(d_model, d_model)
output = W_o(context)
return output, attn_weights
# ========== 对比结果 ==========
print("=== 官方实现 ===")
out1, w1 = use_official_multihead_attention()
print(f"输出形状: {out1.shape}, 注意力权重形状: {w1.shape}")
print("\n=== 手动实现 ===")
out2, w2 = manual_multihead_attention()
print(f"输出形状: {out2.shape}, 注意力权重形状: {w2.shape}")
实战技巧:如何查看每个头在关注什么?
import torch
import torch.nn as nn
# 创建多头注意力
mha = nn.MultiheadAttention(embed_dim=512, num_heads=8, batch_first=True)
# 模拟输入:假设这是日志中的8个token
query = torch.randn(1, 8, 512) # batch=1, seq=8, dim=512
# 获取注意力权重
_, attn_weights = mha(query, query, query)
# attn_weights: (batch, n_heads, seq, seq) = (1, 8, 8, 8)
print(f"第1个头的注意力权重矩阵 (8x8):")
print(attn_weights[0, 0]) # 第1个样本,第0个头
print(f"\n第3个头的注意力权重矩阵 (8x8):")
print(attn_weights[0, 2]) # 第1个样本,第2个头
# 通过可视化可以发现:
# 头0主要关注 位置0 和 位置5 的关系
# 头2主要关注 位置2 的自相关
# 头5主要关注 位置7 的上下文
四、What-If — 如果只有一个头会怎样?
What-If 实验一:只有一个头
假设我们把 8 个头换成 1 个头,会发生什么?
维度分配:8个头 × 64维 = 512维
1个头 × 512维 = 512维
看起来总维度没变,但实际上:
- 表示空间从 8×64 降到了 1×512
- 模型必须在 512 维空间里同时编码所有类型的相关性
- 不同类型的相关性会相互干扰
What-If 实验二:头太多会怎样?
如果我们把头数从 8 增加到 64,每个头只有 512/64=8 维:
- 优点:每个子空间更专门化,可能学到更细粒度的模式
- 缺点:每个头太窄,表达能力受限
- 实际问题:头之间可能学不到有用的差异,都是在拟合噪声
所以 8 个头是一个经验性的好选择,平衡了"分工"和"表达能力"。
What-If 实验三:不同头学到了什么?
学界对 Transformer 不同头的作用有一些有趣的发现:
- 有些头专注于语法关系:主语-动词一致、修饰关系
- 有些头专注于语义关系:同义词、指代关系
- 有些头专注于位置关系:相邻词、远距离依赖
- 有些头几乎是"空的":注意力分布很均匀,没什么信息
这说明多头注意力确实在分工,而且有些分工是天生的、不需要人为干预。
本节要点
- 单头的问题:表示空间竞争,不同类型信息相互干扰
- 头太多的风险:每个头太窄,表达能力受限,可能过拟合噪声
- 8头的经验:平衡了分工专门化和整体表达能力
- 头的多样性:不同头会自然学会关注不同类型的相关性
FAQ(20问)
以下是关于多头注意力机制的常见问题:
Q1. 多头注意力中的"头"是什么概念?
一句话结论:头是一组独立的注意力计算单元,每个头有自己的一套 QKV 投影。展开:想象你有 8 个侦探,每个人都看同一份情报,但用不同的"眼镜"(投影矩阵)来看。有的看人际关系,有的看时间线,最后汇总大家的发现。
Q2. 多头注意力和单头注意力在数学上有什么区别?
一句话结论:单头是 d_model×d_model 的投影,多头是 h 个 d_k×d_k 的投影拼接。展开:单头的 QKV 投影把 512 维变成 512 维。多头把 512 维拆成 8 组 64 维,每组独立计算后再拼接。数学上表达能力是等价的,但分头让优化更容易。
Q3. 头数(num_heads)通常怎么选?
一句话结论:d_model=512 时用 8 头,d_model=768 时用 12 头,原则是 d_model % num_heads == 0。展开:常见选择是让每个头的维度 d_k 在 32-128 之间。BERT-base 用 12 头(每头 64 维),BERT-large 用 16 头(每头 64 维)。
Q4. 每个头到底学到了什么?
一句话结论:不同头会自然学会关注不同类型的相关性。展开:学界发现有些头专注语法,有些专注语义,有些专注位置。但具体每个头学到了什么,通常需要可视化注意力权重来观察。
Q5. 多头注意力能并行计算吗?
一句话结论:可以,所有头的计算可以合并成矩阵运算一次性完成。展开:虽然逻辑上是分头计算,但实际实现中会把所有头的 QKV 拼接成大矩阵,用一次矩阵乘法搞定所有头的计算,效率很高。
Q6. 多头注意力和卷积有什么区别?
一句话结论:卷积是局部连接,多头注意力是全局连接。展开:CNN 的卷积核只看局部窗口(如 3×3),而注意力能看到序列中任意位置的关系。多头让注意力可以同时学到多种类型的全局关系。
Q7. 多头注意力和自注意力是什么关系?
一句话结论:多头注意力是自注意力的扩展,自注意力是单头注意力的特例。展开:自注意力(Self-Attention)指 Q=K=V 的注意力。多头自注意力(Multi-Head Self-Attention)就是用多个头来计算自注意力。
Q8. 为什么多头注意力比单头好?
一句话结论:分而治之,让每个头专门学习一种类型的相关性。展开:就像一个团队比一个人强,每个头专注于自己的任务,最后汇总比单打独斗效果更好。
Q9. 多头注意力会增加多少计算量?
一句话结论:QKV 投影和输出投影的计算量不变,只是实现方式不同。展开:从矩阵乘法的角度看,多头和单头的 FLOPs 差不多。但多头让模型更容易优化,间接提升了效果。
Q10. 多头注意力可以用在哪些地方?
一句话结论:任何需要建模序列内关系的地方。展开:文本分类、机器翻译、问答系统、图像处理、语音识别等。Transformer 的编码器和解码器都大量使用多头注意力。
Q11. 如何判断多头注意力是否过拟合?
一句话结论:看验证集损失是否上升,或者不同头的注意力分布是否变得很相似。展开:如果所有头都学到了相似的东西,说明模型在"偷懒",没有充分利用多头的分工优势。
Q12. 可以动态调整头数吗?
一句话结论:可以,用剪枝技术移除不重要的头。展开:有些研究发现部分头是冗余的,可以通过剪枝移除而不影响效果。这对于模型压缩很有用。
Q13. 不同层的头会学不同的东西吗?
一句话结论:是的,浅层头学语法,深层头学语义。展开:和 CNN 类似,Transformer 的底层头关注词汇和语法,高层头关注语义和推理。
Q14. 多头注意力和 Cross Attention 有什么区别?
一句话结论:Self-Attention 的 Q=K=V,Cross Attention 的 Q≠K≠V。展开:Cross Attention 用于解码器和编码器之间,让解码器"看"编码器的输出。Query 来自解码器,Key 和 Value 来自编码器。
Q15. 为什么有些头几乎什么都不学?
一句话结论:这些头可能是冗余的,或者任务不需要这种类型的相关性。展开:学界研究发现部分头确实学到的模式很稀疏(注意力分布接近均匀),这可能是因为:1)这些头在做冗余备份;2)其他头已经学到了这种模式;3)任务本身不需要这种相关性。可以通过剪枝移除不重要的头来压缩模型。
Q16. 多头注意力如何调试?
一句话结论:可视化注意力权重,观察每个头关注的位置。展开:画出每个头的注意力热力图,看是否和预期一致。比如"主语"的头应该更关注"动词"。
Q17. 多头注意力可以用在非序列数据上吗?
一句话结论:可以,把任意数据看作"序列"即可。展开:图像可以看作像素的序列,图数据可以看作节点和边的序列。多头注意力在这些领域都有应用(如 Vision Transformer)。
Q18. 头的数量和模型大小有关系吗?
一句话结论:有,大模型通常用更多的头。展开:BERT-base(110M 参数)用 12 头,BERT-large(340M 参数)用 16 头。GPT-3(175B 参数)用 96 头。
Q19. 多头注意力在 SRE 场景有什么用?
一句话结论:分析日志时,不同头关注不同类型的故障信号。展开:头0关注错误关键词,头1关注服务调用关系,头2关注时间模式。这样可以更全面地理解故障根因。
Q20. 多头注意力有哪些变体?
一句话结论:分组注意力、稀疏注意力、线性注意力等。展开:为了降低计算复杂度,出现了各种优化版本。分组注意力把头分组计算,稀疏注意力只关注部分位置,线性注意力用核函数近似。
核心结论
- 分而治之:多个头分工合作,每个头学一种类型的相关性
- 8头经验:d_model/8 是常用选择,平衡专门化和表达力
- 自然分工:不同头会自己学会关注不同类型的模式
- 生产用官方:实际开发直接用 torch.nn.MultiheadAttention
Roadmap预告
下期预告:《残差连接与层归一化:为什么 Transformer 能堆这么深?》
预告内容:
- 残差连接:让信息可以直接"跳过"层传播
- 层归一化:稳定训练、加速收敛
- 为什么 ResNet 的思想可以用在 Transformer 上
- What-If:如果没有残差,Transformer 能训练吗?
敬请期待!

浙公网安备 33010602011771号