共计 2213 个字符,预计需要花费 6 分钟才能阅读完成。
背景痛点
Attention 机制已经成为现代深度学习模型的基石,从 NLP 领域的 Transformer 到 CV 领域的 Vision Transformer,其重要性不言而喻。然而,在使用现成框架(如 PyTorch)时,由于自动求导机制的便利性,许多开发者对 Attention 的反向传播细节一知半解。这种理解上的模糊性不仅限制了模型的调试能力,还可能影响性能优化和问题排查。

数学推导
正向计算过程
Attention 的正向计算可以表示为以下公式:
$$
\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V
$$
其中,$Q$, $K$, $V$ 分别代表查询(Query)、键(Key)和值(Value)矩阵,$d_k$ 是键的维度。
反向传播推导
为了理解反向传播,我们需要计算 Attention score 对 $Q$, $K$, $V$ 的偏导数。具体步骤如下:
- 计算 Attention score $A = \frac{QK^T}{\sqrt{d_k}}$
- 计算 softmax 输出 $S = \text{softmax}(A)$
- 最终输出 $O = SV$
在反向传播过程中,我们需要计算 $\frac{\partial L}{\partial Q}$, $\frac{\partial L}{\partial K}$, 和 $\frac{\partial L}{\partial V}$,其中 $L$ 是损失函数。
Scaled Dot-Product 的影响
scale 因子 $\sqrt{d_k}$ 在反向传播中起到了稳定梯度幅度的作用,防止梯度爆炸或消失。其反向传播贡献可以通过链式法则计算:
$$
\frac{\partial L}{\partial Q} = \frac{\partial L}{\partial A} \cdot \frac{\partial A}{\partial Q} = \frac{\partial L}{\partial A} \cdot \frac{K}{\sqrt{d_k}}
$$
PyTorch 实现
下面是一个手动实现 Attention 前向和反向传播的 PyTorch 代码示例:
import torch
import torch.nn.functional as F
class ManualAttention(torch.autograd.Function):
@staticmethod
def forward(ctx, Q, K, V, scale_factor):
# 计算 scaled dot-product attention
scores = torch.matmul(Q, K.transpose(-2, -1)) / scale_factor
attn_weights = F.softmax(scores, dim=-1)
output = torch.matmul(attn_weights, V)
# 保存反向传播需要的中间结果
ctx.save_for_backward(Q, K, V, attn_weights, scale_factor)
return output
@staticmethod
def backward(ctx, grad_output):
Q, K, V, attn_weights, scale_factor = ctx.saved_tensors
# 计算对 V 的梯度
grad_V = torch.matmul(attn_weights.transpose(-2, -1), grad_output)
# 计算对 attention weights 的梯度
grad_attn = torch.matmul(grad_output, V.transpose(-2, -1))
# 计算对 scores 的梯度(包含 softmax 梯度)grad_scores = grad_attn * attn_weights
grad_scores = grad_scores - attn_weights * torch.sum(grad_attn * attn_weights, dim=-1, keepdim=True)
# 计算对 Q 和 K 的梯度
grad_Q = torch.matmul(grad_scores, K) / scale_factor
grad_K = torch.matmul(grad_scores.transpose(-2, -1), Q) / scale_factor
return grad_Q, grad_K, grad_V, None
性能优化
Flash Attention
Flash Attention 通过重新计算 attention 分数而不是存储中间结果来优化内存使用,这对反向传播特别有益,因为它减少了内存带宽需求。
梯度检查点
对于长序列处理,梯度检查点技术可以通过牺牲部分计算时间来节省内存,这对反向传播过程中的内存优化至关重要。
避坑指南
- Softmax 梯度计算 :确保正确实现 softmax 的梯度计算,这是常见的错误来源。
- 数值稳定性 :在计算 softmax 时考虑数值稳定性,可以使用 log-softmax 技巧。
- 混合精度训练 :在使用混合精度训练时,注意保持足够的精度在关键计算步骤,如 softmax 和矩阵乘法。
思考题
- 如何修改上述实现来支持多头注意力(Multi-Head Attention)的反向传播?
- 在极端长序列(如 10000+ tokens)场景下,哪些技术可以优化 Attention 反向传播的内存使用?
- 如何验证你的手动实现的反向传播与 PyTorch 自动求导结果完全一致?
