共计 2681 个字符,预计需要花费 7 分钟才能阅读完成。
问题背景:为什么需要自注意力机制?
在自然语言处理中,传统的 RNN 和 LSTM 模型在处理长序列时存在两个主要问题:

- 梯度消失 / 爆炸问题:随着序列长度增加,RNN 在反向传播时梯度难以有效传递
- 顺序计算限制:无法并行处理序列,训练速度慢
而自注意力机制通过计算序列中所有位置之间的关系,一次性获取全局信息。2017 年提出的 Transformer 模型首次完全基于注意力机制,彻底解决了这些问题。
数学原理:Scaled Dot-Product Attention
自注意力核心计算公式如下:
$$Attention(Q,K,V)=softmax(\frac{QK^T}{\sqrt{d_k}})V$$
其中:
- $Q$ (Query)、$K$ (Key)、$V$ (Value) 分别由输入序列通过线性变换得到
- $d_k$ 是 Key 向量的维度
- 除以 $\sqrt{d_k}$ 是为了防止点积结果过大导致 softmax 梯度消失
多头注意力机制图解
多头注意力的核心思想是将注意力计算拆分为多个 ” 头 ” 并行处理:
- 线性投影:将 Q、K、V 通过不同的线性层投影到低维空间
- 原始维度:[batch_size, seq_len, model_dim]
-
投影后:[batch_size, seq_len, num_heads * head_dim]
-
头拆分:使用 view 操作重组张量
- 变形后:[batch_size, seq_len, num_heads, head_dim]
-
转置后:[batch_size, num_heads, seq_len, head_dim]
-
并行计算:每个头独立计算注意力权重
PyTorch 实现详解
import torch
import torch.nn as nn
import torch.nn.functional as F
class MultiHeadAttention(nn.Module):
def __init__(self, model_dim=512, num_heads=8, dropout=0.1):
super().__init__()
assert model_dim % num_heads == 0
self.head_dim = model_dim // num_heads
self.num_heads = num_heads
# QKV 投影矩阵
self.q_linear = nn.Linear(model_dim, model_dim)
self.k_linear = nn.Linear(model_dim, model_dim)
self.v_linear = nn.Linear(model_dim, model_dim)
self.dropout = nn.Dropout(dropout)
self.out_linear = nn.Linear(model_dim, model_dim)
def forward(self, x, mask=None):
batch_size = x.size(0)
# 1. 线性投影 [batch, seq_len, model_dim] -> [batch, seq_len, model_dim]
q = self.q_linear(x)
k = self.k_linear(x)
v = self.v_linear(x)
# 2. 头拆分和维度调整
q = q.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2) # [batch, heads, seq_len, head_dim]
k = k.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
v = v.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
# 3. 计算注意力得分
scores = torch.matmul(q, k.transpose(-2, -1)) / (self.head_dim ** 0.5)
# 4. 应用 mask(如 padding mask)if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
# 5. softmax 归一化
attn_weights = F.softmax(scores, dim=-1)
attn_weights = self.dropout(attn_weights)
# 6. 加权求和
output = torch.matmul(attn_weights, v) # [batch, heads, seq_len, head_dim]
# 7. 合并多头输出
output = output.transpose(1, 2).contiguous() # [batch, seq_len, heads, head_dim]
output = output.view(batch_size, -1, self.num_heads * self.head_dim)
return self.out_linear(output)
性能优化技术
Flash Attention
传统注意力实现存在两大瓶颈:
- 内存访问瓶颈:标准实现需要多次读写 HBM(高带宽内存)
- 冗余计算:softmax 需要分步计算
Flash Attention 通过以下技术提升性能:
- 平铺(Tiling):将注意力计算分块处理,减少内存访问
- 重计算(Recomputation):反向传播时重新计算中间结果,减少内存占用
- 融合内核(Fused Kernel):将多个操作合并为一个 CUDA 内核
实际测试表明,Flash Attention 可以带来 2 - 4 倍的训练加速。
常见错误与避坑指南
- 维度转置错误
- 错误:计算 QK^T 时忘记转置 K 矩阵
- 现象:注意力权重矩阵形状不匹配
-
修复:确保使用 k.transpose(-2, -1)
-
mask 应用错误
- 错误:padding mask 应用在 softmax 之后
- 现象:padding 位置仍然参与计算
-
修复:在 softmax 前应用 mask,并用极大负值填充
-
维度拆分错误
- 错误:head_dim * num_heads != model_dim
- 现象:view 操作抛出形状错误
- 修复:初始化时检查维度可整除性
开放性问题
- 头数是否越多越好?
- 实验表明,头数超过一定阈值后性能提升有限
-
不同任务可能需要不同的头数配置
-
如何动态分配头维度?
- 是否可以学习不同头的重要性
- 混合专家 (MoE) 思路是否适用
结语
多头注意力机制是 Transformer 架构的核心创新,理解其实现细节对于优化模型性能至关重要。希望本文的数学推导和代码实现能帮助你深入理解这一关键技术。在实际应用中,建议从小规模实验开始,逐步验证各组件的行为,避免常见实现陷阱。
正文完
