- 机器学习
- 深度学习
- AutoML
- 大数据
- 后端
【免费下载链接】h2o-3
H2O 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.
本文围绕 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.java):
- Available in: GBM(由 SharedTree.java 中的默认实现
throw new UnsupportedOperationException可推断,基类默认不支持,仅 GBM 覆写了该逻辑); - Hyperparameter: 否——它不属于网格搜索的超参数(Schema 中标记为
gridable = false),且属于 expert(专家级)参数,通常无需调优; - 配合参数:
in_training_checkpoints_tree_interval,两者必须搭配使用才有效。
核心语义:训练中检查点与“完整模型”的本质区别
该选项的核心用途是:在训练仍处于运行状态时,自动把“未训练完的模型”落盘到指定目录。这样一旦集群意外关闭,你可以借助这些检查点手动重启训练,而不是从头再来。
需要特别强调的是(原文明确指出的行为语义):
- 检查点不是完整训练模型:检查点只是某个中间时刻的快照。例如一个仅包含 4 棵树的检查点,其预测结果可能与一个完整训练出来、恰好也包含 4 棵树的模型不同——因为树的分裂过程(如采样、得分时机)会受到后续训练过程的影响,中间快照并不等价于“以 4 棵树为终点的完整训练”。
- 一致性保证是有条件的:一个没有任何训练中断的完整训练模型,与“从检查点继续训练得到的模型”,只要给定相同的数据和超参数,最终结果始终一致。这一保证在源码测试中得到了直接验证(见下文测试用例)。
参数与默认值速查
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
in_training_checkpoints_dir | String | 空(未启用) | 训练中检查点的落盘目录路径;必须指向可写目录 |
in_training_checkpoints_tree_interval | int | 1 | 每训练多少棵树保存一次检查点;必须> 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.java):
protected 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-3);testPartialCheckpointAreProperlyExported:ntrees=4时,目录中应存在 1、2、3 棵树的检查点文件;testPartialCheckpointAreProperlyExported_definedInterval:interval=3、ntrees=10时,仅 3、6、9 棵树的检查点存在,其余不存在;testPartialCheckpointAreProperlyExported_restart:从ntrees_3检查点重启训练(_checkpoint指向该检查点),新的检查点目录中只包含第 4 棵树之后的检查点;testPartialCheckpointAreProperlyExported_restartWithDefinedInterval与..._restartAndChangeInterval:验证重启后间隔参数沿用/变更时的文件生成行为;testUsageOfPartialCheckpointGivesTheSameModelPrediction:从ntrees_2检查点续训到 6 棵树,最终模型预测与参考模型一致;testCheckpointingDoesNotChangeModel:开启检查点与不开启检查点训练出的模型打分一致,证明落盘动作本身不改变模型训练结果。
实战示例:R 与 Python
R 示例
原文给出的 R 完整示例(使用 prostate 数据集,ntrees=10,检查点目录为checkpoints):
library(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 棵树不落盘,ntrees=10最终得到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(path="http://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_id="gbm_model", ntrees=10, seed=1111, in_training_checkpoints_dir=checkpoints_dir) pros_gbm.train(x=predictors, y=response, training_frame=prostate) # 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_id="gbm_model", ntrees=10, seed=1111, checkpoint=checkpoint, in_training_checkpoints_dir=checkpoints_dir) pros_gbm_restarted.train(x=predictors, y=response, training_frame=prostate) pros_gbm_restarted # this model is equal to pros_gbm示例中还演示了两个关键操作:
- 加载中间检查点:
h2o.load_model(path.join(checkpoints_dir, pros_gbm.model_id + ".ntrees_3"))按命名约定定位并加载 3 棵树的检查点; - 从检查点重启训练:把
checkpoint=checkpoint传给新的 estimator,配合相同的数据、ntrees、seed等超参数重新训练,最终模型与原来的pros_gbm等价——这与前述“完整训练模型与从检查点续训模型在相同数据/超参数下一致”的保证吻合,也被 testUsageOfPartialCheckpointGivesTheSameModelPrediction 等测试锁定。
控制检查点体积:in_training_checkpoints_tree_interval 实战
当训练树数较多(如数千棵)时,默认每棵树都落盘会产生大量文件。通过in_training_checkpoints_tree_interval可降低落盘频率,例如每 2 棵树保存一次(见 in_training_checkpoints_tree_interval.rst):
pros_gbm = H2OGradientBoostingEstimator(model_id="gbm_model", ntrees=10, seed=1111, in_training_checkpoints_dir=checkpoints_dir, in_training_checkpoints_tree_interval=2) pros_gbm.train(x=predictors, y=response, training_frame=prostate) # 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完全一致(interval=3, ntrees=10时仅存在 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-3
H2O 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.
相关推荐
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分布式训练微调
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考