共计 2407 个字符,预计需要花费 7 分钟才能阅读完成。
背景:为什么需要自注意力机制
在 Transformer 出现之前,RNN 和 LSTM 是处理序列数据的主流架构。但它们存在两个致命缺陷:

- 顺序计算:必须逐个处理时间步,难以并行化
- 长程依赖衰减:信息传递路径过长时,梯度容易消失 / 爆炸
自注意力机制 (Self-Attention) 通过计算序列元素间的关联权重,实现了:
- 任意位置直接交互(解决长程依赖)
- 矩阵运算天然可并行(提升计算效率)
- 动态权重分配(比固定窗口的 CNN 更灵活)
数学原理:缩放点积注意力
核心公式如下:
$$\text{Attention}(Q,K,V) = \text{softmax}(\frac{QK^T}{\sqrt{d_k}})V$$
其中:
– $Q$ (Query), $K$ (Key), $V$ (Value) 分别由输入线性变换得到
– $d_k$ 是 key 向量的维度,缩放因子防止点积过大导致 softmax 饱和
具体计算步骤:
- 计算相似度分数:$S = QK^T$(形状:[batch, heads, seq_len, seq_len])
- 缩放并归一化:$P = \text{softmax}(S/\sqrt{d_k})$
- 加权求和:$O = PV$
PyTorch 完整实现
import torch
import torch.nn as nn
import einops
from torch.nn import 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
# 合并计算 QKV 的线性变换
self.qkv_proj = nn.Linear(d_model, 3*d_model)
self.out_proj = nn.Linear(d_model, d_model)
def forward(self, x, mask=None):
batch_size, seq_len = x.shape[:2]
# 生成 QKV 并分头 [B,L,3*D] -> [B,L,H,3*D/H]
qkv = self.qkv_proj(x)
qkv = einops.rearrange(qkv,
'b l (three h d) -> three b h l d',
three=3, h=self.n_heads)
q, k, v = qkv[0], qkv[1], qkv[2] # 各 [B,H,L,D_k]
# 计算注意力分数
scores = torch.matmul(q, k.transpose(-2,-1)) / (self.d_k ** 0.5)
# 处理 mask(如因果掩码)if mask is not None:
scores = scores.masked_fill(mask==0, float('-inf'))
# 注意力权重与 value 相乘
attn = F.softmax(scores, dim=-1)
output = torch.matmul(attn, v) # [B,H,L,D_k]
# 合并多头输出
output = einops.rearrange(output,
'b h l d -> b l (h d)')
return self.out_proj(output)
关键实现技巧:
- 使用
einops库简化张量 reshape 操作 - 合并 QKV 的线性投影减少一次矩阵乘法
- 支持传入 mask 处理不同注意力模式
工业级优化策略
1. Flash Attention
通过融合 kernel 技术,将注意力计算中的:
- 矩阵乘法
- Mask 处理
- Softmax
- 加权求和
合并为单个 CUDA kernel,显著减少内存读写次数。实测在 A100 上可获得 3 - 5 倍加速。
2. KV 缓存(Decoder 优化)
在自回归生成时,先前时间步的 KV 矩阵可缓存复用:
# 推理时缓存实现示例
kv_cache = None
def process_step(new_x):
global kv_cache
q = calc_q(new_x)
if kv_cache is None:
k, v = calc_kv(new_x)
kv_cache = (k, v)
else:
new_k, new_v = calc_kv(new_x)
k = torch.cat([kv_cache[0], new_k], dim=1)
v = torch.cat([kv_cache[1], new_v], dim=1)
kv_cache = (k, v)
# 计算当前步输出...
3. 显存与序列长度
注意力矩阵的内存占用为 $O(L^2)$,处理长序列时可考虑:
- 梯度检查点(trade-off 计算与显存)
- 块稀疏注意力(如 Longformer 的滑动窗口)
- 内存高效的注意力变体(如 Linformer)
避坑指南
梯度爆炸预防
- 初始化 QK 投影矩阵时缩小方差(如使用 $1/\sqrt{d_k}$ 缩放)
- 添加 LayerNorm 稳定训练
- 梯度裁剪(
torch.nn.utils.clip_grad_norm_)
混合精度训练
with torch.autocast(device_type='cuda', dtype=torch.float16):
output = attn_layer(inputs)
loss = criterion(output, targets)
scaler.scale(loss).backward() # 使用 GradScaler
scaler.step(optimizer)
scaler.update()
注意事项:
– 在 softmax 前保持 float32 精度
– 使用 AMP 自动管理精度转换
开放性问题
- 稀疏注意力设计:
- 如何平衡局部敏感性与全局信息捕获?
-
动态稀疏模式(如根据内容路由)是否比固定模式更优?
-
线性注意力工程化:
- 近似方法(如核函数技巧)在哪些场景下会失效?
- 如何避免特征映射带来的计算开销抵消收益?
自注意力机制仍在快速发展,期待更多创新解决其计算效率与表达能力之间的平衡问题。
正文完
发表至: 人工智能
近一天内
