共计 2598 个字符,预计需要花费 7 分钟才能阅读完成。
背景与痛点:为什么需要注意力机制?
传统 RNN/LSTM 处理序列数据时存在两大瓶颈:

- 长程依赖丢失:随着序列长度增加,早期信息在反向传播时梯度逐渐消失(vanishing gradient 问题)
- 固定编码瓶颈:Encoder 必须将整个输入序列压缩为固定长度的上下文向量,信息压缩必然导致细节丢失
2014 年《Neural Machine Translation by Jointly Learning to Align and Translate》论文首次提出注意力机制,核心思想是让模型动态关注当前任务相关的输入部分。例如翻译 ”Hello world” 时,生成 ”world” 只需聚焦第二个单词而非整个句子。
核心概念:注意力机制的本质
注意力本质是一种 可学习的权重分配策略,其数学表达包含三个核心组件:
- Query(Q):当前需要计算输出的目标位置(如解码器当前时间步)
- Key(K):输入序列的各个位置标识(如编码器所有时间步)
- Value(V):对应 Key 的实际内容信息
计算分为两步:
- 通过 Q 与 K 的相似度计算注意力权重(常见方法见下节)
- 对 V 进行加权求和得到输出
公式化表示为:
Attention(Q, K, V) = softmax(QK^T/√d_k)V
其中 d_k 是 Key 的维度,缩放因子用于防止点积过大导致 softmax 梯度消失。
常见注意力机制对比
1. 加性注意力(Additive Attention)
- 最早出现在 Bahdanau 的 NMT 论文
- 计算方式:
score(q,k) = v^T tanh(W_q q + W_k k) - 优点:适用于 query 和 key 维度不同的场景
- 缺点:需学习额外参数矩阵,计算量较大
2. 点积注意力(Dot-Product Attention)
- Vaswani 在 Transformer 中推广
- 计算方式:
score(q,k) = q^T k - 优点:计算高效,无需额外参数
- 缺点:需保证 q 和 k 维度相同,当 d_k 较大时方差增大需缩放
3. 多头注意力(Multi-Head Attention)
- Transformer 的核心创新
- 并行计算多组注意力并将结果拼接
- 优势:
- 允许模型同时关注不同子空间的信息
- 类似 CNN 中多通道的概念
自注意力机制深度解析
自注意力(Self-Attention)是 Transformer 的基石,其特殊性在于:
- Q,K,V 均来自同一输入序列(区别于传统注意力中 Q 来自解码器)
- 通过三个可学习矩阵 W_Q, W_K, W_V 实现线性变换
具体计算流程:
- 将输入序列 X(n×d_model)分别乘以 W_Q,W_K,W_V 得到 Q,K,V
- 计算 QK^T 并缩放,得到 n×n 的注意力矩阵
- 按行 softmax 归一化后乘以 V
- 多头情况下拼接各头结果并通过 WO 矩阵融合
优势体现在:
- 全局依赖建模:每个位置直接关联所有其他位置
- 并行计算:摆脱 RNN 的序列依赖
- 可解释性:注意力权重可视化显示特征关联
PyTorch 实现关键代码
import torch
import torch.nn as nn
import torch.nn.functional as F
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, n_heads=8):
super().__init__()
assert d_model % n_heads == 0
self.d_k = d_model // n_heads
self.n_heads = n_heads
# 线性变换层
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.W_o = nn.Linear(d_model, d_model)
def forward(self, x):
batch_size = x.size(0)
# 线性变换并分头 [batch, seq_len, d_model] -> [batch, seq_len, n_heads, d_k]
q = self.W_q(x).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
k = self.W_k(x).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
v = self.W_v(x).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
# 计算缩放点积注意力
scores = torch.matmul(q, k.transpose(-2, -1)) / (self.d_k ** 0.5)
attn = F.softmax(scores, dim=-1)
# 加权求和并合并多头
out = torch.matmul(attn, v).transpose(1, 2).contiguous()
out = out.view(batch_size, -1, self.n_heads * self.d_k)
return self.W_o(out)
实际应用案例
BERT 中的注意力机制
- 使用 12/24 层 Transformer 编码器
- 每层包含 12/16 个注意力头
- 采用全连接自注意力(未使用因果掩码)
- 预训练时通过 MLM 任务学习双向表征
GPT 系列模型
- 使用解码器结构的 Transformer
- 通过注意力掩码实现自回归生成
- GPT- 3 的单头注意力维度达到 128
避坑指南
- 梯度消失问题:
- 现象:深层 Transformer 训练困难
-
解决:使用残差连接 +LayerNorm(Pre-LN 结构效果更佳)
-
注意力矩阵爆炸:
- 现象:长序列时 QK^T 值过大导致 softmax 饱和
-
解决:务必使用缩放因子 1 /√d_k
-
内存溢出:
- 现象:处理长文本时 O(n^2)复杂度耗尽显存
- 解决:
- 采用内存高效的注意力实现(如 FlashAttention)
- 使用稀疏注意力或分块计算
性能优化考量
- 计算复杂度:
- 自注意力:O(n^2 d)(n 为序列长度,d 为特征维度)
- RNN:O(n d^2)
-
当 n < d 时注意力更高效(这也是 GPT 处理长文本的挑战)
-
内存占用:
- 注意力矩阵需要存储 n×n 的中间结果
-
1K 长度的序列单精度浮点就需要 4MB 显存
-
工程优化方向:
- 混合精度训练
- 内核融合技术
- 分布式计算
开放性问题
- 如何设计更高效的稀疏注意力模式?
- 动态调整注意力头数是否能提升模型适应性?
- 生物神经系统中的注意力机制对 AI 有何启发?
正文完
发表至: 未分类
近一天内
