Skip to content

训练优化技术

训练大语言模型需要巨大的计算资源,各种优化技术可以显著降低训练成本。

2.3.1 混合精度训练(Mixed Precision Training)

混合精度训练使用FP16/BF16进行前向/反向传播,,在保持精度的同时加速训练。FP16 vs BF16

特性FP16BF16
指数位5 bit8 bit
尾数位10 bit7 bit
数值范围±65,504±3.4×10³⁸
精度较高较低
适用场景推理训练

精度较高较低

适用场景推理训练

latex
FP16: [sign:1][exponent:5][mantissa:10]  范围小但精度高 →
 BF16: [sign:1][exponent:8][mantissa:7]    范围大但精度低 →
 FP32: [sign:1][exponent:8][mantissa:23] → 完整精度

混合精度训练流程

python
from torch.cuda.amp import autocast, GradScaler

 scaler = GradScaler()

 for batch in dataloader:
     optimizer.zero_grad()
     
    #   前向传播使用FP16
     with autocast():
          outputs = model(batch)
          loss = criterion(outputs, targets)


     #   缩放损失并反向传播
     scaler.scale(loss).backward()


     #   梯度裁剪(在缩放空间进行)
     scaler.unscale_(optimizer)
     torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)


     #   更新参数(使用FP32)
     scaler.step(optimizer)
     scaler.update()

性能提升

  • 内存节省:约50%(FP16 vs FP32)
  • 速度提升:1.5-3×(取决于GPU架构,Tensor Core加速)
  • A100/H100:BF16是推荐格式,硬件原生支持

2.3.2 梯度累积(Gradient Accumulation)

原理

将大batch拆分为多个小batch,累积梯度后再更新:

python
accumulation_steps = 4    #    累积4个小batch
effective_batch_size = batch_size * accumulation_steps


  optimizer.zero_grad()
  for i, batch in enumerate(dataloader):
        loss = model(batch) / accumulation_steps    #   缩放损失
        loss.backward()


        if (i + 1) % accumulation_steps == 0:
            optimizer.step() # 根据当前累积的梯度更新模型参数。
            optimizer.zero_grad() # 将模型参数的梯度清零,为下一组 accumulation_steps个 batch 的累积做准备。

内存vs速度的权衡

配置内存使用训练速度适用场景
无累积基准最快内存充足
4步累积25%75%内存受限
8步累积12.5%50%严重受限

2.3.3 ZeRO 优化器(Zero Redundancy Optimizer)

ZeRO是DeepSpeed 提出的数据并行优化技术,将优化器状态、梯度和参数分片到多个GPU

