训练优化技术
训练大语言模型需要巨大的计算资源,各种优化技术可以显著降低训练成本。
2.3.1 混合精度训练(Mixed Precision Training)
混合精度训练使用FP16/BF16进行前向/反向传播,,在保持精度的同时加速训练。FP16 vs BF16
| 特性 | FP16 | BF16 |
|---|---|---|
| 指数位 | 5 bit | 8 bit |
| 尾数位 | 10 bit | 7 bit |
| 数值范围 | ±65,504 | ±3.4×10³⁸ |
| 精度 | 较高 | 较低 |
| 适用场景 | 推理 | 训练 |
精度较高较低
适用场景推理训练
FP16: [sign:1][exponent:5][mantissa:10] 范围小但精度高 →
BF16: [sign:1][exponent:8][mantissa:7] 范围大但精度低 →
FP32: [sign:1][exponent:8][mantissa:23] → 完整精度混合精度训练流程
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,累积梯度后再更新:
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(基准) | 完整 | 完整 | 完整 | 1× | 基线 |
| ZeRO-1 | 完整 | 完整 | 1/N | ~4x→~3x | 很低 |
| ZeRO-2 | 完整 | 1/N | 1/N | ~3x→~2x | 中(≈DDP) |
| ZeRO-3 | 1/N | 1/N | 1/N | ~N× | 较高 |
内存节省效果
对于参数量为 \(\Psi\) 的模型:
| ZeRO Stage | 每GPU内存 | 相对DDP节省 |
|---|---|---|
| DDP | \(16\Psi\) | 1× |
| ZeRO-1 | \(16\Psi\) | 1× |
| ZeRO-2 | \(2\Psi + 14\Psi/N\) | ~8× (N=8) |
| ZeRO-3 | \(16\Psi/N\) | N× |
其中 \(N\) 是GPU数量。
# 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% | 平衡方案 |
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