Skip to content

分布式检查点保存与恢复

大模型训练的检查点保存和恢复是工程中不可忽视的问题。分布式训练下,检查点涉及多个 GPU 的状态协调,需要高效的存储格式和恢复机制。

分布式检查点保存与恢复

检查点内容

一个完整的训练检查点包含:

组件大小 (7B FP16)说明
模型参数14GB各 GPU 的参数分片
优化器状态28GBAdam m/v(FP32)
梯度14GB当前步梯度
RNG 状态数 KB随机数生成器状态
训练状态数 KBstep、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)。

相关资源

最近更新