1. 优化器

优化器是管理反向传播的类,每次反向传播计算出一个梯度,并不是直接加到参数上,而是要考虑乘上一个系数,这个系数具体怎么选择,就是优化器负责的。

几种常见 optimizer 的关系可以粗略理解为:

momentum:SGD 加上 gradient 的指数移动平均。
AdaGrad:用历史 squared gradient 的累计值缩放更新。
RMSProp:对 squared gradient 使用指数移动平均。
Adam:RMSProp 加 momentum。

1.1 momentum

这里有一个比较好的比喻是,参数看成坐标,梯度是参数的变化量,看成速度。对于梯度,并不是直接加到参数上,而是也会乘上系数,可以看成速度也在变化,变化率由这个系数决定,那么这个系数可以看成加速度。

SGD(随机梯度下降)是最基础的,就是直接把梯度加到参数上。这相当于一个没有加速度项,且没有惯性的物体。这可能导致的结果是,开始步长太小,收敛太慢,后期陷入局部最优解出不去。

动量优化的思想是,真实物体有惯性项,有保持历史速度的倾向,或者说有动量,想改变动量需要付出条件。这样可以让参数带着动量冲出局部最优解,并且动量叠加,可以让学习率变大,加速收敛。

具体实现上,速度或者说动量,采用历史梯度的指数移动平均,这样越近的梯度权重越大,越远的梯度影响越小,考虑了历史梯度信息,但不会使得很久以前的梯度一直叠加。

速度 vt=β⋅vt−1+(1−β)⋅gt\text{速度 } v_t = \beta \cdot v_{t-1} + (1-\beta) \cdot g_t速度 vt=βvt1+(1β)gt
参数更新 wt=wt−1−η⋅vt\text{参数更新 } w_t = w_{t-1} - \eta \cdot v_t参数更新 wt=wt1ηvt

其中

  • gtg_tgt 是当前的梯度。
  • vtv_tvt 是积累的“速度”(即一阶矩,通常超参数 β=0.9\beta = 0.9β=0.9)。
  • η\etaη 是学习率。

1.2 AdaGrad(自适应梯度)

动量会记住历史梯度信息,但是不会改学习率,所有参数共享一个学习率η\etaη。但不同参数需要的学习率可能不同,并且不固定,不能开始写死。于是一个想法是给每个参数自适应学习率。根据历史梯度信息调整学习率。

具体公式为
历史梯度平方累加 st=st−1+gt2\text{历史梯度平方累加 } s_t = s_{t-1} + g_t^2历史梯度平方累加 st=st1+gt2
参数更新 wt=wt−1−ηst+ϵ⋅gt\text{参数更新 } w_t = w_{t-1} - \frac{\eta}{\sqrt{s_t} + \epsilon} \cdot g_t参数更新 wt=wt1st+ϵηgt

其中

  • sts_tst 是从第一天到今天所有梯度的平方累加和。
  • ϵ\epsilonϵ 是一个极小的数(比如 10−810^{-8}108),防止分母为 0。

这里缩放学习率,相当于一个加速度项,能调整速度项的大小。放在学习率的分母上,如果历史梯度太大,则起到缩小学习率的作用,这样可以稳定参数更新,避免更新步长太大或太小。

但是注意到实现方法是累加平方和,这会导致历史信息只会越来越大,最后使得学习率趋于0,但此时模型可能还没收敛。

1.3 RMSProp(均方根传播)

这是为了解决自适应梯度,历史梯度平方和越来越大的问题。累加历史梯度平方和时,采用类似动量优化器里的指数移动平均,这样可以遗忘早期信息,避免一直累加越来越大。这样如果一段时间内梯度都很大,会缩小学习率,如果后面梯度又太小了,还能再放大学习率。

具体公式为
平方梯度的均值 st=β⋅st−1+(1−β)⋅gt2\text{平方梯度的均值 } s_t = \beta \cdot s_{t-1} + (1-\beta) \cdot g_t^2平方梯度的均值 st=βst1+(1β)gt2

