共计 2516 个字符,预计需要花费 7 分钟才能阅读完成。
背景痛点
在自然语言处理任务中,传统的注意力机制存在两个主要瓶颈:

-
计算复杂度高:原始注意力矩阵的计算复杂度为 O(n²),当处理长序列时(如文档级文本),显存和计算时间会急剧增加。
-
表征能力有限:单头注意力只能学习一种注意力模式,难以捕捉词语间复杂的依赖关系(如同义词、指代关系等)。
技术对比
参数量分析
- 单头注意力参数量:$3d_{model}^2$(Q/K/ V 三个矩阵)
- 多头注意力参数量:$3d_{model}^2$(与头数 h 无关,因为通过分片实现)
计算效率
数学上,多头注意力的核心思想是将 Q /K/ V 拆分为 h 个子空间:
$$\text{MultiHead}(Q,K,V) = \text{Concat}(head_1,…,head_h)W^O$$
其中每个头的计算保持 $d_k = d_{model}/h$,总计算量与单头相当但表征能力更强。
核心实现
完整 PyTorch 实现(2.4.3 风格)
import torch
import torch.nn as nn
import math
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, h=8, dropout=0.1):
super().__init__()
assert d_model % h == 0, "d_model 必须能被 h 整除"
self.d_k = d_model // h
self.h = h
# 线性变换层(未分头前的全连接)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)
# 输出层和 Dropout
self.W_o = nn.Linear(d_model, d_model)
self.dropout = nn.Dropout(dropout)
# 缩放因子
self.scale = 1 / math.sqrt(self.d_k)
def forward(self, x, mask=None):
"""
输入:
x - [batch_size, seq_len, d_model]
mask - [batch_size, seq_len, seq_len]
输出:
[batch_size, seq_len, d_model]
"""
batch_size = x.size(0)
# 1. 线性变换并分头
Q = self.W_q(x) # [B,T,d_model]
K = self.W_k(x)
V = self.W_v(x)
# 分头操作(使用 chunk 实现)Q = torch.chunk(Q, self.h, dim=-1) # h 个[B,T,d_k]
K = torch.chunk(K, self.h, dim=-1)
V = torch.chunk(V, self.h, dim=-1)
# 2. 缩放点积注意力
attention_outputs = []
for q, k, v in zip(Q, K, V):
# 计算注意力分数
scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale # [B,T,T]
# 掩码处理(可选)if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
# Softmax 和 Dropout
attn = torch.softmax(scores, dim=-1)
attn = self.dropout(attn)
# 加权求和
out = torch.matmul(attn, v) # [B,T,d_k]
attention_outputs.append(out)
# 3. 合并多头结果
concat = torch.cat(attention_outputs, dim=-1) # [B,T,d_model]
output = self.W_o(concat)
return output
性能优化
GPU 显存占用分析
- 当 h 从 8 增加到 16 时(保持 d_model=512):
- 每个头的 d_k 从 64 降到 32
- 显存占用增加约 15%(主要来自中间注意力矩阵)
并行计算优化
使用 torch.baddbmm 替代循环计算:
# 原版(循环计算)for q, k, v in zip(Q, K, V):
scores = torch.matmul(q, k.transpose(-2, -1))
# 优化版(批量计算)Q = torch.stack(Q, dim=1) # [B,h,T,d_k]
K = torch.stack(K, dim=1).transpose(-2, -1) # [B,h,d_k,T]
scores = torch.baddbmm(torch.zeros(batch_size, self.h, seq_len, seq_len, device=x.device),
Q, K,
beta=0, alpha=self.scale
)
避坑指南
梯度爆炸预防
-
初始化策略:对 Q / K 矩阵使用 Xavier 初始化
nn.init.xavier_uniform_(self.W_q.weight, gain=1/math.sqrt(2)) nn.init.xavier_uniform_(self.W_k.weight, gain=1/math.sqrt(2)) -
缩放因子:必须保持 $1/\sqrt{d_k}$ 的比例
变长序列处理
对 padding 部分添加负无穷掩码:
mask = (x != 0).unsqueeze(1).unsqueeze(2) # [B,1,1,T]
scores = scores.masked_fill(mask == 0, -1e9)
思考题
- 头数增加是否总能提升模型性能?如何通过实验确定最优头数?
- 不同注意力头是否真的学习了不同的模式?如何可视化验证?
- 当 d_model 固定时,是否存在头数 h 的理论上限?
实践建议
在实际项目中,建议先使用 h = 8 的默认配置,然后:
1. 监控各头的注意力分布(使用 torch.topk 统计)
2. 测试不同头数下的验证集表现
3. 对长序列任务(如文档分类),可适当减少头数以降低显存消耗
通过本文的实现,读者应能理解多头注意力的核心思想,并掌握工业级实现的优化技巧。建议在 Transformer 架构中实际测试该模块,观察不同超参数对最终效果的影响。
正文完
发表至: 未分类
近两天内
