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 分布式检查点的高层封装,实现并行保存/加载。