大模型中的数学(002):DP 中的梯度缩放

数据并行中的梯度缩放

考虑一个大小为 NN 的全局 batch B={(xi,yi)}i=1N\mathcal{B} = \{(\mathbf{x}_i, \mathbf{y}_i)\}_{i=1}^N。全局目标损失函数采用 mean reduction:

L(W)=1Ni=1Ni(W)L(\mathbf{W}) = \frac{1}{N} \sum_{i=1}^{N} \ell_i(\mathbf{W})

其梯度为基准真值:

g=L(W)=1Ni=1Ni(W)(1)\mathbf{g}^\star = \nabla L(\mathbf{W}) = \frac{1}{N} \sum_{i=1}^{N} \nabla \ell_i(\mathbf{W}) \tag{1}


1. 单卡梯度累加(Gradient Accumulation)

B\mathcal{B} 拆分为 MM 个 micro-batch {Bm}m=1M\{\mathcal{B}_m\}_{m=1}^M,每个大小为 B=N/MB = N/M。对第 mm 个 micro-batch 计算局部损失(mean reduction):

Lm(W)=1BmiBmi(W)=MNiBmi(W)L_m(\mathbf{W}) = \frac{1}{|\mathcal{B}_m|} \sum_{i \in \mathcal{B}_m} \ell_i(\mathbf{W}) = \frac{M}{N} \sum_{i \in \mathcal{B}_m} \ell_i(\mathbf{W})

反向传播得到局部梯度:

gm=Lm(W)=MNiBmi(W)\mathbf{g}_m = \nabla L_m(\mathbf{W}) = \frac{M}{N} \sum_{i \in \mathcal{B}_m} \nabla \ell_i(\mathbf{W})

在同一张卡上累加 MM 个 micro-batch 的梯度:

m=1Mgm=m=1MMNiBmi(W)=MNi=1Ni(W)=Mg(2)\sum_{m=1}^{M} \mathbf{g}_m = \sum_{m=1}^{M} \frac{M}{N} \sum_{i \in \mathcal{B}_m} \nabla \ell_i(\mathbf{W}) = \frac{M}{N} \sum_{i=1}^{N} \nabla \ell_i(\mathbf{W}) = M \cdot \mathbf{g}^\star \tag{2}

结论:累加后的梯度是基准真值的 MM 倍。为保证等价性,必须执行 缩放。有两种等价实现:

  • 方式一(缩放梯度):在优化器更新前将累加梯度除以 MM

    gacc=1Mm=1Mgm=g\mathbf{g}_{\text{acc}} = \frac{1}{M} \sum_{m=1}^{M} \mathbf{g}_m = \mathbf{g}^\star

  • 方式二(缩放损失):将每个 micro-batch 的损失预先缩放

    L~m(W)=1MLm(W)=1NiBmi(W)\tilde{L}_m(\mathbf{W}) = \frac{1}{M} \cdot L_m(\mathbf{W}) = \frac{1}{N} \sum_{i \in \mathcal{B}_m} \ell_i(\mathbf{W})

    此时 g~m=1Mgm\tilde{\mathbf{g}}_m = \frac{1}{M} \mathbf{g}_m,累加后自然得到 m=1Mg~m=g\sum_{m=1}^M \tilde{\mathbf{g}}_m = \mathbf{g}^\star,无需额外缩放梯度。

注:工程上,一般更常用的是缩放损失:在 forward 阶段对 loss 做缩放,backward 自动传播正确梯度,避免在优化器步骤额外处理。


2. 多卡数据并行(Data Parallel)

B\mathcal{B} 拆分为 DD 份(DD 为数据并行维度,即 DP size),卡 kk 持有子集 Bk\mathcal{B}_k,满足 Bk=N/D|\mathcal{B}_k| = N/Dk=1DBk=B\bigcup_{k=1}^D \mathcal{B}_k = \mathcal{B}。卡 kk 的局部损失:

Lk(W)=1BkiBki(W)=DNiBki(W)L_k(\mathbf{W}) = \frac{1}{|\mathcal{B}_k|} \sum_{i \in \mathcal{B}_k} \ell_i(\mathbf{W}) = \frac{D}{N} \sum_{i \in \mathcal{B}_k} \ell_i(\mathbf{W})

局部梯度:

gk=Lk(W)=DNiBki(W)\mathbf{g}_k = \nabla L_k(\mathbf{W}) = \frac{D}{N} \sum_{i \in \mathcal{B}_k} \nabla \ell_i(\mathbf{W})

各卡执行 all_reduce 求和(SUM):

gsum=k=1Dgk=k=1DDNiBki(W)=DNi=1Ni(W)=Dg(3)\mathbf{g}_{\text{sum}} = \sum_{k=1}^{D} \mathbf{g}_k = \sum_{k=1}^{D} \frac{D}{N} \sum_{i \in \mathcal{B}_k} \nabla \ell_i(\mathbf{W}) = \frac{D}{N} \sum_{i=1}^{N} \nabla \ell_i(\mathbf{W}) = D \cdot \mathbf{g}^\star \tag{3}

结论:all-reduce SUM 后的梯度是基准真值的 DD 倍。为保证等价性,必须执行 平均

gdp=1Dk=1Dgk=g\mathbf{g}_{\text{dp}} = \frac{1}{D} \sum_{k=1}^{D} \mathbf{g}_k = \mathbf{g}^\star


3. 梯度累加与数据并行联合使用

当同时使用梯度累加(MM 个 micro-batch)和数据并行(DD 张卡)时,每张卡上的 N/DN/D 个样本再细分为 MM 个 micro-batch,每个 micro-batch 大小为 N/(DM)N/(DM)。卡 kk 的第 mm 个局部损失:

Lk,m(W)=1Bk,miBk,mi(W)=DMNiBk,mi(W)L_{k,m}(\mathbf{W}) = \frac{1}{|\mathcal{B}_{k,m}|} \sum_{i \in \mathcal{B}_{k,m}} \ell_i(\mathbf{W}) = \frac{DM}{N} \sum_{i \in \mathcal{B}_{k,m}} \ell_i(\mathbf{W})

对应局部梯度:

gk,m=Lk,m(W)=DMNiBk,mi(W)\mathbf{g}_{k,m} = \nabla L_{k,m}(\mathbf{W}) = \frac{DM}{N} \sum_{i \in \mathcal{B}_{k,m}} \nabla \ell_i(\mathbf{W})

先进行单卡上的梯度累加,再进行卡间的 all-reduce SUM:

k=1Dm=1Mgk,m=DMNk=1Dm=1MiBk,mi(W)=DMNi=1Ni(W)=DMg\sum_{k=1}^{D} \sum_{m=1}^{M} \mathbf{g}_{k,m} = \frac{DM}{N} \sum_{k=1}^{D} \sum_{m=1}^{M} \sum_{i \in \mathcal{B}_{k,m}} \nabla \ell_i(\mathbf{W}) = \frac{DM}{N} \sum_{i=1}^{N} \nabla \ell_i(\mathbf{W}) = DM \cdot \mathbf{g}^\star

结论:联合累加后的梯度是基准真值的 MDM \cdot D 倍。最终缩放应为:

gtotal=1MDk=1Dm=1Mgk,m=g\mathbf{g}_{\text{total}} = \frac{1}{M \cdot D} \sum_{k=1}^{D} \sum_{m=1}^{M} \mathbf{g}_{k,m} = \mathbf{g}^\star

评论