(Checkpointing)完整指南:保存與恢復(fù)訓(xùn)練數(shù)據(jù)管道狀態(tài))
深度學(xué)習(xí)數(shù)據(jù)工程【免費(fèi)下載鏈接】DALIA GPU-accelerated library containing highly optimized building blocks and an execution engine for data processing to accelerate deep learning training and inference applications.項(xiàng)目地址https://gitcode.com/gh_mirrors/da/DALI點(diǎn)擊查看免費(fèi)下載導(dǎo)讀本文以 NVIDIA DALI 的官方文檔 advanced_topics_checkpointing.rst 為骨架系統(tǒng)講解 DALI 檢查點(diǎn)Checkpointing功能如何保存管道Pipeline當(dāng)前狀態(tài)到文件并在之后從檢查點(diǎn)恢復(fù)使新管道與舊管道產(chǎn)生完全一致的輸出。該能力對(duì)可能被中斷的長(zhǎng)時(shí)間訓(xùn)練任務(wù)尤其有價(jià)值。讀完本文你將掌握enable_checkpointing的開啟方式、Pipeline.checkpoint()的保存與checkpoint參數(shù)的恢復(fù)流程、fn.external_source的部分支持限制以及 TensorFlow 插件中與tf.train.checkpoint的集成方式并結(jié)合倉(cāng)庫(kù)源碼了解底層實(shí)現(xiàn)原理。什么是 DALI CheckpointingDALI 的檢查點(diǎn)功能允許你把管道Pipeline的當(dāng)前狀態(tài)保存到一個(gè)文件中之后從該檢查點(diǎn)恢復(fù)管道恢復(fù)后的新管道將產(chǎn)生與舊管道完全相同的輸出。這對(duì)于運(yùn)行時(shí)間較長(zhǎng)、很可能被中途打斷的訓(xùn)練任務(wù)尤其有用。從實(shí)現(xiàn)上看DALI 管道檢查點(diǎn)包含兩類關(guān)鍵信息見 checkpoint.h 與文檔說明管道中所有隨機(jī)數(shù)生成器RNG的狀態(tài)保證恢復(fù)后隨機(jī)算子如fn.random.uniform生成的隨機(jī)序列與中斷前完全一致每個(gè)讀取器Reader的進(jìn)度保證數(shù)據(jù)讀取從中斷時(shí)的 epoch 與迭代位置繼續(xù)而不是從頭開始。在 C 側(cè)這一設(shè)計(jì)體現(xiàn)為Checkpoint類——它是整個(gè)管道級(jí)狀態(tài)的聚合通過AddOperator(instance_name)為每個(gè)算子注冊(cè)獨(dú)立的OpCheckpoint并用name2id_映射把算子實(shí)例名與檢查點(diǎn)條目關(guān)聯(lián)起來Checkpoint還持有iteration_id_當(dāng)前迭代序號(hào)以及來自 Python 側(cè)的ExternalContextCheckpoint包含pipeline_data與iterator_data見 checkpoint.h。關(guān)鍵設(shè)計(jì)要點(diǎn)檢查點(diǎn)保存的是有狀態(tài)算子的狀態(tài)。那些不維護(hù)用戶可觀察狀態(tài)的算子如解碼器、resize、歸一化等在概念上是無狀態(tài)的不會(huì)進(jìn)入檢查點(diǎn)——這一點(diǎn)在 Dynamic 模式檢查點(diǎn)文檔 中有明確說明。完整實(shí)操示例可參考官方 notebookPipeline checkpointing notebook動(dòng)態(tài)模式Dynamic API的等價(jià)流程見 Dynamic mode checkpointing。Checkpointing API 使用詳解開啟檢查點(diǎn)enable_checkpointingTrue要啟用檢查點(diǎn)功能在創(chuàng)建管道時(shí)將enable_checkpointing設(shè)為True。開啟后DALI 會(huì)跟蹤每個(gè)算子的狀態(tài)以便按需保存。官方文檔明確指出開啟檢查點(diǎn)不應(yīng)影響性能。pipeline_def(..., enable_checkpointingTrue) def pipeline(): ... p pipeline()在 Python API 中enable_checkpointing是Pipeline.__init__的可選參數(shù)默認(rèn)值為False見 pipeline.py。該參數(shù)會(huì)一路傳遞到 C 側(cè)Python 層在build()時(shí)把該參數(shù)寫入PipelineParams見 pipeline.pyC 層在管道反序列化/構(gòu)建時(shí)調(diào)用this-EnableCheckpointing()并在執(zhí)行器構(gòu)建階段通過executor_-EnableCheckpointing(checkpointing_enabled())傳遞給執(zhí)行器見 pipeline.cc 與 pipeline.cc在 ProtoBuf 的PipelineDef消息中也保留了一個(gè)enable_checkpointing字段默認(rèn)false見 dali.proto。從執(zhí)行器內(nèi)部看啟用檢查點(diǎn)后Executor2會(huì)為每次迭代的IterationData預(yù)先創(chuàng)建Checkpoint對(duì)象InitIterationData中if (config_.checkpointing) iter_data-checkpoint CreateCheckpoint(...)見 exec2.cc每個(gè)算子節(jié)點(diǎn)執(zhí)行完成后會(huì)把自身的狀態(tài)寫入該迭代對(duì)應(yīng)的檢查點(diǎn)見 exec_node_task.cc。這就是按需保存能夠隨時(shí)取到最新狀態(tài)的原因。注意shuffle_after_epochTrue的讀取器在啟用檢查點(diǎn)后樣本打亂的方式可能與未啟用時(shí)略有不同。原因是啟用檢查點(diǎn)后讀取器必須在每個(gè) epoch 保存一份初始順序的備份以便恢復(fù)詳見下文讀取器如何支持檢查點(diǎn)一節(jié)見 file_label_loader.h。保存檢查點(diǎn)Pipeline.checkpoint()保存檢查點(diǎn)需要調(diào)用Pipeline.checkpoint()方法它返回一個(gè)序列化后的檢查點(diǎn)字符串內(nèi)部為序列化后的 Protobuf 消息。也可以傳入文件名作為參數(shù)DALI 會(huì)直接把檢查點(diǎn)寫入該文件文件內(nèi)容將被覆蓋。for _ in range(iters): output p.run() # 把檢查點(diǎn)寫入文件 checkpoint p.checkpoint() open(checkpoint_file.cpt, wb) # 或者更簡(jiǎn)單 checkpoint p.checkpoint(checkpoint_file.cpt)從源碼看checkpoint()的實(shí)現(xiàn)流程是見 pipeline.py先調(diào)用self.build()若尚未構(gòu)建確保管道已就緒調(diào)用_get_checkpoint()通過b.ExternalContextCheckpoint()把 Python 側(cè)的迭代上下文{iter: self._consumer_iter, epoch_idx: self._epoch_idx}JSON 序列化后放入pipeline_data一起打包見 pipeline.py通過self._pipe.GetSerializedCheckpoint(external_ctx_cpt)觸發(fā) C 側(cè)序列化若傳入了filename則以二進(jìn)制寫模式wb把序列化結(jié)果寫入文件并返回該字符串。在 C 側(cè)Checkpoint::SerializeToProtobuf會(huì)遍歷所有算子的OpCheckpoint對(duì)每個(gè)算子調(diào)用op-SerializeCheckpoint(cpt)取回其序列化狀態(tài)再連同external_ctx_cpt一起打包成 Protobuf 消息返回見 checkpoint.cc。序列化格式定義在 dali.proto 的Checkpoint消息中每個(gè)OpCheckpoint包含operator_name算子實(shí)例名與operator_state算子狀態(tài)字節(jié)流。注意調(diào)用Pipeline.checkpoint()可能會(huì)引入可觀測(cè)的開銷。官方建議不要過于頻繁地調(diào)用它。從源碼也可以看出原因每次調(diào)用都需要遍歷整張算子圖、逐個(gè)算子收集并序列化狀態(tài)且 GPU 算子在保存狀態(tài)時(shí)還可能涉及流同步OpCheckpoint::SetOrder(AccessOrder)用于保證異步保存的 GPU 狀態(tài)在主機(jī)側(cè)可見見 op_checkpoint.h。從檢查點(diǎn)恢復(fù)checkpoint參數(shù)之后可以從已保存的檢查點(diǎn)恢復(fù)管道狀態(tài)。做法是在構(gòu)造Pipeline時(shí)傳入checkpoint參數(shù)?;謴?fù)后的管道應(yīng)當(dāng)產(chǎn)生與原始管道完全一致的輸出。checkpoint open(checkpoint_file.cpt, rb).read() p_restored pipeline(checkpointcheckpoint)在 Python API 中checkpoint是Pipeline.__init__的另一個(gè)可選參數(shù)默認(rèn)值為None見 pipeline.py。其恢復(fù)流程為build()過程中調(diào)用_restore_state_from_checkpoint()見 pipeline.py若self._checkpoint is not None則調(diào)用 C 側(cè)self._pipe.RestoreFromSerializedCheckpoint(self._checkpoint)并把is_restored_from_checkpoint置為True見 pipeline.py。該屬性可通過Pipeline.is_restored_from_checkpoint查詢?nèi)魝魅氲臋z查點(diǎn)不是合法的 JSON/Protobuf 數(shù)據(jù)會(huì)拋出錯(cuò)誤提示請(qǐng)確保檢查點(diǎn)是由相同版本的 DALI 創(chuàng)建的。在 C 側(cè)Executor2::RestoreFromCheckpoint會(huì)遍歷算子圖中的每個(gè)算子節(jié)點(diǎn)逐個(gè)調(diào)用n.op-RestoreState(cpt.GetOpCheckpoint(n.instance_name))恢復(fù)狀態(tài)并把執(zhí)行器的迭代計(jì)數(shù)設(shè)為檢查點(diǎn)中保存的iteration_id如果檢查點(diǎn)中的算子狀態(tài)數(shù)量多于當(dāng)前圖中的算子還會(huì)拋出檢查點(diǎn)包含多余算子狀態(tài)的運(yùn)行時(shí)錯(cuò)誤見 exec2.cc。反序列化側(cè)的對(duì)應(yīng)實(shí)現(xiàn)是Checkpoint::DeserializeFromProtobuf它會(huì)按算子名逐一匹配并把狀態(tài)交還給對(duì)應(yīng)算子見 checkpoint.cc。警告恢復(fù)時(shí)必須保證恢復(fù)的管道與原始管道相同即包含相同的算子、相同的參數(shù)。用不同管道創(chuàng)建的檢查點(diǎn)去恢復(fù)將導(dǎo)致未定義行為undefined behavior。源碼中也能看到對(duì)應(yīng)的防御邏輯RestoreFromCheckpoint要求檢查點(diǎn)中的每個(gè)算子名都能在當(dāng)前圖中找到找不到或多余都會(huì)報(bào)錯(cuò)見 exec2.cc。算子層面的檢查點(diǎn)接口從算子基類看DALI 為每個(gè)算子定義了 4 個(gè)與檢查點(diǎn)相關(guān)的虛函數(shù)見 operator.hSaveState(OpCheckpoint cpt, AccessOrder order)把算子狀態(tài)保存到檢查點(diǎn)對(duì)象RestoreState(const OpCheckpoint cpt)從檢查點(diǎn)恢復(fù)算子狀態(tài)SerializeCheckpoint(const OpCheckpoint cpt)把算子狀態(tài)序列化為字符串DeserializeCheckpoint(OpCheckpoint cpt, const std::string data)反序列化并填充檢查點(diǎn)對(duì)象。默認(rèn)實(shí)現(xiàn)會(huì)調(diào)用CheckpointingUnsupportedError()——即該算子未實(shí)現(xiàn)檢查點(diǎn)。只有真正有狀態(tài)、且實(shí)現(xiàn)了這些接口的算子讀取器、RNG 相關(guān)算子等才參與檢查點(diǎn)。單算子級(jí)別的測(cè)試覆蓋可參考 checkpoint_test.cc其中分別驗(yàn)證了 CPU-only、GPU-only、混合Mixed三種圖結(jié)構(gòu)以及序列化/反序列化往返CheckpointTest.CPUOnly、GPUOnly、Mixed、Serialize見 checkpoint_test.cc。讀取器如何支持檢查點(diǎn)原理剖析讀取器是檢查點(diǎn)最主要的受益者。DALI 的讀取器由 Loader負(fù)責(zé)實(shí)際讀樣本驅(qū)動(dòng)其檢查點(diǎn)支持在 loader.h 中定義LoaderStateSnapshot保存了讀取器在 epoch 開始時(shí)的基礎(chǔ)狀態(tài)——rng隨機(jī)數(shù)引擎、current_epoch當(dāng)前 epoch 數(shù)和age樣本年齡計(jì)數(shù)見 loader.h。Loader 構(gòu)造時(shí)通過options.GetArgumentbool(checkpointing)讀取檢查點(diǎn)開關(guān)見 loader.h并在Init()階段就保存一份初始快照。具體到FileReader基于file_label_loader.h啟用檢查點(diǎn)且設(shè)置了shuffle_after_epoch時(shí)PrepareMetadata會(huì)先保存一份初始文件順序的備份backup_file_label_entries_見 file_label_loader.h在每個(gè) epoch 的Reset中啟用檢查點(diǎn)時(shí)從備份順序重新洗牌file_label_entries_ backup_file_label_entries_洗牌種子由shuffle_after_epoch_seed_ (current_epoch_ 32)推導(dǎo)——每個(gè) epoch 用不同種子因此恢復(fù)后仍能保證隨機(jī)分布不受影響同時(shí)順序可復(fù)現(xiàn)見 file_label_loader.hRestoreStateImpl只需恢復(fù)current_epoch其余索引狀態(tài)由 loader 基類的快照機(jī)制統(tǒng)一處理見 file_label_loader.h。這正好解釋了文檔中的那條注意事項(xiàng)shuffle_after_epochTrue時(shí)啟用檢查點(diǎn)后打亂方式可能略有不同——因?yàn)榭苫謴?fù)性優(yōu)先打亂順序的推導(dǎo)方式被調(diào)整了。另外file_reader_op.cc中shuffle_after_epoch的文檔還提到使用shuffle_after_epoch時(shí)不能同時(shí)使用stick_to_shard和random_shuffle多 GPU 場(chǎng)景下所有管道實(shí)例應(yīng)使用相同的shuffle_after_epoch_seed以保證全局一致的洗牌見 file_reader_op.cc。External Source 的檢查點(diǎn)支持部分支持fn.external_source算子僅部分支持檢查點(diǎn)。支持的場(chǎng)景只有當(dāng)source是單參數(shù)可調(diào)用對(duì)象callable且該參數(shù)為以下三者之一時(shí)檢查點(diǎn)才受支持批次索引batch indexBatchInfoSampleInfo。對(duì)于這類source恢復(fù)檢查點(diǎn)后查詢會(huì)從檢查點(diǎn)中保存的位置繼續(xù)epoch 與迭代都會(huì)對(duì)齊。從源碼看Python 側(cè)_check_checkpointing_support的實(shí)現(xiàn)邏輯正是如此只有kind _SourceKind.CALLABLE and has_inputs即可調(diào)用且?guī)?shù)才算支持檢查點(diǎn)否則會(huì)發(fā)出警告見 pipeline.py。external_source的文檔也明確注明恢復(fù)檢查點(diǎn)后單參數(shù)可調(diào)用 source 的查詢會(huì)從檢查點(diǎn)保存的 epoch 和迭代繼續(xù)見 external_source.py。其底層依靠callback_args中基于current_iter/epoch_idx的索引推導(dǎo)見 external_source.py。不支持的場(chǎng)景其他類型的source如無參可調(diào)用、迭代器等不支持檢查點(diǎn)。它們的狀態(tài)不會(huì)被保存進(jìn)檢查點(diǎn)恢復(fù)后這些 source 會(huì)從頭開始。如果與從中間恢復(fù)的讀取器搭配使用可能導(dǎo)致數(shù)據(jù)錯(cuò)位。官方建議如果你需要使用檢查點(diǎn)推薦把 source 改寫成受支持的單參數(shù)可調(diào)用形式。例如def my_source(sample_info: SampleInfo): # 依據(jù) sample_info.idx_in_epoch / iteration 返回對(duì)應(yīng)樣本 return data[sample_info.idx_in_epoch] pipe pipeline_def(..., enable_checkpointingTrue)(...)TensorFlow 插件中的檢查點(diǎn)nvidia.dali.plugin.tf.DALIDataset與 TensorFlow 的tf.train.checkpoint機(jī)制深度集成——這意味著你可以用 TensorFlow 標(biāo)準(zhǔn)的檢查點(diǎn) API如tf.train.Checkpoint手動(dòng)保存/恢復(fù)來同時(shí)保存 DALI 管道的狀態(tài)與模型權(quán)重?zé)o需額外的 DALI 專屬代碼路徑。在插件實(shí)現(xiàn)中DALIDataset的保存save與恢復(fù)restore鉤子分別通過 C API 完成保存時(shí)調(diào)用daliPipelineGetCheckpoint拿到檢查點(diǎn)句柄再經(jīng)daliPipelineSerializeCheckpoint序列化以名為checkpoint的 Tensor 寫入 TensorFlow 的檢查點(diǎn)文件見 dali_dataset_op.cc恢復(fù)時(shí)從檢查點(diǎn)文件讀取checkpointTensor經(jīng)daliPipelineDeserializeCheckpointdaliPipelineRestoreCheckpoint恢復(fù)到管道見 dali_dataset_op.cc。重要限制插件層面DALIDatasetWithInputs暫不支持檢查點(diǎn)。其checkCheckpointingSupport()會(huì)直接返回Unimplemented錯(cuò)誤Checkpointing is not supported for DALI dataset with inputs見 dali_dataset_op.ccGPU 數(shù)據(jù)集暫不支持檢查點(diǎn)。同樣由checkCheckpointingSupport()拋出 Checkpointing is not supported for DALI GPU dataset見 dali_dataset_op.cc。警告使用 TensorFlow 插件時(shí)請(qǐng)確保滿足上述兩個(gè)前提非DALIDatasetWithInputs、非 GPU 數(shù)據(jù)集否則會(huì)在保存/恢復(fù)階段直接報(bào)錯(cuò)。常見問題與最佳實(shí)踐1. 檢查點(diǎn)保存多久調(diào)用一次合適checkpoint()有可觀測(cè)的開銷需要遍歷整張算子圖、序列化所有算子狀態(tài)不要每輪迭代都調(diào)用。合理做法是每隔固定的若干迭代如每 N 個(gè) epoch或結(jié)合訓(xùn)練框架的定期保存如 TensorFlow 的tf.train.CheckpointManager、PyTorch 的torch.utils.checkpoint風(fēng)格定期存盤調(diào)用一次。2. 恢復(fù)后的管道輸出一定一致嗎只要恢復(fù)的管道與原始管道結(jié)構(gòu)、參數(shù)完全相同且 source 是受支持的單參數(shù) callable恢復(fù)后的輸出應(yīng)當(dāng)與原始管道完全一致——這正是檢查點(diǎn)保存 RNG 狀態(tài)與 reader 進(jìn)度的意義。反之管道不同、source 不支持則可能產(chǎn)生未定義行為或數(shù)據(jù)錯(cuò)位。3. 檢查點(diǎn)文件是什么格式checkpoint()返回的是一個(gè)序列化后的 Protobuf 字符串其中包含每個(gè)算子的狀態(tài)operator_nameoperator_state以及 Python 側(cè)的pipeline_data/iterator_data。請(qǐng)用二進(jìn)制模式讀寫示例中open(checkpoint_file.cpt, wb)/rb。4. 啟用檢查點(diǎn)影響性能嗎官方文檔明確說明不應(yīng)有任何影響。從實(shí)現(xiàn)看啟用后執(zhí)行器只是為每次迭代額外維護(hù)一個(gè)Checkpoint對(duì)象并在算子執(zhí)行后寫入狀態(tài)見 exec2.cc開銷主要體現(xiàn)在調(diào)用checkpoint()主動(dòng)保存的那一刻。5. 多管道 / 多 GPU 場(chǎng)景如果使用多個(gè)管道如base_iterator.py中的多管道迭代器所有管道必須設(shè)置相同的enable_checkpointing值否則會(huì)拋ValueError見 base_iterator.py。若部分管道從檢查點(diǎn)恢復(fù)而部分沒有迭代器會(huì)發(fā)出警告并可能出現(xiàn)意外結(jié)果見 base_iterator.py??偨Y(jié)DALI 的檢查點(diǎn)功能為長(zhǎng)時(shí)間訓(xùn)練任務(wù)提供了可靠的中斷恢復(fù)手段開啟pipeline_def(..., enable_checkpointingTrue)保存p.checkpoint()或p.checkpoint(file.cpt)返回序列化 Protobuf 字符串恢復(fù)構(gòu)造時(shí)傳入checkpoint...恢復(fù)后的管道輸出與原管道完全一致核心內(nèi)容所有 RNG 狀態(tài) 每個(gè)讀取器的進(jìn)度限制fn.external_source僅支持單參數(shù) callableTensorFlow 插件的DALIDatasetWithInputs與 GPU 數(shù)據(jù)集暫不支持。結(jié)合源碼可以確認(rèn)底層由執(zhí)行器Executor2逐迭代維護(hù)檢查點(diǎn)、每個(gè)算子通過SaveState/RestoreState/SerializeCheckpoint/DeserializeCheckpoint四個(gè)接口參與狀態(tài)保存與恢復(fù)最終以 Protobuf 消息整體序列化。相關(guān)實(shí)現(xiàn)與測(cè)試文件包括pipeline.pyPython APIcheckpoint/_get_checkpoint/_restore_state_from_checkpointcheckpoint.h 與 checkpoint.cc管道級(jí)檢查點(diǎn)聚合與序列化exec2.cc執(zhí)行器側(cè)保存/恢復(fù)loader.h 與 file_label_loader.h讀取器狀態(tài)快照checkpoint_test.ccCPU/GPU/混合圖及序列化往返測(cè)試dali_dataset_op.ccTensorFlow 插件集成贊分享深度學(xué)習(xí)數(shù)據(jù)工程【免費(fèi)下載鏈接】DALIA GPU-accelerated library containing highly optimized building blocks and an execution engine for data processing to accelerate deep learning training and inference applications.項(xiàng)目地址https://gitcode.com/gh_mirrors/da/DALI點(diǎn)擊查看免費(fèi)下載相關(guān)推薦verl檢查點(diǎn)管理訓(xùn)練狀態(tài)保存與恢復(fù)verl檢查點(diǎn)管理訓(xùn)練狀態(tài)保存與恢復(fù) 概述 在大規(guī)模語(yǔ)言模型LLM的強(qiáng)化學(xué)習(xí)訓(xùn)練過程中verlVolcano Engine Reinforcement人工智能大模型強(qiáng)化學(xué)習(xí)RLHF分布式訓(xùn)練微調(diào)LeRobot機(jī)器人學(xué)習(xí)3步構(gòu)建你的第一個(gè)AI機(jī)器人控制模型LeRobot機(jī)器人學(xué)習(xí)3步構(gòu)建你的第一個(gè)AI機(jī)器人控制模型 想不想讓機(jī)器人像人一樣學(xué)習(xí)新技能 你是否曾夢(mèng)想過讓機(jī)械臂學(xué)會(huì)抓取物體、讓機(jī)器人自主完成復(fù)雜人工智能機(jī)器學(xué)習(xí)深度學(xué)習(xí)機(jī)器人具身智能強(qiáng)化學(xué)習(xí)55項(xiàng)功能全面升級(jí)HsMod插件讓你的爐石傳說體驗(yàn)飛升8倍速55項(xiàng)功能全面升級(jí)HsMod插件讓你的爐石傳說體驗(yàn)飛升8倍速 HsMod是一款基于BepInEx框架開發(fā)的爐石傳說游戲增強(qiáng)插件為玩家提供了從游戲性能優(yōu)化到社游戲開發(fā)上一篇在VS Code中使用Windows Subsystem for Linux(WSL)進(jìn)行開發(fā)下一篇FlyEnv v4.9.7 版本更新優(yōu)化 FTP 服務(wù)與 PHP 安全配置創(chuàng)作聲明:本文部分內(nèi)容由AI輔助生成(AIGC),僅供參考