如何正确保存mlr3框架下的lightgbm模型以解决加载后预测报错
问题根因
LightGBM的R包底层依赖C++实现的模型对象,对外暴露的是外部指针结构,R默认的saveRDS序列化机制无法保留外部指针的有效内存地址,序列化再读取后指针失效,就会触发预测报错,该问题和mlr3的封装逻辑无关,是LightGBM R接口本身的特性导致的。
正确保存方案
有两种常用的实现方式,都需要配合LightGBM官方提供的序列化接口使用:
方案1:分开保存管线结构与模型文件
保存步骤
library(mlr3) library(mlr3pipelines) library(mlr3extralearners) library(lightgbm) # 原有训练逻辑不变 data = tsk("german_credit")$data() data = data[, c("credit_risk", "amount", "purpose", "age")] task = TaskClassif$new("boston", backend = data, target = "credit_risk") g = po("imputemedian") %>>% po("imputeoor") %>>% po("fixfactors") %>>% po("encodeimpact") %>>% lrn("classif.lightgbm") gl = GraphLearner$new(g) gl$train(task) # 测试原地预测正常 newdata <- data[1,] gl$predict_newdata(newdata) # 1. 单独提取LightGBM原始模型,用官方接口保存 lgb_model = gl$graph$pipeops$classif.lightgbm$learner$model lgb.save(lgb_model, "lgb_model.model") # 2. 保存GraphLearner的完整管线结构 saveRDS(gl, "gl_pipeline.rds")
加载预测步骤
library(mlr3) library(mlr3pipelines) library(mlr3extralearners) library(lightgbm) # 分别读取管线和模型 gl = readRDS("gl_pipeline.rds") lgb_model = lgb.load("lgb_model.model") # 替换管线中的失效模型为新加载的有效模型 gl$graph$pipeops$classif.lightgbm$learner$model = lgb_model # 预测可正常运行 newdata <- data[1,] gl$predict_newdata(newdata)
方案2:单文件保存(序列化模型为raw向量)
如果不想维护两个独立文件,可以将LightGBM模型序列化为R原生的raw向量,和管线存在同一个RDS中:
保存步骤
# 训练逻辑和之前一致,训练完成后执行: lgb_raw = lgb.serialize(gl$graph$pipeops$classif.lightgbm$learner$model) # 打包为列表保存 saveRDS(list(pipeline = gl, lgb_raw = lgb_raw), "gl_full.rds")
加载预测步骤
res = readRDS("gl_full.rds") gl = res$pipeline # 反序列化raw向量为可用的LightGBM模型,替换回管线 gl$graph$pipeops$classif.lightgbm$learner$model = lgb.unserialize(res$lgb_raw) # 正常预测 newdata <- data[1,] gl$predict_newdata(newdata)
注意事项
如果自定义了LightGBM学习器的id参数,替换模型时需要将代码中的classif.lightgbm修改为你自己设置的id值即可。
内容的提问来源于stack exchange,提问作者BinhNN
相关产品推荐
相关产品推荐