参数更新 wt=wt−1−ηst+ϵ⋅gt\text{参数更新 } w_t = w_{t-1} - \frac{\eta}{\sqrt{s_t} + \epsilon} \cdot g_t参数更新 wt=wt1st+ϵηgt

β\betaβ 是衰减率(通常设为 0.999),这意味着它更关注最近几步的梯度大小,很久以前的梯度就被“遗忘了”。

1.4 Adam(自适应矩估计)

融合了动量和RMSPRop,用 Momentum 算一阶矩(确定动量)。用 RMSProp 算二阶矩(每个参数自适应学习率)。

具体公式为
mt=β1⋅mt−1+(1−β1)⋅gtm_t = \beta_1 \cdot m_{t-1} + (1-\beta_1) \cdot g_tmt=β1mt1+(1β1)gt
vt=β2⋅vt−1+(1−β2)⋅gt2v_t = \beta_2 \cdot v_{t-1} + (1-\beta_2) \cdot g_t^2vt=β2vt1+(1β2)gt2
m^t=mt1−β1t,v^t=vt1−β2t\hat{m}_t = \frac{m_t}{1-\beta_1^t}, \quad \hat{v}_t = \frac{v_t}{1-\beta_2^t}m^t=1β1tmt,v^t=1β2tvt
wt=wt−1−ηv^t+ϵ⋅m^tw_t = w_{t-1} - \frac{\eta}{\sqrt{\hat{v}_t} + \epsilon} \cdot \hat{m}_twt=wt1v^t+ϵηm^t

中间第三行,是因为初始m,vm,vm,v都是0,会导致开始的更新步长太小,因此做一个缩放,下面β1t\beta_1^tβ1t就是β\betaβttt次方,在开始可以让参数更大,随着训练步数的增加,这一项会趋近于0,缩放分母趋近于1,也就是不产生影响。

2. 内存估计

内存占用大致可以估计为以下几项加起来。

    num_parameters = D * D * L
    parameter_memory = 2 * num_parameters  # (2 bytes for bf16) 
    gradient_memory = 2 * num_parameters  # (2 bytes for bf16) 
    optimizer_state_memory = 4 * num_parameters  # (4 bytes for fp32) 
    activation_memory = 2 * (B * D * L)  # (2 bytes for bf16) 

模型参数和梯度用fp16,每个元素2B,优化器要求高精度,用fp32,每个元素4B。激活值也用fp32。

优化器使用高精度的原因是,反向传播更新时,会用学习率乘上梯度,这两个值都很小,乘起来可能会小于fp16的最小可表示值,变成0,那更新相当于什么都没做。用fp32则可以正确表示这个极小值,正常更新。

其中参数,梯度,优化器的元素个数都是一样的,都等于模型总参数。模型总参数的估计方法为,假设有L层,每层最主要的是全连接层,模型隐藏层维度D,那么每个全连接层参数都是D×DD \times DD×D,故总参数量D×D×LD \times D\times LD×D×L

当然,优化器的元素个数估计为等于参数量是一个保守的估计,前面Adam就保存了一阶二阶矩。严格来说应该是和参数量一次方成正比。

激活值的B是batch size,D是隐藏层维度,也是每个token的维度,L是模型层数,反向传播中,每一层的激活值都要保存,故乘上L。

3. 计算量估计

前一节估计了内存占用,计算量也是另一个重要的估计指标。

对于模型中的某一层,前向传播最主要的计算量都在全连接层,全连接层计算是(B,D)的激活值,乘上(D,D)的全连接层矩阵,矩阵乘法每个位置都需要一次加法,一次乘法,总计算量大概为
2BD2 FLOPs. 2BD^2\ \text{FLOPs}. 2BD2 FLOPs.

backward 需要计算两个 matmul,其中对h1的偏导,是为了反向传播到更前面的层,对W的偏导,是为了计算当前层的梯度,更新当前层。

∂L∂h1=∂L∂h2W2T, \frac{\partial L}{\partial h_1} =\frac{\partial L}{\partial h_2}W_2^T, h1L=h2LW2T,

