ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

H2O GBM 训练中检查点 in_training_checkpoints_dir 参数详解:集群中断后的断点续训方案

H2O GBM 训练中检查点 in_training_checkpoints_dir 参数详解:集群中断后的断点续训方案 机器学习深度学习AutoML大数据后端【免费下载链接】h2o-3H2O is an Open Source, Distributed, Fast Scalable Machine Learning Platform: Deep Learning, Gradient Boosting (GBM) XGBoost, Random Forest, Generalized Linear Modeling (GLM with Elastic Net), K-Means, PCA, Generalized Additive Models (GAM), RuleFit, Support Vector Machine (SVM), Stacked Ensembles, Automatic Machine Learning (AutoML), etc.项目地址https://gitcode.com/gh_mirrors/h2/h2o-3点击查看免费下载本文围绕 H2O-3 中 GBM 模型的in_training_checkpoints_dir参数展开讲解如何在训练尚未结束时自动将“半成品”模型落盘到指定目录以及配合in_training_checkpoints_tree_interval控制落盘频率从而在集群宕机后手动从最近检查点恢复训练。读完本文你将掌握该参数的行为语义、源码级实现原理、R/Python 完整实战代码以及“检查点模型”与“完整训练模型”之间的预测一致性边界。参数定位适用算法与网格搜索属性in_training_checkpoints_dir是 H2O-3 中仅适用于 GBM的参数用于在训练进行过程中自动将尚未完成的模型写入一个指定目录In-Training Checkpoints训练中检查点。其关键属性见 关联文档 与 REST 层定义 SharedTreeV3.javaAvailable in: GBM由 SharedTree.java 中的默认实现throw new UnsupportedOperationException可推断基类默认不支持仅 GBM 覆写了该逻辑Hyperparameter: 否——它不属于网格搜索的超参数Schema 中标记为gridable false且属于 expert专家级参数通常无需调优配合参数:in_training_checkpoints_tree_interval两者必须搭配使用才有效。核心语义训练中检查点与“完整模型”的本质区别该选项的核心用途是在训练仍处于运行状态时自动把“未训练完的模型”落盘到指定目录。这样一旦集群意外关闭你可以借助这些检查点手动重启训练而不是从头再来。需要特别强调的是原文明确指出的行为语义检查点不是完整训练模型检查点只是某个中间时刻的快照。例如一个仅包含 4 棵树的检查点其预测结果可能与一个完整训练出来、恰好也包含 4 棵树的模型不同——因为树的分裂过程如采样、得分时机会受到后续训练过程的影响中间快照并不等价于“以 4 棵树为终点的完整训练”。一致性保证是有条件的一个没有任何训练中断的完整训练模型与“从检查点继续训练得到的模型”只要给定相同的数据和超参数最终结果始终一致。这一保证在源码测试中得到了直接验证见下文测试用例。参数与默认值速查参数类型默认值说明in_training_checkpoints_dirString空未启用训练中检查点的落盘目录路径必须指向可写目录in_training_checkpoints_tree_intervalint1每训练多少棵树保存一次检查点必须 0从参数定义源码 SharedTreeModel.java 可以看到in_training_checkpoints_tree_interval的默认值是1即默认每训练一棵树就落盘一次当你希望减小in_training_checkpoints_dir目录体积时可以调大该间隔例如设为2、3具体使用说明见 in_training_checkpoints_tree_interval.rst。in_training_checkpoints_tree_interval只有在in_training_checkpoints_dir已定义时才会生效见 SharedTreeV3.java 的说明与 SharedTree.java 的判断逻辑。源码级实现检查点如何写入与校验参数校验逻辑在 SharedTree.java 中训练开始前会校验间隔参数if (_parms._in_training_checkpoints_tree_interval 0) error(_in_training_checkpoints_tree_interval, _in_training_checkpoints_tree_interval must be 0.);即in_training_checkpoints_tree_interval必须大于 0否则直接报错。同时SharedTree.java若指定了检查点目录会通过H2O.getPM().isWritableDirectory(...)校验该路径必须是可写目录否则报错if (!StringUtils.isNullOrEmpty(_parms._in_training_checkpoints_dir)) { if (!H2O.getPM().isWritableDirectory(_parms._in_training_checkpoints_dir)) { error(_in_training_checkpoints_dir, In training checkpoints directory path must point to a writable path.); } }落盘时机与文件命名在每轮树构建的主循环中SharedTree.java通过取模判断是否到达检查点时机boolean manualCheckpointsInterval tid 0 tid % _parms._in_training_checkpoints_tree_interval 0; if (!StringUtils.isNullOrEmpty(_parms._in_training_checkpoints_dir) manualCheckpointsInterval) { doInTrainingCheckpoint(); }注意tid 0的条件第 0 棵树训练初始状态不会触发检查点落盘。真正执行写入的是 GBM 对doInTrainingCheckpoint()的覆写GBM.javaprotected void doInTrainingCheckpoint() { try { String modelFile _parms._in_training_checkpoints_dir / _model._key.toString() .ntrees_ _model._output._ntrees; GBMModel modelClone _model.clone(); modelClone.setInputParms(_parms); modelClone._key Key.make(_model._key . _model._output._ntrees); modelClone._output (GBMModel.GBMOutput) _model._output.clone(); modelClone._output.changeModelMetricsKey(modelClone._key); modelClone.exportBinaryModel(modelFile, true); } catch (IOException e) { throw new RuntimeException(Failed to write GBM checkpoint _model._key.toString(), e); } }由此可以明确检查点文件的命名约定model_id.ntrees_树数量例如model_id为gbm_model、当前已建 3 棵树时检查点文件名为gbm_model.ntrees_3文件以二进制模型格式exportBinaryModel导出可通过h2o.load_model/Model.importBinaryModel直接加载。当文件已存在时以true覆盖写入。测试用例佐证仓库测试 GBMCheckpointInTrainingTest.java 对上述行为做了系统验证可以作为理解该参数行为的权威参考testPartialCheckpointGivesTheSameResultAsTheFinalModel加载ntrees_3检查点并打分与 3 棵树的参考模型打分一致容差1e-3testPartialCheckpointAreProperlyExportedntrees4时目录中应存在 1、2、3 棵树的检查点文件testPartialCheckpointAreProperlyExported_definedIntervalinterval3、ntrees10时仅 3、6、9 棵树的检查点存在其余不存在testPartialCheckpointAreProperlyExported_restart从ntrees_3检查点重启训练_checkpoint指向该检查点新的检查点目录中只包含第 4 棵树之后的检查点testPartialCheckpointAreProperlyExported_restartWithDefinedInterval与..._restartAndChangeInterval验证重启后间隔参数沿用/变更时的文件生成行为testUsageOfPartialCheckpointGivesTheSameModelPrediction从ntrees_2检查点续训到 6 棵树最终模型预测与参考模型一致testCheckpointingDoesNotChangeModel开启检查点与不开启检查点训练出的模型打分一致证明落盘动作本身不改变模型训练结果。实战示例R 与 PythonR 示例原文给出的 R 完整示例使用 prostate 数据集ntrees10检查点目录为checkpointslibrary(h2o) h2o.init() # import the prostate dataset: prostate h2o.importFile(http://s3.amazonaws.com/h2o-public-test-data/smalldata/prostate/prostate.csv) # set the predictors, response, and categorical features: prostate$RACE - as.factor(prostate$RACE) prostate$CAPSULE - as.factor(prostate$CAPSULE) predictors - c(ID, AGE, RACE, DPROS, DCAPS, PSA, VOL, GLEASON) response - CAPSULE # specify directory for training checkpoints: checkpoints_dir - checkpoints # train the model and provide checkpoints in training process: pros_gbm - h2o.gbm(x predictors, y response, model_id gbm-model, ntrees 10, seed 1111, training_frame prostate, in_training_checkpoints_dir checkpoints_dir) # retrieve the number of files in the exported checkpoints directory: num_files - length(list.files(checkpoints_dir)) num_files # 9由于默认in_training_checkpoints_tree_interval 1每棵树落盘一次且第 0 棵树不落盘ntrees10最终得到9 个检查点文件对应树 1~9。Python 示例原文给出的 Python 完整示例# import necessary modules: import h2o from h2o.estimators.gbm import H2OGradientBoostingEstimator import tempfile from os import listdir, path # start h2o: h2o.init() # import the prostate dataset: prostate h2o.import_file(pathhttp://s3.amazonaws.com/h2o-public-test-data/smalldata/prostate/prostate.csv) # set the predictors, response, and categorical features: prostate[CAPSULE] prostate[CAPSULE].asfactor() prostate[RACE] prostate[RACE].asfactor() predictors [ID, AGE, RACE, DPROS, DCAPS, PSA, VOL, GLEASON] response CAPSULE # specify directory for training checkpoints: checkpoints_dir tempfile.mkdtemp() # train the model and export checkpoints in training process: pros_gbm H2OGradientBoostingEstimator(model_idgbm_model, ntrees10, seed1111, in_training_checkpoints_dircheckpoints_dir) pros_gbm.train(xpredictors, yresponse, training_frameprostate) # retrieve the number of files in the exported checkpoints directory: checkpoints listdir(checkpoints_dir) print(checkpoints) num_files len(listdir(checkpoints_dir)) print(num_files) # 9 # load checkpoint containing 3. trees: checkpoint h2o.load_model(path.join(checkpoints_dir, pros_gbm.model_id .ntrees_3)) display(Checkpoint:, checkpoint) # restart from checkpoint containing 3. trees: pros_gbm_restarted H2OGradientBoostingEstimator(model_idgbm_model, ntrees10, seed1111, checkpointcheckpoint, in_training_checkpoints_dircheckpoints_dir) pros_gbm_restarted.train(xpredictors, yresponse, training_frameprostate) pros_gbm_restarted # this model is equal to pros_gbm示例中还演示了两个关键操作加载中间检查点h2o.load_model(path.join(checkpoints_dir, pros_gbm.model_id .ntrees_3))按命名约定定位并加载 3 棵树的检查点从检查点重启训练把checkpointcheckpoint传给新的 estimator配合相同的数据、ntrees、seed等超参数重新训练最终模型与原来的pros_gbm等价——这与前述“完整训练模型与从检查点续训模型在相同数据/超参数下一致”的保证吻合也被 testUsageOfPartialCheckpointGivesTheSameModelPrediction 等测试锁定。控制检查点体积in_training_checkpoints_tree_interval 实战当训练树数较多如数千棵时默认每棵树都落盘会产生大量文件。通过in_training_checkpoints_tree_interval可降低落盘频率例如每 2 棵树保存一次见 in_training_checkpoints_tree_interval.rstpros_gbm H2OGradientBoostingEstimator(model_idgbm_model, ntrees10, seed1111, in_training_checkpoints_dircheckpoints_dir, in_training_checkpoints_tree_interval2) pros_gbm.train(xpredictors, yresponse, training_frameprostate) # 10 棵树、间隔 2、第 0 棵不落盘 → 共 4 个检查点树 2、4、6、8 num_files len(listdir(checkpoints_dir)) print(num_files) # 4对应 R 写法pros_gbm - h2o.gbm(x predictors, y response, model_id gbm-model, ntrees 10, seed 1111, training_frame prostate, in_training_checkpoints_dir checkpoints_dir, in_training_checkpoints_tree_interval 2) num_files - length(list.files(checkpoints_dir)) num_files # 4该行为与测试用例testPartialCheckpointAreProperlyExported_definedInterval完全一致interval3, ntrees10时仅存在 3、6、9 三个检查点。使用建议与注意事项指定可写目录目录路径必须在所有节点上可写校验见上文isWritableDirectory否则训练直接报错负载权衡默认每棵树落盘一次会产生大量文件生产环境建议结合in_training_checkpoints_tree_interval调大间隔平衡恢复粒度与磁盘占用恢复粒度间隔越大集群中断后可恢复的最小粒度越粗若中断发生在两个检查点之间只能恢复到最近的上一个检查点续训一致性前提从检查点重启时务必传入与原始训练相同的数据与超参数尤其seed、ntrees、各类采样率才能获得与无中断完整训练等价的模型即使中间快照本身不等价于同等树数的完整模型续训到相同总树数后结果一致参数适用范围该功能仅 GBM 支持且两个参数均不是网格搜索超参数属于专家级开关普通调参流程中无需主动设置。小结in_training_checkpoints_dir配合in_training_checkpoints_tree_interval为 H2O GBM 提供了面向长时训练与不稳定集群环境的训练中断恢复机制训练期间按树间隔将模型二进制快照写入指定目录集群宕机后可从最近检查点手动重启。其“中间快照不等价于完整模型、续训到相同终点则结果一致”的语义已在源码与测试中得到严格保证是构建高可靠训练流水线时值得掌握的一组专家级参数。赞分享机器学习深度学习AutoML大数据后端【免费下载链接】h2o-3H2O is an Open Source, Distributed, Fast Scalable Machine Learning Platform: Deep Learning, Gradient Boosting (GBM) XGBoost, Random Forest, Generalized Linear Modeling (GLM with Elastic Net), K-Means, PCA, Generalized Additive Models (GAM), RuleFit, Support Vector Machine (SVM), Stacked Ensembles, Automatic Machine Learning (AutoML), etc.项目地址https://gitcode.com/gh_mirrors/h2/h2o-3点击查看免费下载相关推荐H2O-3 GBM 训练中检查点间隔控制in_training_checkpoints_tree_interval 参数完全指南H2O 3 GBM 训练中检查点间隔控制in_training_checkpoints_tree_interval 参数完全指南 H2O 3 的 in_tra机器学习深度学习AutoML大数据后端终极指南Kohya_SS训练中断后从检查点恢复的完整方法终极指南Kohya_SS训练中断后从检查点恢复的完整方法 在AI模型训练过程中训练中断是每个用户都可能遇到的问题。Kohya_SS作为当前最流行的Stabl人工智能微调LoRA深度学习计算机视觉AI 应用训练中断不用慌OpenRLHF断点续训全攻略训练中断不用慌OpenRLHF断点续训全攻略 在大模型训练过程中意外断电、显存溢出或网络中断等问题时常导致训练被迫中止。重新开始不仅浪费算力更可能错过最佳人工智能大模型强化学习RLHF分布式训练微调上一篇Pixelle-Video5分钟掌握AI短视频生成的终极完整指南下一篇百度网盘秒传链接工具5分钟掌握高效文件管理方法创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表