分布式检查点保存与恢复
大模型训练的检查点保存和恢复是工程中不可忽视的问题。分布式训练下,检查点涉及多个 GPU 的状态协调,需要高效的存储格式和恢复机制。
检查点内容
一个完整的训练检查点包含:
| 组件 | 大小 (7B FP16) | 说明 |
|---|---|---|
| 模型参数 | 14GB | 各 GPU 的参数分片 |
| 优化器状态 | 28GB | Adam m/v(FP32) |
| 梯度 | 14GB | 当前步梯度 |
| RNG 状态 | 数 KB | 随机数生成器状态 |
| 训练状态 | 数 KB | step、lr、loss 等 |
| 数据加载器 | 数 KB | 采样位置 |
PyTorch 分布式检查点
python
from torch.distributed.checkpoint import save, load
from torch.distributed.checkpoint.state_dict import get_state_dict, set_state_dict
# 保存检查点
def save_checkpoint(model, optimizer, step, path):
state_dict = {
"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
"step": step,
}
save(state_dict, checkpoint_path=path)
# 加载检查点
def load_checkpoint(model, optimizer, path):
state_dict = {
"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
}
load(state_dict, checkpoint_path=path)
set_state_dict(model, optimizer, state_dict["model"], state_dict["optimizer"])DeepSpeed 检查点
python
# DeepSpeed 自动保存所有状态
model_engine.save_checkpoint(save_dir, tag=f"step-{step}")
# 恢复训练
model_engine.load_checkpoint(load_dir, tag="step-1000")检查点结构
checkpoints/
└── step-1000/
├── zero_pp_rank_0_mp_rank_00_model_states.pt
├── zero_pp_rank_1_mp_rank_00_model_states.pt
├── zero_pp_rank_0_mp_rank_00_optim_states.pt
├── zero_pp_rank_1_mp_rank_00_optim_states.pt
└── latest异步保存
同步保存会阻塞训练,异步保存将 I/O 操作移到后台:
python
# 使用独立进程异步保存
import subprocess
def async_save_checkpoint(state_dict, path):
# 先保存到本地临时路径
temp_path = f"/tmp/checkpoint_{step}"
torch.save(state_dict, temp_path)
# 后台拷贝到共享存储
subprocess.Popen(
["cp", "-r", temp_path, path],
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
)检查点频率
- 保存频率:每 1000-5000 步保存一次
- 保留数量:保留最近 3-5 个检查点
- 关键节点:训练阶段性节点(如学习率变化前)务必保存
- 预估大小:7B 模型单次检查点约 42GB,70B 约 420GB
并行度变更
从检查点恢复时,并行度(TP/PP/DP)必须与保存时相同。如需变更并行度,需使用支持重分片的检查点格式(如 PyTorch DCP 或 HuggingFace SafeTensors)。