跳转至

dist_checkpointing.py — 分布式检查点

文件路径: verl/utils/megatron/dist_checkpointing.py

文件概述

封装 Megatron Core 的分布式检查点 API,提供完全并行的保存和加载功能。

核心函数

1. 保存

def save_dist_checkpointing(sharded_state_dict, ckpt_path, async_save=False, content_metadata=None):
    save_strategy = get_default_save_sharded_strategy("torch_dist")
    save_strategy = FullyParallelSaveStrategyWrapper(
        save_strategy, mpu.get_data_parallel_group(with_context_parallel=True)
    )
    return dist_checkpointing.save(sharded_state_dict, ckpt_path, sharded_strategy=save_strategy, ...)

使用 FullyParallelSaveStrategyWrapper,让数据并行组内的所有 rank 同时写入,大幅加速保存。

2. 加载

def load_dist_checkpointing(sharded_state_dict, ckpt_dir):
    load_strategy = get_default_load_sharded_strategy(ckpt_dir)
    load_strategy = FullyParallelLoadStrategyWrapper(
        load_strategy, mpu.get_data_parallel_group(with_context_parallel=True)
    )
    return dist_checkpointing.load(sharded_state_dict, ckpt_dir, sharded_strategy=load_strategy)

与其他模块的关系

  • 被 checkpoint/megatron_checkpoint_manager.py 调用
  • 依赖 megatron.core.dist_checkpointing 模块

小结

Megatron 分布式检查点的高层封装,实现并行保存/加载。