共计 2243 个字符,预计需要花费 6 分钟才能阅读完成。
背景:为什么需要理解 Attention 的反向传播
Attention 机制是 Transformer 模型的核心组件,它通过计算查询(Query)、键(Key)和值(Value)之间的相关性,实现了对输入序列的动态加权。在训练过程中,反向传播算法的效率直接影响模型的收敛速度。理解 Attention 的反向传播不仅有助于调试模型,还能为自定义 Attention 变体提供理论基础。

数学推导:从标量到矩阵
1. 前向传播回顾
Attention 的前向传播可以表示为:
$$
\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V
$$
其中 $Q \in \mathbb{R}^{n\times d_k}$, $K \in \mathbb{R}^{m\times d_k}$, $V \in \mathbb{R}^{m\times d_v}$。
2. 标量分量求导
假设损失函数为 $L$,我们先对单个元素求导。令 $S = QK^T/\sqrt{d_k}$,$P = \text{softmax}(S)$,则输出 $O = PV$。
对于 $O_{ij}$ 的梯度 $\frac{\partial L}{\partial O_{ij}}$,我们需要计算:
$$
\frac{\partial L}{\partial V_{kl}} = \sum_{i,j} \frac{\partial L}{\partial O_{ij}} \frac{\partial O_{ij}}{\partial V_{kl}} = \sum_i \frac{\partial L}{\partial O_{il}} P_{ik}
$$
3. Softmax 梯度计算
Softmax 的梯度需要特别注意。对于 $P_{ij} = \frac{e^{S_{ij}}}{\sum_k e^{S_{ik}}}$,其导数为:
$$
\frac{\partial P_{ij}}{\partial S_{kl}} = P_{ij}(\delta_{ik}\delta_{jl} – P_{il})
$$
其中 $\delta$ 是 Kronecker delta 函数。
4. 矩阵化梯度公式
将上述标量结果转换为矩阵运算:
$$
\frac{\partial L}{\partial V} = P^T \frac{\partial L}{\partial O}
$$
$$
\frac{\partial L}{\partial S} = \frac{\partial L}{\partial P} \odot (P – P \otimes P)
$$
其中 $\odot$ 表示逐元素乘法,$\otimes$ 表示外积。
代码验证:PyTorch 实现
import torch
import torch.nn.functional as F
def attention_backward(dO, Q, K, V, mask=None):
"""
dO: gradient of loss w.r.t output [n, d_v]
Q: query matrix [n, d_k]
K: key matrix [m, d_k]
V: value matrix [m, d_v]
"""
dk = Q.size(-1)
S = torch.matmul(Q, K.transpose(-2, -1)) / (dk ** 0.5) # [n, m]
if mask is not None:
S = S.masked_fill(mask == 0, -1e9)
P = F.softmax(S, dim=-1) # [n, m]
# Gradient w.r.t V
dV = torch.matmul(P.transpose(-2, -1), dO) # [m, d_v]
# Gradient w.r.t P
dP = torch.matmul(dO, V.transpose(-2, -1)) # [n, m]
# Gradient w.r.t S
dS = P * (dP - torch.sum(P * dP, dim=-1, keepdim=True)) # [n, m]
# Gradient w.r.t Q and K
dQ = torch.matmul(dS, K) / (dk ** 0.5) # [n, d_k]
dK = torch.matmul(dS.transpose(-2, -1), Q) / (dk ** 0.5) # [m, d_k]
return dQ, dK, dV
性能优化与避坑指南
计算复杂度分析
- 原始实现复杂度为 $O(n^2d + nmd)$,其中 $n$ 是序列长度,$d$ 是特征维度
- 内存消耗主要来自存储中间矩阵 $S$ 和 $P$,大小为 $O(nm)$
常见实现错误
- 梯度爆炸 :未对 $QK^T$ 进行缩放,导致 softmax 输入值过大
- Mask 处理不当 :在计算 softmax 前未正确应用 mask,导致无效位置参与计算
- 数值稳定性 :未对 softmax 做 log-sum-exp 优化,长序列时可能出现 NaN
Flash Attention 优化
Flash Attention 通过以下技术提升效率:
1. 分块计算,减少内存访问
2. 融合 kernel,减少中间结果存储
3. 在线 softmax,避免存储完整的 attention 矩阵
思考题
如何将上述推导扩展到多头 Attention 情况?需要考虑:
1. 每个头的梯度如何聚合
2. 投影矩阵的梯度计算
3. 不同头之间的梯度流动路径
总结
理解 Attention 的反向传播需要耐心地拆解每一步的矩阵运算。通过标量推导到矩阵化的转换,我们不仅能验证 autograd 的结果,还能针对特定场景进行优化。建议读者在理解单头 Attention 的基础上,尝试推导多头情况的梯度计算,这将帮助深入理解 Transformer 的运作机制。
