FlexCheckPoint框架笔记整理
1.问题背景
Checkpoint负责对模型参数、优化器、数据流、随即状态以及徐连所需配置信息进行持久化保存,以便在训练任务出现故障中断后可以重新恢复状态(续训)。
在大模型相同训练任务的不同阶段以及不同任务之间,由于上下游的分布式并行策略以及模型结构可能发生改变往往需要进行checkpoint的转化和迁移,然而很多场景下都是针对不同的模型和任务 case by case手动定制checkpoint切分转换脚本,脚本的开发和维护成本高,且难以复用。
- 预训练场景
- 变换并行策略(无模型结构变化)
- 变换sharding数:由于训练卡量变化,在保持其他并行策略不变的情况下,做sharding reshard,这是比较常见的场景。
- 进行不同并行策略的转换:在dp、pp、tp、sharding、ep之间互转,例如以dp2pp2训练保存checkpoint,而用sharding4做接续训练。
- 训练的ckpt离线转推理做下游评估
- 训练ckpt转推理:合pp和sharding,保留tp的模型结构做推理
- 推理不关心训练的并行方法:推理组只想拿到参数完整的safetensors格式的权重文件,用最原始的模型做推理,而不需要关心,训练用的什么策略,去对组网做特殊处理。
- 跨模型结构的参数融合:
- MOE模型融合不同模态的专家
- 训练中途括tokenizer的词表维度,训练段Linear参数的shape发生变化
- fused_qkv,fused_ffn在不同tp并行数下,需要重排做二次划分。
- fused_qkv,fused_ffn转非fused_qkv,fused_ffn做训练。
- 变换并行策略(无模型结构变化)
- 开源场景
- 飞桨的checkpoint数据使用pickel存储,是pdparam格式,不支持safetensors格式,有一套UC框架,是直接训练结束时保存成safetensors格式,但无法直接转换。开源时每次要准配两份权重,并且要case by case写转换脚本。
- 部分参数存储形式和开源社区习惯不同,例如Linear的weight存储transpose形式等。
因此为了实现解决上述问题,需要一套新的checkpoint系统来降低转换开销,支持任意并行策略的互相转换,支持参数融合,模型结构的改变,并且在线转换的时间控制在秒级。
2.FlexCheckpoint设计整体框架

以上是FC的整体架构图,主要由两个部分组成,一部分是DCP,另一部分是AOA。
考虑到DCP协议存储和转换过程与分布式策略解耦,其灵活性和通用性更强,支持零冗余加载,并且能兼容飞桨自动并行生态,因此FC的底层复用DCP协议实现切分转换,在切分标记上,FC借鉴Megatron的思想,由用户在Layer中标记ShardedStateDict,与Megatron一样,如果用户调用框架提供的分布式API,例如像ColumnParallel、RowParallel、VocabEmbedingParallel等并行组件,无需用户做额外操作,框架自动标记好。
同时为了实现跨模型结构转换,比如参数的一些合并、拆分、转置、替换等等,我们进一步提出了AOA(all-in-One-arrow),该协议允许用户在单卡视角下使用非常简洁的箭头语言表达跨模型结构转换的信息,而不需要关心参数的具体切分方式。AOA提供了7种操作原语来表达所有的转换操作,包括split、merge、rename、add、remove、transpose、cast,用户只需要用简单的箭头语言表达参数转换逻辑,AOA协议会自动将这些语义翻译成底层参数切片的转换操作,从而实现跨模型结构转换的机制。
3.跨并行策略参数转换(DCP)
