数据并行中的梯度缩放
考虑一个大小为 N 的全局 batch B={(xi,yi)}i=1N。全局目标损失函数采用 mean reduction:
L(W)=N1i=1∑Nℓi(W)
其梯度为基准真值:
g⋆=∇L(W)=N1i=1∑N∇ℓi(W)(1)
1. 单卡梯度累加(Gradient Accumulation)
将 B 拆分为 M 个 micro-batch {Bm}m=1M,每个大小为 B=N/M。对第 m 个 micro-batch 计算局部损失(mean reduction):
Lm(W)=∣Bm∣1i∈Bm∑ℓi(W)=NMi∈Bm∑ℓi(W)
反向传播得到局部梯度:
gm=∇Lm(W)=NMi∈Bm∑∇ℓi(W)
在同一张卡上累加 M 个 micro-batch 的梯度:
m=1∑Mgm=m=1∑MNMi∈Bm∑∇ℓi(W)=NMi=1∑N∇ℓi(W)=M⋅g⋆(2)
结论:累加后的梯度是基准真值的 M 倍。为保证等价性,必须执行 缩放。有两种等价实现:
-
方式一(缩放梯度):在优化器更新前将累加梯度除以 M
gacc=M1m=1∑Mgm=g⋆
-
方式二(缩放损失):将每个 micro-batch 的损失预先缩放
L~m(W)=M1⋅Lm(W)=N1i∈Bm∑ℓi(W)
此时 g~m=M1gm,累加后自然得到 ∑m=1Mg~m=g⋆,无需额外缩放梯度。
注:工程上,一般更常用的是缩放损失:在 forward 阶段对 loss 做缩放,backward 自动传播正确梯度,避免在优化器步骤额外处理。
2. 多卡数据并行(Data Parallel)
将 B 拆分为 D 份(D 为数据并行维度,即 DP size),卡 k 持有子集 Bk,满足 ∣Bk∣=N/D 且 ⋃k=1DBk=B。卡 k 的局部损失:
Lk(W)=∣Bk∣1i∈Bk∑ℓi(W)=NDi∈Bk∑ℓi(W)
局部梯度:
gk=∇Lk(W)=NDi∈Bk∑∇ℓi(W)
各卡执行 all_reduce 求和(SUM):
gsum=k=1∑Dgk=k=1∑DNDi∈Bk∑∇ℓi(W)=NDi=1∑N∇ℓi(W)=D⋅g⋆(3)
结论:all-reduce SUM 后的梯度是基准真值的 D 倍。为保证等价性,必须执行 平均:
gdp=D1k=1∑Dgk=g⋆
3. 梯度累加与数据并行联合使用
当同时使用梯度累加(M 个 micro-batch)和数据并行(D 张卡)时,每张卡上的 N/D 个样本再细分为 M 个 micro-batch,每个 micro-batch 大小为 N/(DM)。卡 k 的第 m 个局部损失:
Lk,m(W)=∣Bk,m∣1i∈Bk,m∑ℓi(W)=NDMi∈Bk,m∑ℓi(W)
对应局部梯度:
gk,m=∇Lk,m(W)=NDMi∈Bk,m∑∇ℓi(W)
先进行单卡上的梯度累加,再进行卡间的 all-reduce SUM:
k=1∑Dm=1∑Mgk,m=NDMk=1∑Dm=1∑Mi∈Bk,m∑∇ℓi(W)=NDMi=1∑N∇ℓi(W)=DM⋅g⋆
结论:联合累加后的梯度是基准真值的 M⋅D 倍。最终缩放应为:
gtotal=M⋅D1k=1∑Dm=1∑Mgk,m=g⋆
评论