mlr3中GraphLearner训练后为何仍显示<<UNTRAINED>>状态?
问题原因及解决方法
核心原因是GraphLearner训练时会使用原始graph的副本,而非直接修改你传入的原始graph,所以你直接访问learner$graph看到的还是未训练的原始管道,而训练后的状态存在learner$model里。
具体细节:
- mlr3的GraphLearner设计逻辑是为了隔离训练状态和原始管道定义:当你调用
train()时,它会复制一份传入的graph进行训练,原始graph保持未训练状态,这样你可以后续复用原始管道做修改、重新训练,不会被已训练的状态污染。 - 要查看训练完成的管道状态,不要用
learner$graph,而是访问learner$model,这个属性里存储的是训练后的graph实例,里面的所有PipeOps都会显示<>。
举个实际代码示例验证:
# 构建包含one-hot编码的管道 graph = po("encode", method = "one-hot") %>>% po("classif.rpart") # 封装为GraphLearner gl = GraphLearner$new(graph) # 执行训练 gl$train(tsk("iris")) # 查看原始graph(状态为UNTRAINED) print(gl$graph) # 查看训练后的graph(状态为TRAINED) print(gl$model)
如果非要让原始graph直接被训练,可以跳过GraphLearner,直接调用graph的train()方法,但这种方式会修改原始graph的状态,后续无法再用它构建新的未训练管道,一般不推荐。
内容的提问来源于stack exchange,提问作者atg
相关产品推荐
相关产品推荐

