混合精度训练、显存优化与大模型训练策略深度解析¶
混合精度训练(AMP)的基本原理是什么?FP16 和 FP32 各用于什么部分?¶
混合精度训练(Automatic Mixed Precision, AMP)的核心思想是在训练过程中同时使用半精度(FP16)和单精度(FP32)浮点数,以在保持模型精度的前提下大幅减少显存占用和加速计算。现代 GPU(如 NVIDIA Volta、Turing、Ampere、Hopper 架构)配备了专门的 Tensor Core,可以在 FP16 下进行极高效的矩阵乘加运算,吞吐量是 FP32 的数倍。
AMP 的运作方式是:将大多数前向和反向传播中的激活值、权重和梯度存储在 FP16 中,利用 Tensor Core 加速计算;同时,保持一份 FP32 的主副本权重(master weights)用于参数更新,以避免 FP16 的有限数值范围造成的精度损失。具体来说:
-
FP16 用于:前向传播的矩阵乘法(QKV 投影、FFN 线性层、注意力分数计算等)、反向传播的梯度计算。这些操作在 FP16 下通过 Tensor Core 可以获得数倍的吞吐量提升,且显存占用减半。
-
FP32 用于:权重的主副本(master weights),优化器状态(Adam 的动量 m 和方差 v),以及需要高精度的归约操作(如 BatchNorm 的统计量,但 Transformer 多用 LayerNorm,此处影响较小)。主副本在 FP32 下更新后,再转换为 FP16 用于下一轮前向传播。
此外,AMP 会动态地将一些对数值范围敏感的操作(如 softmax、LayerNorm、损失计算)提升到 FP32,以避免溢出。这些操作虽然计算量不大,但对精度至关重要。通过这种方式,AMP 在几乎不损失模型最终精度的情况下,实现了接近 2 倍的训练加速和约 25-40% 的显存节省。
为什么需要 Loss Scaling?在 FP16 训练中,动态损失缩放是如何工作的?¶
Loss Scaling 是 FP16 混合精度训练中必不可少的技术,旨在解决梯度下溢问题。FP16 的最小正规数约为 6×10⁻⁸,许多梯度值(尤其是训练初期或网络中较远的层)可能远小于这一阈值,导致它们在 FP16 下被截断为零。如果大量梯度变为零,训练将停滞。
Loss Scaling 的核心做法是:在反向传播前,将损失值乘以一个较大的缩放因子(scale factor),然后进行反向计算。根据链式法则,损失被放大后,所有梯度也会同比例放大,从而使那些原本在 FP16 下会下溢的小梯度值被推入 FP16 的可表示范围。在更新参数前,再将梯度除以相同的缩放因子,恢复到正确的数值范围。
动态损失缩放则进一步自动化了这一过程。其算法如下:
-
初始化一个缩放因子(例如 2¹⁶)。
-
每个训练迭代,将损失乘以缩放因子,进行反向传播,得到缩放后的 FP16 梯度。
-
检查缩放后的梯度是否包含 Inf 或 NaN(表示溢出)。如果没有溢出,则将这些梯度除以缩放因子还原,并更新 FP32 的主副本参数。同时,算法可以尝试逐步增大缩放因子(如乘以 1.05),以利用更多的 FP16 动态范围。
-
如果检测到溢出(Inf/NaN),则跳过本次参数更新,将缩放因子减小(如乘以 0.5),并重试该次迭代。
动态损失缩放自动平衡了溢出的风险和梯度的精度,使得 FP16 训练在大多数情况下可以平稳进行。
BF16 相比 FP16 有什么优势?为什么大模型训练更倾向于使用 BF16?¶
BF16(Brain Floating Point 16)和 FP16 都是 16 位浮点数,但它们的位分配不同:FP16 有 1 位符号、5 位指数、10 位尾数;BF16 有 1 位符号、8 位指数、7 位尾数。BF16 的指数位数与 FP32 相同(都是 8 位),因此它的动态范围与 FP32 完全一致,可以表示的数值范围约为 10⁻³⁸ 到 10³⁸,远大于 FP16 的 6×10⁻⁸ 到 65504。
这一特性带来了两大关键优势:
-
无需损失缩放:由于 BF16 的动态范围足够大,训练中很少出现梯度上溢(Inf)或下溢(零)的问题,因此不需要复杂的损失缩放机制,简化了训练流程。
-
与 FP32 的互换性:BF16 和 FP32 之间的转换非常简单——直接截断或填充尾数即可,不需要处理指数偏移。这使得硬件实现更高效。
大模型训练更倾向于 BF16 的原因:
-
数值稳定性:大模型参数多、层数深,梯度容易在传播中产生极端值。BF16 的大动态范围能更好地容忍这些极端值,避免训练中断。
-
简化工程:省去损失缩放,减少了超参数调优和溢出处理的开销。
-
硬件支持:NVIDIA A100、H100 等 GPU 以及 Google TPU 都对 BF16 提供原生支持,其 Tensor Core 在 BF16 下的吞吐量与 FP16 相当甚至更高。
因此,BF16 已成为大模型训练的主流半精度格式。
训练一个 10B 参数的模型,粗略估计需要多少显存?列出参数、梯度、优化器状态、激活值的组成。¶
假设使用混合精度训练(FP16 权重和梯度,FP32 主副本和优化器状态),模型参数量为 10B。
-
模型参数(权重):FP16 存储,每个参数 2 字节,共 10B × 2B = 20 GB。
-
梯度:FP16 存储(实际在反向传播中计算,通常与参数同精度),20 GB。
-
优化器状态:若使用 Adam/AdamW,需要存储 FP32 的动量(m)和方差(v),每个参数 4 字节 × 2 = 8 字节。共 10B × 8B = 80 GB。
-
主副本参数(master weights):FP32,用于累积更新,10B × 4B = 40 GB。
以上三项(模型、梯度、优化器状态)通常需要 20 + 20 + 80 + 40 = 160 GB(注意梯度在反向时存在,优化器状态在更新时存在,可部分复用但峰值需考虑)。此外,还有一项大开销:
- 激活值(activations):前向传播中各层的中间输出,用于反向传播计算梯度。激活值的显存与模型结构、序列长度、批次大小有关。对于 Transformer,激活值主要来自注意力矩阵和 FFN 的中间表示。粗略估计,对于一个 10B 的 Decoder-only 模型,若序列长度 2048、批次大小 1,激活值可能占用约 20-40 GB(视具体实现和重计算策略而定)。
因此,总训练显存峰值可能在 160 GB + 20~40 GB ≈ 180-200 GB。这通常需要多卡分布式训练(如 4 张 A100 80GB)。
梯度累积(Gradient Accumulation)是如何模拟大 batch size 训练的?它会影响 Batch Normalization 吗?¶
梯度累积是一种用时间换空间的技术。当单卡显存无法容纳所需的大批量数据时,可以将一个大批次拆分成多个微批次(micro-batches),依次计算每个微批次的损失和梯度,但不立即更新参数,而是将梯度累加。经过 K 次累加后,模型参数使用累积的总梯度更新,等效于使用 batch size = K × micro_batch_size 进行训练。
具体步骤:
-
初始化梯度累加器为零。
-
对每个微批次:前向计算损失 → 反向传播得到梯度(注意梯度会除以总批大小以保持尺度)→ 将梯度累加到累加器。
-
重复 K 次后,执行优化器步骤更新参数,然后清零累加器。
对 Batch Normalization 的影响:BN 在计算均值和方差时,依赖于批次内的统计量。如果只在每个微批次内计算 BN,这些统计量基于很小的样本数,噪声大,会损害训练。因此,通常有两种应对方式:
-
使用 SyncBN(同步批归一化):在微批次间同步均值和方差,但实现复杂,不常见。
-
改用 Layer Normalization:Transformer 中已普遍使用 LayerNorm,它对每个样本独立归一化,不依赖批次统计,因此梯度累积不会影响它。这也是 Transformer 不受此问题困扰的原因。
-
在需要 BN 的场景(如某些视觉模型),实现正确的梯度累积需要确保 BN 的统计量在等效大批次上计算,通常需要框架支持或通过设置
momentum累积。
梯度检查点(Gradient Checkpointing)的原理是什么?它用计算换空间的具体代价如何?¶
梯度检查点是一种以计算换内存的技术。在标准反向传播中,需要保存前向传播的所有中间激活值,以便在反向时计算梯度。对于深层网络,这些激活值的显存占用可能极大(尤其是注意力矩阵)。梯度检查点会选择性丢弃某些中间激活值,在反向传播需要时,再通过从最近的检查点重新进行前向计算来恢复这些激活值。
具体实现:将网络划分为若干段(checkpoint segments)。前向传播时,只保留每个段的输入激活(检查点),段内部的中间激活全部丢弃。反向传播到某段时,从该段的输入激活出发,重新执行该段的前向计算,得到所需的中间激活,进而计算该段的梯度。这样,内存占用从存储所有中间激活降低为仅存储检查点激活,但代价是需要额外的重计算。
代价:重计算引入约 33% 的额外前向计算量(对于 Transformer,通常对每个 Transformer 层设置检查点,重计算一整层)。对于大模型,这部分的计算开销相对于整体训练是值得的,因为它可以释放大量显存,从而允许更大的批次或更长的序列。
在 Transformer 中,梯度检查点通常作用在哪些子层上?为什么选择这些位置?¶
梯度检查点通常以整个 Transformer 层为单元施加。原因:
-
每层内部包含多头注意力和 FFN,中间会产生大量激活(尤其是注意力概率矩阵和 FFN 隐层输出),占用内存高。
-
以整层为检查点,只需保存该层的输入(即上一层的输出),重计算时重新执行该层的所有运算。这样实现简单,且能有效减少内存。
-
如果粒度更细(如只对注意力或 FFN 做检查点),虽然可能进一步节省一些内存,但会引入更多重计算和实现复杂度。因此,实践中通常在每个 Transformer 层边界设置检查点。
分析一下不使用任何显存优化技术时,一个 7B 模型在 FP16 下的理论训练显存需求。¶
模型 7B 参数,假设优化器为 Adam,无梯度检查点、无 CPU Offload、无模型并行。FP16 混合精度训练(FP16 参数和梯度,FP32 优化器状态和主副本):
-
参数(FP16):7B × 2 = 14 GB
-
梯度(FP16):14 GB(在反向时存在,可与参数同时占用)
-
优化器状态(FP32 momentum + variance):7B × 4 × 2 = 56 GB
-
主副本参数(FP32):7B × 4 = 28 GB
-
激活值:取决于序列长度、批次大小。以 2048 tokens,batch size=1 为例,每层激活约需 50-100 MB,若 32 层,则约 1.6-3.2 GB;注意力矩阵会额外占用较大内存。总激活可能在 5-15 GB。
总计峰值约 14+14+56+28+10 ≈ 122 GB(若梯度与参数不同时,可略少,但优化器和主副本必须常驻)。这显存需求远超单张 A100 80GB,必须使用张量并行或 ZeRO 等分布式技术。
如何通过调整模型架构(如减少层数、维度)来降低显存?在给定显存下如何最大化模型效果?¶
降低显存的架构调整:
-
减少层数:直接降低参数总量,参数、梯度、优化器状态等比例下降。但模型深度对性能影响大,浅而宽的模型可能不如深而窄的模型。
-
减小模型宽度(隐藏维度 d_model):参数随维度平方增长,效果显著。可结合扩中间层维度(FFN)的比例调整。
-
降低注意力头数:在总维度固定下,头数影响多头注意力的参数量和计算,适当减少头数(保持每个头维度合理)可节省参数。
-
使用 GQA 或 MQA:分组查询注意力或 Multi-Query Attention 可以减少 KV 缓存和参数,尤其对推理显存有效,训练时也能减少部分参数。
-
词表压缩:词表嵌入矩阵参数量大,可如 ALBERT 进行矩阵分解。
在给定显存下最大化模型效果的原则:
-
深度优先:在预算内尽可能增加层数,因为深度对模型抽象能力更重要。可适当压缩宽度和 FFN 比例。
-
使用 MoE(混合专家):用稀疏激活的 FFN 增大总参数量而不显著增加计算和显存(专家分布在多卡)。在相同显存下,MoE 可以显著提高模型容量和性能。
-
调整精度和优化器:FP16/BF16 混合精度;使用更省显存的优化器(如 Adafactor,减少动量存储)。
-
结合并行和显存优化:ZeRO 和梯度检查点等可以释放显存,从而允许更大的模型。
ZeRO 和梯度检查点可以同时使用吗?它们优化的是显存中的哪些部分?¶
可以同时使用,它们是互补的显存优化技术。
-
ZeRO(Zero Redundancy Optimizer):主要优化优化器状态、梯度和模型参数的冗余存储。在数据并行中,每张卡都保存完整的上述三者,ZeRO 通过分片(partitioning)将优化器状态(ZeRO-1)、梯度(ZeRO-2)、模型参数(ZeRO-3)分布到不同 GPU,每张卡只存储一部分,需要时通过通信获取。这大幅降低了单卡显存占用。
-
梯度检查点:主要优化激活值的显存。通过重计算来减少需要存储的中间激活。
两者作用的对象不同:ZeRO 针对参数、梯度、优化器状态;检查点针对激活值。同时使用可以成倍地降低总显存需求,使训练超大模型成为可能。例如,训练千亿模型时,ZeRO-3 + 梯度检查点是标准配置。
在训练大模型时,如何平衡微批大小、梯度累积步数和总批大小之间的关系?¶
-
总批大小(Global Batch Size, GBS):由训练任务和模型决定,通常通过实验或遵循扩展定律(如 LLM 的预训练 GBS 常在 1M-4M tokens)。太小的 GBS 训练不稳定,收敛慢;太大的 GBS 会导致每个 batch 的样本数过多,可能需增大学习率,但边际收益递减且硬件利用率下降。
-
微批大小(Micro Batch Size, MBS):每张 GPU 一次前向能处理的最大样本数(受显存限制)。MBS 越大,计算效率越高,但显存占用也越大。
-
梯度累积步数(GAS):为了达到 GBS,需要累积的微批次步数。GBS = MBS × GPU 数量 × GAS。
平衡策略:
-
首先根据单卡显存和模型,确定最大的 MBS(通常 1 或 2,若显存允许可更大)。
-
根据总卡数和期望的 GBS,计算需要的 GAS。GAS 过大(如几百)会导致训练速度线性下降,且由于每次更新间的延迟增加,可能影响收敛(但对 Transformer 影响较小)。在可行的情况下,应优先增加卡数减少 GAS。
-
如果 GAS 必须较大,可适当提高学习率(按 sqrt 或线性缩放)来补偿更新频率的降低,但需谨慎。
-
在无法增加卡数时,也可以考虑使用累积梯度但不立即更新,而是利用分布式优化器(如 ZeRO-2/3)在卡间分担状态,并结合通信优化。
显存不足时,CPU Offload 可以将哪些数据卸载到内存?性能下降的幅度通常有多大?¶
CPU Offload 可以将以下数据从 GPU 显存卸载到 CPU 内存:
-
优化器状态(动量、方差)——这是最常用的卸载,因为它们体积大(参数量 × 8 字节),且在前向/反向中不需要。
-
模型参数(如 ZeRO-Offload 将 FP16 参数也卸载,前向时按需从 CPU 搬到 GPU)。
-
激活值(部分框架支持,但较少见,因需要频繁传输)。
当仅卸载优化器状态时,性能下降通常20%-50%,取决于 PCIe 带宽和计算/通信重叠的效率。因为每一步更新时,需要从 CPU 内存读取 FP32 状态到 GPU,更新后再写回。如果利用 CUDA Stream 实现异步传输,并与计算重叠,可部分隐藏延迟,损失可控制在 30% 左右。
若卸载参数(如 ZeRO-Infinity),由于前向每层都要从 CPU 加载参数,通信量更大,性能下降可能达到 50%-80%,主要限于极大规模模型且 GPU 显存极端不足的场景。
使用 BF16 时,为什么不需要像 FP16 那样做 Loss Scaling?请从表示范围解释。¶
FP16 需要 Loss Scaling 的核心原因是其有限的动态范围:FP16 能表示的最小正规数约为 6×10⁻⁸,许多小梯度值会因小于该值而变成零(下溢)。Loss Scaling 通过放大损失来放大梯度,防止下溢。然而,BF16 的指数位宽与 FP32 相同(8 位),其能表示的最小正规数约为 1.2×10⁻³⁸,与 FP32 相当,远小于实际训练中梯度值的量级。因此,梯度几乎不可能下溢。同时,BF16 的上溢范围也极大(3.4×10³⁸),一般梯度也不会溢出。所以 BF16 在训练中自然避开了 FP16 的数值问题,无需专门的损失缩放机制。
在启用梯度检查点时,哪些中间结果必须保留,哪些可以丢弃再重计算?如何实现最小重计算策略?¶
-
必须保留:每个检查点段的输入激活(即该 Transformer 层的输入张量)。这是重计算该段的起点。此外,如果使用 Pre-LN 架构,LayerNorm 的输入通常也是检查点,因为 LayerNorm 的输出是由输入计算出的,可以重算。
-
可以丢弃再重计算:段内所有的中间激活,包括注意力计算中的 Q, K, V 映射结果、注意力概率矩阵、softmax 前的 logits、FFN 的隐层输出等。
最小重计算策略:以每个 Transformer 层为一个检查点段。前向时只保存该层的输入张量(即上一层的输出 + 残差连接)。反向传播到该层时,从该输入张量开始,重新进行该层的前向计算(包括 LayerNorm、自注意力、FFN、残差相加),得到所有需要的中间激活,进而计算梯度。这种整层重计算方式内存节省最大,重计算量约等于额外一次前向。如果希望减少重计算开销,可以在层内更细粒度设置检查点,但实现复杂且内存节省不多,实践中较少采用。
如果使用 ZeRO-3,模型参数被分片,前向和反向传播时如何获取完整参数?¶
在 ZeRO-3 中,每张 GPU 只持有一层或一部分参数的切片。前向传播时,当需要计算某一层的输出时,该层的所有参数切片会通过 AllGather 通信操作从各 GPU 聚集到每张卡的临时缓冲区,形成完整参数,完成该层计算后,临时缓冲区被释放,只保留自己负责的切片。反向传播同理,需要完整参数来计算梯度,因此再次进行 AllGather。计算完当前层的梯度后,通过 ReduceScatter 将梯度分片,每张卡只保留自己负责的那部分梯度的聚合结果,用于更新自己持有的参数分片。这种“ gather → compute → discard → scatter”的模式,确保了任何时刻每张卡上只有当前层的完整参数,大大节省了显存。
混合精度训练中,权重的主副本(master weights)保存在 FP32,这能防止什么?有没有模型直接用 FP16 做主副本?¶
主副本保存在 FP32 主要防止舍入误差的累积。在 FP16 中,尾数只有 10 位,当用很小的梯度更新权重时(尤其是训练后期或使用较大权重衰减),更新量可能小于 FP16 的最小可表示精度,导致权重无法被更新(称为“更新丢失”)。FP32 主副本累加这些微小更新,保留足够的精度,然后下一轮再将 FP32 主副本转换为 FP16 用于计算。这保证了模型学习到细微的变化。
是否存在直接用 FP16 做主副本的模型?非常少见且不推荐。理论上,如果使用 BF16 做主副本,由于其 7 位尾数精度略低,也可以直接用于更新,配合动态损失缩放或适当的训练策略有时可行,但 FP32 主副本仍是标准。
分析一下在 Transformer 中,不同子层(注意力,FFN,归一化)对显存的占用占比。¶
以 Pre-Norm Transformer 为例,在训练时主要显存占用来自激活值(中间结果)。
-
注意力子层:通常占用激活值的大部分。因为需要存储 Q、K、V 映射后的张量(每个形状 [batch, seq_len, hidden_dim]),以及计算出的注意力概率矩阵([batch, heads, seq_len, seq_len])。当序列长度 L 较大时,注意力矩阵是 L×L 的,显存增长剧烈(O(L²)),成为瓶颈。
-
FFN 子层:需要存储第一个线性层后的激活值(维度通常为 hidden_dim × expansion_ratio,如 4 倍),形状 [batch, seq_len, 4×hidden_dim]。该张量也很大,但通常小于注意力矩阵在长序列下的开销。
-
LayerNorm:参数量和激活值均极小,占用可忽略。
在短序列下,FFN 激活值可能占主导;在长序列(如 >2K)下,注意力矩阵成为绝对主导。这也是为什么长序列训练需要稀疏注意力或 FlashAttention 来减少注意力矩阵的显存。
如果使用 CPU Offload 将优化器状态放到内存,训练速度会下降多少?如何通过重叠计算和传输来隐藏延迟?¶
典型下降 20%-50%,取决于 PCIe 带宽和模型大小。重计算可以通过设计高效的流水线来隐藏延迟:
-
异步传输:使用 CUDA Stream 将优化器状态的读写操作与 GPU 上的前向/反向计算重叠。例如,在当前层计算时,预取下一层所需的优化器状态到 GPU,同时将当前层更新后的状态异步写回 CPU。
-
合并小数据块:优化器状态按层组织,每次传输整层状态而非逐参数,减少传输次数。
-
利用 DeepSpeed ZeRO-Offload 的优化策略:它精心安排了数据在 CPU 和 GPU 之间的驻留位置,将计算密集的操作(如梯度更新)放在 CPU 上执行(因为优化器更新计算量小),而 GPU 专注于大矩阵乘法,并通过多级延迟隐藏技术将通信开销降到最低。
激活值重计算(Activation Recomputation)与梯度检查点有什么区别?是同一件事吗?¶
是同一件事的不同称呼。两者都指在反向传播时,不保存全部中间激活,而是通过从检查点重新前向计算来获取所需的激活值。在 PyTorch 中,通常使用 torch.utils.checkpoint 来实现。有时“激活值重计算”更特指只重计算激活值而保留其他(如权重),与选择性检查点(selective checkpointing)同义。在本语境下,二者等同。
在给定总批大小和硬件的情况下,如何最大化模型规模?哪些并行和显存优化是必选项?¶
为了在有限硬件(例如 8 张 A100 80GB)上训练尽可能大的模型,需综合应用:
必选项:
-
梯度检查点:大幅减少激活内存,几乎零额外通信,是首要手段。
-
混合精度训练(BF16/FP16):参数/梯度内存减半,加速计算。
-
ZeRO 优化:
- ZeRO-1(分片优化器状态)可节省大量显存,配合上述两项已能训练 7B 以上模型。
- ZeRO-2(分片梯度)进一步节省。
- ZeRO-3(分片模型参数)几乎将显存占用均分到所有卡,允许训练接近总显存之和的模型。
根据模型规模可能需要的并行:
-
张量并行(Tensor Parallelism):当单层参数太大,即使 ZeRO-3 分片也无法装入单卡时,使用张量并行将层内矩阵切分到多卡。配合 ZeRO-3 使用。
-
流水线并行(Pipeline Parallelism):将不同层分配到不同设备,结合微批次流水线减少气泡。可与 ZeRO 和张量并行组成 3D 并行。
最大化规模的策略:
-
启用 BF16 + 梯度检查点 + ZeRO-3,这是最高效的显存压缩组合。
-
如果 ZeRO-3 仍不足,加入张量并行(如 Megatron-LM 模式)。
-
在分布式集群中,优先利用数据并行维度上的 ZeRO,因为它通信量相对较小。
-
对于超大模型(如 100B+),采用 3D 混合并行:张量并行 + 流水线并行 + 数据并行(含 ZeRO)。
通过上述组合,可以在有限的硬件上训练远超单卡容量的模型。例如,使用 8 张 A100 80GB 即可训练 LLaMA-7B,甚至 13B 或 70B 若结合适当的并行和显存优化。