共计 2906 个字符,预计需要花费 8 分钟才能阅读完成。
背景介绍
Attention 机制是 Transformer 模型的核心组件,它通过计算输入序列中各个位置的重要性权重,实现对不同位置信息的动态聚焦。理解 Attention 的反向传播过程对于调试和优化模型至关重要。本文将带大家从数学原理出发,逐步推导 Attention 层的反向传播公式,并通过 PyTorch 代码实现来加深理解。

数学推导
前向计算
Attention 的前向计算可以分为以下几步:
-
计算 Query(Q)、Key(K)、Value(V) 矩阵:
$$ Q = XW^Q, K = XW^K, V = XW^V $$ -
计算注意力分数:
$$ S = \frac{QK^T}{\sqrt{d_k}} $$ -
应用 Softmax 得到注意力权重:
$$ A = \text{softmax}(S) $$ -
加权求和得到输出:
$$ O = AV $$
反向传播推导
我们需要计算损失 L 对各个参数的梯度。按照链式法则,从输出逐步回传:
-
计算∂L/∂O,这是来自上一层的梯度
-
∂L/∂V 的计算:
$$ \frac{\partial L}{\partial V} = A^T \frac{\partial L}{\partial O} $$ -
∂L/∂A 的计算:
$$ \frac{\partial L}{\partial A} = \frac{\partial L}{\partial O} V^T $$ -
∂L/∂S 的计算(注意 Softmax 的梯度计算):
$$ \frac{\partial L}{\partial S_{ij}} = A_{ij}(\delta_{ij} – A_{ij}) \frac{\partial L}{\partial A_{ij}} $$ -
∂L/∂Q 和∂L/∂K 的计算:
$$ \frac{\partial L}{\partial Q} = \frac{1}{\sqrt{d_k}} \frac{\partial L}{\partial S} K $$
$$ \frac{\partial L}{\partial K} = \frac{1}{\sqrt{d_k}} Q^T \frac{\partial L}{\partial S} $$ -
最后计算对参数矩阵 W 的梯度:
$$ \frac{\partial L}{\partial W^Q} = X^T \frac{\partial L}{\partial Q} $$
$$ \frac{\partial L}{\partial W^K} = X^T \frac{\partial L}{\partial K} $$
$$ \frac{\partial L}{\partial W^V} = X^T \frac{\partial L}{\partial V} $$
PyTorch 实现
下面我们实现一个简化版的 Attention 层,包含手动反向传播:
import torch
import torch.nn as nn
import torch.nn.functional as F
class ManualAttention(nn.Module):
def __init__(self, d_model, d_k):
super().__init__()
self.W_q = nn.Linear(d_model, d_k, bias=False)
self.W_k = nn.Linear(d_model, d_k, bias=False)
self.W_v = nn.Linear(d_model, d_k, bias=False)
self.d_k = d_k
def forward(self, x):
Q = self.W_q(x)
K = self.W_k(x)
V = self.W_v(x)
S = torch.matmul(Q, K.transpose(-2, -1)) / (self.d_k ** 0.5)
A = F.softmax(S, dim=-1)
O = torch.matmul(A, V)
# 保存中间结果用于反向传播
self.cache = (Q, K, V, A)
return O
def manual_backward(self, dO):
Q, K, V, A = self.cache
# 计算各部分的梯度
dV = torch.matmul(A.transpose(-2, -1), dO)
dA = torch.matmul(dO, V.transpose(-2, -1))
# Softmax 梯度计算
dS = A * (dA - torch.sum(A * dA, dim=-1, keepdim=True))
dQ = torch.matmul(dS, K) / (self.d_k ** 0.5)
dK = torch.matmul(Q.transpose(-2, -1), dS) / (self.d_k ** 0.5)
# 计算参数梯度
dWq = torch.matmul(x.transpose(-2, -1), dQ)
dWk = torch.matmul(x.transpose(-2, -1), dK)
dWv = torch.matmul(x.transpose(-2, -1), dV)
return dWq, dWk, dWv
验证实验
我们可以比较手动计算和自动求导的结果差异:
# 创建测试数据
d_model, d_k = 64, 32
x = torch.randn(16, 10, d_model, requires_grad=True)
model = ManualAttention(d_model, d_k)
# 自动求导
O = model(x)
loss = O.sum()
loss.backward()
# 手动求导
O = model(x)
dO = torch.ones_like(O)
dWq_manual, dWk_manual, dWv_manual = model.manual_backward(dO)
# 比较结果
print("dWq diff:", torch.norm(model.W_q.weight.grad - dWq_manual))
print("dWk diff:", torch.norm(model.W_k.weight.grad - dWk_manual))
print("dWv diff:", torch.norm(model.W_v.weight.grad - dWv_manual))
常见问题与优化
数值稳定性问题
- 梯度消失 / 爆炸 :
- 当 d_k 较大时,点积值可能过大导致 Softmax 饱和
-
解决方案:使用缩放因子 1 /√d_k
-
Softmax 数值溢出 :
- 实现时应对输入做最大值归一化
def stable_softmax(x): x = x - x.max(dim=-1, keepdim=True)[0] return torch.exp(x) / torch.exp(x).sum(dim=-1, keepdim=True)
内存优化
- 使用分块计算减少内存占用
- 对于长序列,考虑稀疏 Attention 或局部 Attention
- 混合精度训练可以减少内存使用
总结与延伸
通过本文的推导和实现,我们深入理解了 Attention 机制的反向传播过程。这为我们调试和优化 Attention 层提供了理论基础。
留给读者的思考题:如何将这里的推导扩展到多头 Attention?提示:需要考虑多个头的并行计算和最后的线性变换。
在实际应用中,理解这些底层机制能帮助我们更好地设计模型架构、诊断训练问题,以及实现各种 Attention 变体。建议读者尝试自己推导多头 Attention 的反向传播,并比较不同实现方式的性能差异。