∂L∂W2=h1T∂L∂h2. \frac{\partial L}{\partial W_2} =h_1^T\frac{\partial L}{\partial h_2}. W2L=h1Th2L.

反向传播整体参数量是前向传播的两倍,计算量
4BD2 FLOPs. 4BD^2\ \text{FLOPs}. 4BD2 FLOPs.
这一层的总计算量
6BD2 FLOPs. 6BD^2\ \text{FLOPs}. 6BD2 FLOPs.
考虑到模型有L层,总计算量
6BD2L FLOPs. 6BD^2L\ \text{FLOPs}. 6BD2L FLOPs.
或者考虑N = num_parameters = D * D * L,也可以写成
6BN FLOPs. 6BN\ \text{FLOPs}. 6BN FLOPs.

4. 计算练习

用这一节学习的估算方法,我们可以估计一些实际训练中的规模问题。

问题一:用 1024 张 H100,在 15T tokens 上训练 70B 模型需要多久?

训练 compute 的常用粗略估计为 6ND6ND6ND,其中 NNN 是参数量,DDD 是训练 token 数:

C=6×70×109×15×1012=6.3×1024 FLOPs. C=6\times70\times10^9\times15\times10^{12} =6.3\times10^{24}\ \text{FLOPs}. C=6×70×109×15×1012=6.3×1024 FLOPs.

若一张 H100 的 dense bf16 峰值约为 989.5989.5989.5 TFLOP/s,MFU 为 0.5,则每天可用计算量为

989.5×1012×0.5×1024×86400, 989.5\times10^{12}\times0.5\times1024\times86400, 989.5×1012×0.5×1024×86400,

最终约需 144 天

问题二:8 张 80 GB H100 使用 AdamW 最多能训练多大的模型?

混合精度训练中,每个参数至少需要:

  • 参数:2 bytes(bf16)。
  • gradient:2 bytes(bf16)。
  • Adam 一阶矩与二阶矩:4+4=84+4=84+4=8 bytes(fp32)。

因此上界为

N=80×109×812≈53.3 B parameters. N=\frac{80\times10^9\times8}{12}\approx53.3\ \text{B parameters}. N=1280×109×853.3 B parameters.

5. 节约显存

5.1 梯度累加 Gradient accumulation

大batch能降低梯度的随机性,提高训练稳定性。但是激活值占用的内存和B成正比,太大会显存放不下。
Mactivation≈2BDL M_{activation}\approx2BDL Mactivation2BDL
一个解决方法是梯度累加Gradient accumulation,把一个batch分成多个小batch,然后累加他们的梯度值,计算多个小batch,才更新一次参数,清空梯度。这样的代价是会增加batch迭代轮数,并且只能影响激活参数占用的内存,其他内存如梯度,优化器都不会收到影响,可优化比例有限。

5.2 激活值检查点 Activation checkpointing

反向传播中需要每一层的激活值,正常做法是全都保存。那么为了节约内存,可以不保存或少保存,现场计算,时间换空间。

具体思想是:

  • forward 只保存部分层的 activation。
  • backward 从最近的 checkpoint 重新计算缺失的 activation

对于 LLL 层网络:

  • 保存每层 activation:memory 为 O(L)O(L)O(L),无需重算。
  • 完全不保存:memory 为 O(1)O(1)O(1),但每层都从头重算,compute 为 O(L2)O(L^2)O(L2),计算思路是:L层每层都重新算,第i层需要前向传播i层,计算量i,总计算量O(∑i=1Li)=O(L2)O(\sum_{i=1}^Li)=O(L^2)O(i=1Li)=O(L2)
  • 每隔 L\sqrt LL 层保存:memory 为 O(L)O(\sqrt L)O(L),额外重算仍为 O(L)O(L)O(L)。计算思路是,从每个检查点开始,计算右侧一个块的激活值并保存,有L/D块,每块长度D,计算量D,总计算量O(∑i=1L/DD)=O(L)O(\sum_{i=1}^{L/D}D)=O(L)O(i=1L/DD)=O(L)

checkpoint 间隔决定了 memory 与 compute 的具体平衡。

更多推荐