R语言中lightGBM模型保存加载后预测失败问题求助
问题现象
使用saveRDS()保存lightGBM模型,readRDS()加载后执行预测时触发错误:
Error in predictor$predict(data = data, start_iteration = start_iteration, : Attempting to use a Booster which no longer exists. This can happen if you have called Booster$finalize() or if this Booster was saved with saveRDS(). To avoid this error in the future, use saveRDS.lgb.Booster() or Booster$save_model() to save lightgbm Boosters.
训练完成后直接使用模型,或在当前会话内加载模型均可正常预测,但在其他应用会话中加载后报错;lasso、xgboost等其他模型无此问题。用户代码如下:
library(tidymodels) library(lightgbm) library(bonsai) TMA_model <- data.frame( age = c(50, 45, 60, 55, 70, 34, 55, 48, 58, 42, 52, 47, 62, 57, 72, 36, 53, 49, 59, 44, 51, 46, 61, 56, 71, 35, 54, 50, 60, 43), delta_creat = c(1.2, 1.0, 1.5, 1.3, 1.8, 1.2, 1.3, 1.4, 1.1, 1.6, 1.2, 1.0, 1.5, 1.3, 1.8, 1.2, 1.3, 1.4, 1.1, 1.6, 1.2, 1.0, 1.5, 1.3, 1.8, 1.2, 1.3, 1.4, 1.1, 1.6), max_LDH = c(300, 280, 320, 310, 330, 295, 325, 290, 315, 305, 305, 315, 290, 320, 300, 325, 280, 330, 310, 295, 310, 295, 330, 280, 320, 305, 310, 325, 315, 290), min_plat = c(150, 140, 160, 155, 170, 145, 165, 150, 135, 160, 160, 155, 170, 145, 150, 160, 140, 170, 155, 145, 155, 145, 170, 140, 160, 150, 155, 160, 165, 135), min_hb = c(12, 11.5, 13, 12.5, 14, 11.8, 13.2, 12.2, 12.6, 11.9, 12.4, 11.7, 13.1, 12.3, 13.5, 11.6, 12.8, 12.1, 13.4, 12.0, 12.9, 11.6, 13.2, 12.8, 13.0, 11.5, 12.7, 12.3, 12.5, 11.8), max_ast = c(40, 38, 42, 41, 45, 37, 43, 39, 44, 40, 42, 38, 44, 41, 40, 39, 43, 38, 45, 37, 44, 37, 42, 38, 41, 40, 43, 39, 44, 38), max_bt = c(37, 37.2, 37.5, 37.3, 37.8, 37.1, 37.4, 37.6, 37.0, 37.7, 37.2, 37.1, 37.5, 37.4, 37.3, 37.0, 37.6, 37.8, 37.2, 37.4, 37.7, 37.1, 37.3, 37.8, 37.0, 37.4, 37.2, 37.6, 37.5, 37.1), max_ttap = c(100, 95, 105, 102, 110, 97, 108, 93, 103, 98, 105, 100, 95, 110, 102, 108, 97, 93, 103, 95, 100, 97, 105, 93, 108, 102, 110, 98, 103, 100), max_tp = c(12, 12.2, 12.5, 12.3, 12.8, 12.1, 12.4, 12.6, 12.0, 12.7, 12.3, 12.2, 12.1, 12.8, 12.4, 12.0, 12.6, 12.7, 12.5, 12.3, 12.1, 12.4, 12.7, 12.8, 12.0, 12.2, 12.3, 12.6, 12.5, 12.4), hypertension = c(1, 0, 1, 1, 0, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 0, 1, 1, 0, 0, 1, 0, 1, 0, 0, 1, 1, 0, 1), MAP = c(90, 92, 88, 91, 89, 93, 87, 94, 86, 90, 92, 94, 88, 90, 91, 89, 87, 93, 86, 92, 91, 90, 93, 94, 86, 88, 89, 87, 92, 91), TMA_class = factor(c(1, 2, 3, 4, 5, 1, 2, 3, 4, 5, 1, 2, 3, 4, 5, 1, 2, 3, 4, 5, 1, 2, 3, 4, 5, 1, 2, 3, 4, 5)) ) new_data <- data.frame( age = c(48, 55, 62, 54, 67), delta_creat = c(1.2, 1.3, 1.1, 1.4, 1.6), max_LDH = c(305, 290, 320, 310, 295), min_plat = c(150, 160, 135, 155, 145), min_hb = c(12.2, 12.6, 11.9, 12.4, 11.7), max_ast = c(38, 43, 39, 44, 40), max_bt = c(37.2, 37.4, 37.0, 37.6, 37.1), max_ttap = c(95, 108, 93, 103, 98), max_tp = c(12.2, 12.4, 12.0, 12.6, 12.7), hypertension = c(1, 0, 1, 0, 1), MAP = c(92, 87, 94, 86, 90) ) model_spec <- boost_tree(mode = 'classification', mtry = 6, trees = 50, min_n = 50, tree_depth = 5) %>% set_engine('lightgbm') recipe <- recipe(TMA_class ~ ., data = TMA_model) model_fit <- workflow() %>% add_recipe(recipe) %>% add_model(model_spec) %>% fit(data = TMA_model) saveRDS(model_fit, "model_fit.rds") model_b <- readRDS("model_fit.rds") new_data_predictions <- predict(model_b, new_data) print(new_data_predictions)
问题原因
lightGBM的Booster对象包含底层C++实现的指针,标准saveRDS()无法正确序列化这些底层对象,导致加载后丢失关键的模型执行信息,进而触发预测错误。
解决方案
方案1:使用lightGBM专用方法保存/加载Booster对象
从tidymodels workflow中提取底层Booster模型,用lightGBM提供的专用方法处理:
# 训练完成后,提取底层Booster对象 booster_obj <- extract_fit_parsnip(model_fit)$fit # 用lightGBM专用方法保存 saveRDS.lgb.Booster(booster_obj, "lightgbm_booster.rds") # 加载模型 loaded_booster <- readRDS("lightgbm_booster.rds") # 将Booster重新整合回parsnip模型,适配workflow使用 loaded_parsnip <- parsnip::new_model_fit( fit = loaded_booster, spec = model_spec, preproc = extract_preprocessor(model_fit) ) # 构建新的workflow并预测 loaded_workflow <- workflow() %>% add_recipe(recipe) %>% add_model(loaded_parsnip) new_data_predictions <- predict(loaded_workflow, new_data)
方案2:用文本格式保存模型(跨会话更稳定)
将Booster保存为文本格式,再加载:
# 提取Booster并保存为文本 booster_obj <- extract_fit_parsnip(model_fit)$fit booster_obj$save_model("lightgbm_model.txt") # 加载文本模型 loaded_booster <- lgb.load("lightgbm_model.txt") # 整合回workflow使用(步骤同方案1)
方案3:用butcher包优化模型序列化(推荐)
butcher包可以清理模型中的冗余信息,同时正确处理lightGBM的外部指针,直接保存整个workflow:
library(butcher) # 清理模型并保存 model_fit_clean <- butcher(model_fit) saveRDS(model_fit_clean, "model_fit_clean.rds") # 加载后直接预测 model_b <- readRDS("model_fit_clean.rds") new_data_predictions <- predict(model_b, new_data)
内容的提问来源于stack exchange,提问作者Abner Mácola

