完整指南:保存与恢复训练数据管道状态)
深度学习数据工程【免费下载链接】DALIA GPU-accelerated library containing highly optimized building blocks and an execution engine for data processing to accelerate deep learning training and inference applications.项目地址https://gitcode.com/gh_mirrors/da/DALI点击查看免费下载导读本文以 NVIDIA DALI 的官方文档 advanced_topics_checkpointing.rst 为骨架系统讲解 DALI 检查点Checkpointing功能如何保存管道Pipeline当前状态到文件并在之后从检查点恢复使新管道与旧管道产生完全一致的输出。该能力对可能被中断的长时间训练任务尤其有价值。读完本文你将掌握enable_checkpointing的开启方式、Pipeline.checkpoint()的保存与checkpoint参数的恢复流程、fn.external_source的部分支持限制以及 TensorFlow 插件中与tf.train.checkpoint的集成方式并结合仓库源码了解底层实现原理。什么是 DALI CheckpointingDALI 的检查点功能允许你把管道Pipeline的当前状态保存到一个文件中之后从该检查点恢复管道恢复后的新管道将产生与旧管道完全相同的输出。这对于运行时间较长、很可能被中途打断的训练任务尤其有用。从实现上看DALI 管道检查点包含两类关键信息见 checkpoint.h 与文档说明管道中所有随机数生成器RNG的状态保证恢复后随机算子如fn.random.uniform生成的随机序列与中断前完全一致每个读取器Reader的进度保证数据读取从中断时的 epoch 与迭代位置继续而不是从头开始。在 C 侧这一设计体现为Checkpoint类——它是整个管道级状态的聚合通过AddOperator(instance_name)为每个算子注册独立的OpCheckpoint并用name2id_映射把算子实例名与检查点条目关联起来Checkpoint还持有iteration_id_当前迭代序号以及来自 Python 侧的ExternalContextCheckpoint包含pipeline_data与iterator_data见 checkpoint.h。关键设计要点检查点保存的是有状态算子的状态。那些不维护用户可观察状态的算子如解码器、resize、归一化等在概念上是无状态的不会进入检查点——这一点在 Dynamic 模式检查点文档 中有明确说明。完整实操示例可参考官方 notebookPipeline checkpointing notebook动态模式Dynamic API的等价流程见 Dynamic mode checkpointing。Checkpointing API 使用详解开启检查点enable_checkpointingTrue要启用检查点功能在创建管道时将enable_checkpointing设为True。开启后DALI 会跟踪每个算子的状态以便按需保存。官方文档明确指出开启检查点不应影响性能。pipeline_def(..., enable_checkpointingTrue) def pipeline(): ... p pipeline()在 Python API 中enable_checkpointing是Pipeline.__init__的可选参数默认值为False见 pipeline.py。该参数会一路传递到 C 侧Python 层在build()时把该参数写入PipelineParams见 pipeline.pyC 层在管道反序列化/构建时调用this-EnableCheckpointing()并在执行器构建阶段通过executor_-EnableCheckpointing(checkpointing_enabled())传递给执行器见 pipeline.cc 与 pipeline.cc在 ProtoBuf 的PipelineDef消息中也保留了一个enable_checkpointing字段默认false见 dali.proto。从执行器内部看启用检查点后Executor2会为每次迭代的IterationData预先创建Checkpoint对象InitIterationData中if (config_.checkpointing) iter_data-checkpoint CreateCheckpoint(...)见 exec2.cc每个算子节点执行完成后会把自身的状态写入该迭代对应的检查点见 exec_node_task.cc。这就是按需保存能够随时取到最新状态的原因。注意shuffle_after_epochTrue的读取器在启用检查点后样本打乱的方式可能与未启用时略有不同。原因是启用检查点后读取器必须在每个 epoch 保存一份初始顺序的备份以便恢复详见下文读取器如何支持检查点一节见 file_label_loader.h。保存检查点Pipeline.checkpoint()保存检查点需要调用Pipeline.checkpoint()方法它返回一个序列化后的检查点字符串内部为序列化后的 Protobuf 消息。也可以传入文件名作为参数DALI 会直接把检查点写入该文件文件内容将被覆盖。for _ in range(iters): output p.run() # 把检查点写入文件 checkpoint p.checkpoint() open(checkpoint_file.cpt, wb) # 或者更简单 checkpoint p.checkpoint(checkpoint_file.cpt)从源码看checkpoint()的实现流程是见 pipeline.py先调用self.build()若尚未构建确保管道已就绪调用_get_checkpoint()通过b.ExternalContextCheckpoint()把 Python 侧的迭代上下文{iter: self._consumer_iter, epoch_idx: self._epoch_idx}JSON 序列化后放入pipeline_data一起打包见 pipeline.py通过self._pipe.GetSerializedCheckpoint(external_ctx_cpt)触发 C 侧序列化若传入了filename则以二进制写模式wb把序列化结果写入文件并返回该字符串。在 C 侧Checkpoint::SerializeToProtobuf会遍历所有算子的OpCheckpoint对每个算子调用op-SerializeCheckpoint(cpt)取回其序列化状态再连同external_ctx_cpt一起打包成 Protobuf 消息返回见 checkpoint.cc。序列化格式定义在 dali.proto 的Checkpoint消息中每个OpCheckpoint包含operator_name算子实例名与operator_state算子状态字节流。注意调用Pipeline.checkpoint()可能会引入可观测的开销。官方建议不要过于频繁地调用它。从源码也可以看出原因每次调用都需要遍历整张算子图、逐个算子收集并序列化状态且 GPU 算子在保存状态时还可能涉及流同步OpCheckpoint::SetOrder(AccessOrder)用于保证异步保存的 GPU 状态在主机侧可见见 op_checkpoint.h。从检查点恢复checkpoint参数之后可以从已保存的检查点恢复管道状态。做法是在构造Pipeline时传入checkpoint参数。恢复后的管道应当产生与原始管道完全一致的输出。checkpoint open(checkpoint_file.cpt, rb).read() p_restored pipeline(checkpointcheckpoint)在 Python API 中checkpoint是Pipeline.__init__的另一个可选参数默认值为None见 pipeline.py。其恢复流程为build()过程中调用_restore_state_from_checkpoint()见 pipeline.py若self._checkpoint is not None则调用 C 侧self._pipe.RestoreFromSerializedCheckpoint(self._checkpoint)并把is_restored_from_checkpoint置为True见 pipeline.py。该属性可通过Pipeline.is_restored_from_checkpoint查询若传入的检查点不是合法的 JSON/Protobuf 数据会抛出错误提示请确保检查点是由相同版本的 DALI 创建的。在 C 侧Executor2::RestoreFromCheckpoint会遍历算子图中的每个算子节点逐个调用n.op-RestoreState(cpt.GetOpCheckpoint(n.instance_name))恢复状态并把执行器的迭代计数设为检查点中保存的iteration_id如果检查点中的算子状态数量多于当前图中的算子还会抛出检查点包含多余算子状态的运行时错误见 exec2.cc。反序列化侧的对应实现是Checkpoint::DeserializeFromProtobuf它会按算子名逐一匹配并把状态交还给对应算子见 checkpoint.cc。警告恢复时必须保证恢复的管道与原始管道相同即包含相同的算子、相同的参数。用不同管道创建的检查点去恢复将导致未定义行为undefined behavior。源码中也能看到对应的防御逻辑RestoreFromCheckpoint要求检查点中的每个算子名都能在当前图中找到找不到或多余都会报错见 exec2.cc。算子层面的检查点接口从算子基类看DALI 为每个算子定义了 4 个与检查点相关的虚函数见 operator.hSaveState(OpCheckpoint cpt, AccessOrder order)把算子状态保存到检查点对象RestoreState(const OpCheckpoint cpt)从检查点恢复算子状态SerializeCheckpoint(const OpCheckpoint cpt)把算子状态序列化为字符串DeserializeCheckpoint(OpCheckpoint cpt, const std::string data)反序列化并填充检查点对象。默认实现会调用CheckpointingUnsupportedError()——即该算子未实现检查点。只有真正有状态、且实现了这些接口的算子读取器、RNG 相关算子等才参与检查点。单算子级别的测试覆盖可参考 checkpoint_test.cc其中分别验证了 CPU-only、GPU-only、混合Mixed三种图结构以及序列化/反序列化往返CheckpointTest.CPUOnly、GPUOnly、Mixed、Serialize见 checkpoint_test.cc。读取器如何支持检查点原理剖析读取器是检查点最主要的受益者。DALI 的读取器由 Loader负责实际读样本驱动其检查点支持在 loader.h 中定义LoaderStateSnapshot保存了读取器在 epoch 开始时的基础状态——rng随机数引擎、current_epoch当前 epoch 数和age样本年龄计数见 loader.h。Loader 构造时通过options.GetArgumentbool(checkpointing)读取检查点开关见 loader.h并在Init()阶段就保存一份初始快照。具体到FileReader基于file_label_loader.h启用检查点且设置了shuffle_after_epoch时PrepareMetadata会先保存一份初始文件顺序的备份backup_file_label_entries_见 file_label_loader.h在每个 epoch 的Reset中启用检查点时从备份顺序重新洗牌file_label_entries_ backup_file_label_entries_洗牌种子由shuffle_after_epoch_seed_ (current_epoch_ 32)推导——每个 epoch 用不同种子因此恢复后仍能保证随机分布不受影响同时顺序可复现见 file_label_loader.hRestoreStateImpl只需恢复current_epoch其余索引状态由 loader 基类的快照机制统一处理见 file_label_loader.h。这正好解释了文档中的那条注意事项shuffle_after_epochTrue时启用检查点后打乱方式可能略有不同——因为可恢复性优先打乱顺序的推导方式被调整了。另外file_reader_op.cc中shuffle_after_epoch的文档还提到使用shuffle_after_epoch时不能同时使用stick_to_shard和random_shuffle多 GPU 场景下所有管道实例应使用相同的shuffle_after_epoch_seed以保证全局一致的洗牌见 file_reader_op.cc。External Source 的检查点支持部分支持fn.external_source算子仅部分支持检查点。支持的场景只有当source是单参数可调用对象callable且该参数为以下三者之一时检查点才受支持批次索引batch indexBatchInfoSampleInfo。对于这类source恢复检查点后查询会从检查点中保存的位置继续epoch 与迭代都会对齐。从源码看Python 侧_check_checkpointing_support的实现逻辑正是如此只有kind _SourceKind.CALLABLE and has_inputs即可调用且带参数才算支持检查点否则会发出警告见 pipeline.py。external_source的文档也明确注明恢复检查点后单参数可调用 source 的查询会从检查点保存的 epoch 和迭代继续见 external_source.py。其底层依靠callback_args中基于current_iter/epoch_idx的索引推导见 external_source.py。不支持的场景其他类型的source如无参可调用、迭代器等不支持检查点。它们的状态不会被保存进检查点恢复后这些 source 会从头开始。如果与从中间恢复的读取器搭配使用可能导致数据错位。官方建议如果你需要使用检查点推荐把 source 改写成受支持的单参数可调用形式。例如def my_source(sample_info: SampleInfo): # 依据 sample_info.idx_in_epoch / iteration 返回对应样本 return data[sample_info.idx_in_epoch] pipe pipeline_def(..., enable_checkpointingTrue)(...)TensorFlow 插件中的检查点nvidia.dali.plugin.tf.DALIDataset与 TensorFlow 的tf.train.checkpoint机制深度集成——这意味着你可以用 TensorFlow 标准的检查点 API如tf.train.Checkpoint手动保存/恢复来同时保存 DALI 管道的状态与模型权重无需额外的 DALI 专属代码路径。在插件实现中DALIDataset的保存save与恢复restore钩子分别通过 C API 完成保存时调用daliPipelineGetCheckpoint拿到检查点句柄再经daliPipelineSerializeCheckpoint序列化以名为checkpoint的 Tensor 写入 TensorFlow 的检查点文件见 dali_dataset_op.cc恢复时从检查点文件读取checkpointTensor经daliPipelineDeserializeCheckpointdaliPipelineRestoreCheckpoint恢复到管道见 dali_dataset_op.cc。重要限制插件层面DALIDatasetWithInputs暂不支持检查点。其checkCheckpointingSupport()会直接返回Unimplemented错误Checkpointing is not supported for DALI dataset with inputs见 dali_dataset_op.ccGPU 数据集暂不支持检查点。同样由checkCheckpointingSupport()抛出 Checkpointing is not supported for DALI GPU dataset见 dali_dataset_op.cc。警告使用 TensorFlow 插件时请确保满足上述两个前提非DALIDatasetWithInputs、非 GPU 数据集否则会在保存/恢复阶段直接报错。常见问题与最佳实践1. 检查点保存多久调用一次合适checkpoint()有可观测的开销需要遍历整张算子图、序列化所有算子状态不要每轮迭代都调用。合理做法是每隔固定的若干迭代如每 N 个 epoch或结合训练框架的定期保存如 TensorFlow 的tf.train.CheckpointManager、PyTorch 的torch.utils.checkpoint风格定期存盘调用一次。2. 恢复后的管道输出一定一致吗只要恢复的管道与原始管道结构、参数完全相同且 source 是受支持的单参数 callable恢复后的输出应当与原始管道完全一致——这正是检查点保存 RNG 状态与 reader 进度的意义。反之管道不同、source 不支持则可能产生未定义行为或数据错位。3. 检查点文件是什么格式checkpoint()返回的是一个序列化后的 Protobuf 字符串其中包含每个算子的状态operator_nameoperator_state以及 Python 侧的pipeline_data/iterator_data。请用二进制模式读写示例中open(checkpoint_file.cpt, wb)/rb。4. 启用检查点影响性能吗官方文档明确说明不应有任何影响。从实现看启用后执行器只是为每次迭代额外维护一个Checkpoint对象并在算子执行后写入状态见 exec2.cc开销主要体现在调用checkpoint()主动保存的那一刻。5. 多管道 / 多 GPU 场景如果使用多个管道如base_iterator.py中的多管道迭代器所有管道必须设置相同的enable_checkpointing值否则会抛ValueError见 base_iterator.py。若部分管道从检查点恢复而部分没有迭代器会发出警告并可能出现意外结果见 base_iterator.py。总结DALI 的检查点功能为长时间训练任务提供了可靠的中断恢复手段开启pipeline_def(..., enable_checkpointingTrue)保存p.checkpoint()或p.checkpoint(file.cpt)返回序列化 Protobuf 字符串恢复构造时传入checkpoint...恢复后的管道输出与原管道完全一致核心内容所有 RNG 状态 每个读取器的进度限制fn.external_source仅支持单参数 callableTensorFlow 插件的DALIDatasetWithInputs与 GPU 数据集暂不支持。结合源码可以确认底层由执行器Executor2逐迭代维护检查点、每个算子通过SaveState/RestoreState/SerializeCheckpoint/DeserializeCheckpoint四个接口参与状态保存与恢复最终以 Protobuf 消息整体序列化。相关实现与测试文件包括pipeline.pyPython APIcheckpoint/_get_checkpoint/_restore_state_from_checkpointcheckpoint.h 与 checkpoint.cc管道级检查点聚合与序列化exec2.cc执行器侧保存/恢复loader.h 与 file_label_loader.h读取器状态快照checkpoint_test.ccCPU/GPU/混合图及序列化往返测试dali_dataset_op.ccTensorFlow 插件集成赞分享深度学习数据工程【免费下载链接】DALIA GPU-accelerated library containing highly optimized building blocks and an execution engine for data processing to accelerate deep learning training and inference applications.项目地址https://gitcode.com/gh_mirrors/da/DALI点击查看免费下载相关推荐verl检查点管理训练状态保存与恢复verl检查点管理训练状态保存与恢复 概述 在大规模语言模型LLM的强化学习训练过程中verlVolcano Engine Reinforcement人工智能大模型强化学习RLHF分布式训练微调LeRobot机器人学习3步构建你的第一个AI机器人控制模型LeRobot机器人学习3步构建你的第一个AI机器人控制模型 想不想让机器人像人一样学习新技能 你是否曾梦想过让机械臂学会抓取物体、让机器人自主完成复杂人工智能机器学习深度学习机器人具身智能强化学习55项功能全面升级HsMod插件让你的炉石传说体验飞升8倍速55项功能全面升级HsMod插件让你的炉石传说体验飞升8倍速 HsMod是一款基于BepInEx框架开发的炉石传说游戏增强插件为玩家提供了从游戏性能优化到社游戏开发上一篇在VS Code中使用Windows Subsystem for Linux(WSL)进行开发下一篇FlyEnv v4.9.7 版本更新优化 FTP 服务与 PHP 安全配置创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考