共计 2768 个字符,预计需要花费 7 分钟才能阅读完成。
背景介绍
多头注意力机制(Multi-head Attention)是 Transformer 架构中的核心组件,最早由 Vaswani 等人在 2017 年的论文《Attention is All You Need》中提出。它的核心思想是将注意力机制并行化,通过多个 ” 头 ”(head)来捕捉输入序列中不同位置之间的多种依赖关系。

与传统 RNN 或 CNN 相比,多头注意力机制具有几个显著优势:
- 并行计算能力强,适合现代 GPU 架构
- 能够直接建模长距离依赖关系
- 每个注意力头可以学习不同的关注模式
- 计算复杂度相对序列长度是线性的
技术对比:多头注意力 vs 单头注意力
传统单头注意力机制可以看作是多头注意力的特例(head_num=1)。两者主要区别在于:
- 表示能力:多头注意力可以同时关注不同位置的多种关系模式,而单头只能学习一种模式
- 计算效率:虽然多头增加了参数,但通过并行计算可以保持高效
- 泛化性能:多头机制通常表现出更好的泛化能力
- 解释性:可以分析不同头学习到的注意力模式
核心实现解析
下面我们分步骤解析多头注意力的实现过程,并提供详细的 PyTorch 代码示例。
1. 输入投影
首先,我们需要将输入投影到查询 (Q)、键(K) 和值 (V) 空间:
import torch
import torch.nn as nn
class MultiHeadAttention(nn.Module):
def __init__(self, embed_dim, num_heads):
super().__init__()
self.embed_dim = embed_dim
self.num_heads = num_heads
self.head_dim = embed_dim // num_heads
# 确保 embed_dim 能被 num_heads 整除
assert self.head_dim * num_heads == embed_dim, "embed_dim 必须能被 num_heads 整除"
# 线性变换层
self.q_proj = nn.Linear(embed_dim, embed_dim)
self.k_proj = nn.Linear(embed_dim, embed_dim)
self.v_proj = nn.Linear(embed_dim, embed_dim)
# 输出投影
self.out_proj = nn.Linear(embed_dim, embed_dim)
2. 拆分多头
将投影后的 Q、K、V 矩阵拆分为多个头:
def split_heads(self, x, batch_size):
"""将 embed_dim 拆分为 num_heads x head_dim"""
x = x.view(batch_size, -1, self.num_heads, self.head_dim)
return x.permute(0, 2, 1, 3) # (batch_size, num_heads, seq_len, head_dim)
3. 注意力计算
计算缩放点积注意力:
def scaled_dot_product_attention(q, k, v, mask=None):
"""计算注意力权重和上下文向量"""
d_k = q.size(-1) # 获取 head_dim
# 计算 QK^T/sqrt(d_k)
scores = torch.matmul(q, k.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k))
# 可选:应用 mask
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
# softmax 归一化
attn_weights = torch.softmax(scores, dim=-1)
# 加权求和
output = torch.matmul(attn_weights, v)
return output, attn_weights
4. 合并多头输出
def combine_heads(self, x, batch_size):
"""合并多头输出"""
x = x.permute(0, 2, 1, 3) # (batch_size, seq_len, num_heads, head_dim)
return x.reshape(batch_size, -1, self.embed_dim)
5. 完整前向传播
def forward(self, query, key, value, mask=None):
batch_size = query.size(0)
# 线性投影
q = self.q_proj(query)
k = self.k_proj(key)
v = self.v_proj(value)
# 拆分多头
q = self.split_heads(q, batch_size)
k = self.split_heads(k, batch_size)
v = self.split_heads(v, batch_size)
# 计算注意力
attn_output, attn_weights = scaled_dot_product_attention(q, k, v, mask)
# 合并多头
attn_output = self.combine_heads(attn_output, batch_size)
# 输出投影
output = self.out_proj(attn_output)
return output, attn_weights
性能考量
头数选择的影响
- 模型容量:头数越多,模型表示能力越强,但可能导致过拟合
- 计算开销:头数增加会线性增加计算量
- 并行效率:合理头数可以利用 GPU 并行计算优势
- 经验法则:通常 embed_dim 的平方根是合理的头数起点
维度设计准则
embed_dim应该能被num_heads整除head_dim通常选择 32-128 之间的值- 总计算量与
embed_dim和序列长度的平方成正比
避坑指南
- 维度不匹配错误
- 确保输入维度与 embed_dim 一致
- 检查拆分 / 合并操作后的维度
-
使用
assert语句验证关键维度 -
梯度消失 / 爆炸
- 适当初始化权重(如 Xavier 初始化)
- 使用 Layer Normalization
-
添加残差连接
-
计算效率问题
- 避免不必要的张量复制
- 使用
torch.baddbmm优化大矩阵乘法 -
考虑使用 Flash Attention 等优化实现
-
数值稳定性
- 确保 softmax 前的数值范围合理
- 添加微小 epsilon 防止除零错误
思考与实践
为了加深理解,建议尝试以下实验:
- 固定
embed_dim=512,比较num_heads=[4,8,16]时的模型效果和计算时间 - 可视化不同头的注意力权重,分析它们关注的不同模式
- 尝试实现带掩码的多头注意力,用于解码器自回归生成
- 将实现与 PyTorch 内置的
nn.MultiheadAttention进行性能对比
多头注意力机制虽然概念简单,但在实际实现中有许多细节需要考虑。希望通过本文的讲解,你能掌握其核心原理并能够灵活应用到自己的项目中。
正文完
发表至: 人工智能
近两天内
