共计 2349 个字符,预计需要花费 6 分钟才能阅读完成。
从 RNN 到 Transformer 的进化之路
在自然语言处理领域,处理长序列数据一直是个难题。传统 RNN(循环神经网络)虽然能处理序列,但存在两个致命缺陷:

- 梯度消失 / 爆炸问题:随着序列长度增加,RNN 难以有效传递长期依赖信息
- 顺序计算限制:必须逐个处理序列元素,无法并行化计算
Transformer 的提出彻底改变了这一局面。它的核心创新就是引入了 多头注意力机制(Multi-Head Attention),允许模型:
- 同时关注序列的所有位置
- 自动学习不同位置间的关系
- 实现完全并行的计算
理解 b,t,c 三个关键维度
在多头注意力实现中,所有张量都有三个基本维度:
- b(batch):一次处理的样本数量
- t(sequence length):序列的长度
- c(channel/dim):每个位置的特征维度
用一个简单的例子说明:假设我们处理一批英文句子,每个句子有 10 个单词,每个单词用 512 维向量表示,batch size 为 32,那么输入张量的形状就是(32, 10, 512)。
多头注意力的核心计算可以用公式表示:
$$Attention(Q,K,V) = softmax(\frac{QK^T}{\sqrt{d_k}})V$$
其中:
– $Q$ (Query),$K$ (Key),$V$ (Value)都是输入的不同线性变换
– $d_k$ 是 Key 的维度,用于缩放点积结果
PyTorch 实现详解
下面是一个完整的 MultiHeadAttention 类实现,包含详细的维度注释:
import torch
import torch.nn as nn
import torch.nn.functional as F
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
# 线性变换层
self.qkv_proj = nn.Linear(embed_dim, 3*embed_dim)
self.out_proj = nn.Linear(embed_dim, embed_dim)
def forward(self, x, mask=None):
"""
输入:
x: (batch_size, seq_len, embed_dim)
mask: (batch_size, seq_len, seq_len)
输出:
(batch_size, seq_len, embed_dim)
"""
batch_size, seq_len, embed_dim = x.shape
# 生成 Q,K,V [b,t,c] -> [b,t,3c]
qkv = self.qkv_proj(x)
# 分割成多头 [b,t,3c] -> [b,t,num_heads,3*head_dim]
qkv = qkv.reshape(batch_size, seq_len, self.num_heads, 3*self.head_dim)
# 分离 Q,K,V [b,t,num_heads,3*head_dim] -> 3*[b,num_heads,t,head_dim]
q, k, v = torch.chunk(qkv, 3, dim=-1)
q, k, v = [x.transpose(1, 2) for x in (q, k, v)]
# 计算注意力分数 [b,num_heads,t,t]
scores = torch.matmul(q, k.transpose(-2, -1)) / torch.sqrt(torch.tensor(self.head_dim))
# 应用 mask(如需要)if mask is not None:
scores = scores.masked_fill(mask == 0, float('-inf'))
# Softmax 归一化
attn_weights = F.softmax(scores, dim=-1)
# 加权求和 [b,num_heads,t,head_dim]
output = torch.matmul(attn_weights, v)
# 合并多头 [b,num_heads,t,head_dim] -> [b,t,num_heads*head_dim]
output = output.transpose(1, 2).reshape(batch_size, seq_len, embed_dim)
# 最终线性变换
output = self.out_proj(output)
return output
性能优化技巧
实际应用中,注意力计算可能成为性能瓶颈。以下是两个关键优化方向:
- Flash Attention:通过分块计算和 IO 感知算法,显著减少 GPU 显存访问
- 稀疏注意力:对长序列只计算关键位置间的注意力
Flash Attention 特别适合以下场景:
- 处理超长序列(>1024 tokens)
- 在有限显存的 GPU 上训练大模型
- 需要高吞吐量的推理场景
新手常见错误及解决方法
- 维度不匹配:
- 症状:RuntimeError 提示形状不兼容
-
解决:仔细检查所有 reshape 和 transpose 操作后的维度
-
忘记 scale:
- 症状:训练初期出现 NaN 损失
-
解决:确保除以 $\sqrt{d_k}$
-
mask 应用错误:
- 症状:模型在验证集表现异常
- 解决:确认 mask 在正确位置填充了
-inf
延伸思考与实验
- 头的数量如何影响性能?
-
实验:尝试将 num_heads 从 1 逐渐增加到 embed_dim 大小,观察验证集准确率变化
-
注意力模式的可视化
- 实验:选择特定输入,绘制不同头的注意力热力图,分析模式差异
多头注意力机制看似复杂,但通过拆解维度操作和逐步实现,完全可以掌握其精髓。建议读者动手实现一个简化版 Transformer,在实践中深化理解。