微软 DeepSpeed 团队提出(论文:Rajbhandari et al., SC'20, arXiv:1910.02054),是打破大模型训练"显存墙"最核心的技术之一。


传统 Data Parallel (DDP) 里,每张 GPU 都存一份完整的:

模型状态每张卡都有一份7B 模型/BF16/Adam 举例
参数 W✅ 完整副本~14 GB
梯度 ∇W✅ 完整副本~14 GB
优化器状态(Adam 的 m, v + FP32 master weights)✅ 完整副本~56 GB

三者加起来一张卡就要 ~84 GB+,哪怕模型本身才 14 GB——绝大多数显存是被"冗余副本"吃掉的

ZeRO 的核心 insight 就一句话:

数据并行需要的是"计算结果一致",不是每张卡非得存完整副本。把这些状态切分到 N 张卡上,谁需要谁临时收集,用通信换显存。


ZeRO的三个阶段

设你有 N 张 GPU(数据并行度 = N):

Stage 1 —— 分片优化器状态(OS, Optimizer States)

  • 把 Adam 的 m / v / FP32 master weights 等优化器状态 均分成 N 份
  • 每张 GPU 只存 1/N 的优化器状态
  • 参数 W 和梯度 ∇W 仍然全量存在每张卡上(前反向不变)
  • 参数更新时: 收集需要的状态 → 更新 → 丢弃别人的部分

收益:优化器状态那一大坨从 "Ψ×2×4B" 降到 "Ψ×2×4B / N"

通信开销:很低(每个 step 只在 optimizer.step() 附近触发)

适合:模型刚有点撑不住、但又不想改太多东西时,优先试这个。


Stage 2 —— 再分片梯度(OS + Gradients)

  • 在 Stage 1 基础上,反向传播后对梯度做 Reduce-Scatter(而非 All-Reduce)
  • 每张卡 只保留与自己负责的参数分区对应的梯度(1/N)
  • 参数 W 仍然全量副本在各卡

收益:梯度那份显存也从 Ψ 降到 Ψ/N

通信量:和 DDP 差不多持平(All-Reduce → Reduce-Scatter,实际上是更高效的形式)

适合:10B 量级、显存紧张但还没到"参数塞不进"的程度。


Stage 3 —— 连参数也分片(OS + G + Parameters)⭐

  • 参数本身也被切分:每张 GPU 只常住 1/N 的 W
  • 前向/反向需要某层参数时:All-Gather 临时拼出完整层参数 → 算完立刻释放
  • 这就是所谓的 Fully Sharded Data Parallel 思路(PyTorch FSDP 就是同一思想的原生实现)

理论单卡模型状态显存 ≈ 总模型状态 / N

通信开销:最高(参数收集频繁,但 DeepSpeed 做了 prefetch / overlap_comm 来掩盖)

适合:≥7B~13B+ 或 GPU 数较多时,甚至单卡装不下完整参数。


每张卡存 W?每张卡存 ∇W?每张卡存 OptState?显存节省通信开销
DDP(基准)完整完整完整基线
ZeRO-1完整完整1/N~4x→~3x很低
ZeRO-2完整1/N1/N~3x→~2x中(≈DDP)
ZeRO-31/N1/N1/N~N×较高

内存节省效果

对于参数量为 \(\Psi\) 的模型:

ZeRO Stage每GPU内存相对DDP节省
DDP\(16\Psi\)
ZeRO-1\(16\Psi\)
ZeRO-2\(2\Psi + 14\Psi/N\)~8× (N=8)
ZeRO-3\(16\Psi/N\)

其中 \(N\) 是GPU数量。

python
# DeepSpeed ZeRO    配置示例
  deepspeed_config = {
        "zero_optimization": {
             "stage": 2,   #   或 1, 3
             "offload_optimizer": {
                  "device": "cpu",
                  "pin_memory": True
             },
             "allgather_partitions": True,
             "allgather_bucket_size": 2e8,
             "overlap_comm": True,
             "reduce_scatter": True,
        },
        "train_batch_size": "auto",
        "train_micro_batch_size_per_gpu": "auto",
        "gradient_accumulation_steps": "auto",
        "fp16": {
             "enabled": True
        }
  }

2.3.4 梯度检查点(Gradient Checkpointing)

梯度检查点通过牺牲计算来换取内存:

原理

内存 VS 计算权衡

检查点策略内存节省额外计算适用场景
无检查点0%0%内存充足
每层检查点~50%~20%标准选择
每2层检查点~25%~10%平衡方案
python
from torch.utils.checkpoint import checkpoint

  class CheckpointedTransformerLayer(nn.Module):
        def __init__(self, layer):
            super().__init__()
            self.layer = layer

        def forward(self, x):
            #   使用梯度检查点
            return checkpoint(self.layer, x)

2.3.5 训练优化技术总结

技术内存节省速度影响实现复杂度
混合精度~50%+50-200%
梯度累积可调与累积步数成反比
ZeRO-1~0%-5%
ZeRO-2~60-80%-10%
ZeRO-3~90%+-15%
梯度检查点~50%+20%

实际组合示例:

  • 训练7B模型:混合精度 + ZeRO-2 + 梯度检查点
  • 训练70B模型:混合精度 + ZeRO-3 + 梯度检查点 + CPU Offload

用心记录,持续成长