如何在mlr3的GraphLearner场景下用auto_tune()调优XGBoost的early-stopping?
解决mlr3中GraphLearner结合预处理超参数调优与Early Stopping的方案
方案1:自定义Learner封装预处理与模型
把Savitzky-Golay滤波和XGBoost打包成一个自定义Learner,这样就能直接使用mlr3的early stopping机制,同时支持超参数调优。
实现代码
library(mlr3) library(mlr3learners) library(signal) library(data.table) # 自定义二分类Learner,可扩展为多分类 LearnerClassifSGXGB <- R6::R6Class("LearnerClassifSGXGB", inherit = LearnerClassif, public = list( initialize = function() { super$initialize( id = "classif.sgxgb", # 定义所有可调超参数:SG滤波+XGBoost param_set = paradox::ps( sg_window = paradox::p_int(lower = 3, upper = 15, odd = TRUE, default = 5), sg_poly = paradox::p_int(lower = 1, upper = 3, default = 2), eta = paradox::p_dbl(lower = 0.01, upper = 0.3, default = 0.1), max_depth = paradox::p_int(lower = 3, upper = 10, default = 6), early_stopping_rounds = paradox::p_int(lower = 5, upper = 50, default = 10), nrounds = paradox::p_int(lower = 10, upper = 1000, default = 100) ), predict_types = c("response", "prob"), feature_types = c("numeric"), properties = c("twoclass", "multiclass", "weights") ) } ), private = list( .train = function(task) { # 拆分超参数:SG和XGBoost各自的参数 pars <- self$param_set$get_values(tags = "train") sg_pars <- pars[names(pars) %in% c("sg_window", "sg_poly")] xgb_pars <- pars[names(pars) %in% c("eta", "max_depth", "early_stopping_rounds", "nrounds")] # 对特征应用Savitzky-Golay滤波 feat_data <- task$data(cols = task$feature_names) filtered_feats <- apply(feat_data, 2, function(col) { signal::sgolayfilt(col, p = sg_pars$sg_poly, n = sg_pars$sg_window) }) # 创建滤波后的新任务 task_filtered <- task$clone() task_filtered$select(task$feature_names) task_filtered$cbind(data.table(filtered_feats)) # 拆分训练/验证集用于early stopping train_valid_split <- partition(task_filtered, ratio = 0.8) # 初始化并训练XGBoost,启用early stopping xgb_learner <- lrn("classif.xgboost", !!!xgb_pars) xgb_learner$train(task_filtered, row_ids = train_valid_split$train, early_stopping_set = train_valid_split$test) # 返回包含SG参数和XGB模型的对象 list( sg_params = sg_pars, xgb_model = xgb_learner$model ) }, .predict = function(task) { # 对测试特征应用相同参数的SG滤波 feat_data <- task$data(cols = task$feature_names) filtered_feats <- apply(feat_data, 2, function(col) { signal::sgolayfilt(col, p = self$model$sg_params$sg_poly, n = self$model$sg_params$sg_window) }) # 创建滤波后的测试任务 task_filtered <- task$clone()$cbind(data.table(filtered_feats)) # 用训练好的XGB模型预测 xgb_learner <- lrn("classif.xgboost") xgb_learner$model <- self$model$xgb_model xgb_learner$predict(task_filtered) } ) ) # 实例化自定义Learner sg_xgb_learner <- LearnerClassifSGXGB$new() # 构建调优流程(替换为你的任务) task <- tsk("your_spectral_task") resampling <- rsmp("cv", folds = 5) measure <- msr("classif.acc") # 定义超参数搜索空间 search_space <- paradox::ps( sg_window = paradox::p_int(lower = 3, upper = 11, odd = TRUE), sg_poly = paradox::p_int(lower = 1, upper = 3), eta = paradox::p_dbl(lower = 0.05, upper = 0.2), max_depth = paradox::p_int(lower = 4, upper = 8), early_stopping_rounds = paradox::p_int(lower = 10, upper = 30) ) # 配置调优器与终止条件 tuner <- tnr("grid_search", resolution = 5) terminator <- trm("evals", n_evals = 20) # 自动调优 auto_tuner <- AutoTuner$new( learner = sg_xgb_learner, resampling = resampling, measure = measure, search_space = search_space, tuner = tuner, terminator = terminator ) # 运行调优 auto_tuner$train(task)
优缺点
- 优势:完全兼容mlr3的AutoTuner和回调系统,early stopping逻辑直接复用XGBoost原生机制,稳定性高;超参数调优覆盖预处理和模型全流程。
- 劣势:需要编写R6类,对R6语法有一定要求。
方案2:嵌套Resampling手动实现Early Stopping
直接使用GraphLearner构建预处理+模型的 pipeline,通过嵌套交叉验证手动拆分验证集,实现early stopping逻辑。
实现代码
library(mlr3) library(mlr3pipelines) library(mlr3learners) library(signal) # 构建SG滤波+XGBoost的Graph sg_pipe <- po("colapply", applicator = function(x, sg_window, sg_poly) { signal::sgolayfilt(x, p = sg_poly, n = sg_window) }, param_vals = list(sg_window = 5, sg_poly = 2) ) xgb_learner <- lrn("classif.xgboost", nrounds = 1000) xgb_pipe <- po("learner", xgb_learner) graph <- sg_pipe %>>% xgb_pipe glrn <- GraphLearner$new(graph) # 设置可调超参数 glrn$param_set$values$colapply.sg_window <- to_tune(p_int(3, 15, odd = TRUE)) glrn$param_set$values$colapply.sg_poly <- to_tune(p_int(1, 3)) glrn$param_set$values$classif.xgboost.eta <- to_tune(p_dbl(0.01, 0.3)) glrn$param_set$values$classif.xgboost.max_depth <- to_tune(p_int(3, 10)) glrn$param_set$values$classif.xgboost.early_stopping_rounds <- to_tune(p_int(5, 50)) # 定义嵌套Resampling:外层CV评估,内层拆分验证集做early stopping outer_resampling <- rsmp("cv", folds = 5) inner_resampling <- rsmp("holdout", ratio = 0.8) # 自定义评估函数,手动处理early stopping custom_eval <- function(learner, task, resampling) { fold_scores <- c() for (fold in seq_len(resampling$iters)) { # 外层拆分训练/测试集 outer_split <- resampling$instantiate(task)$split(fold) outer_train_task <- task$clone()$filter(outer_split$train) # 内层拆分训练/验证集,用于early stopping inner_split <- inner_resampling$instantiate(outer_train_task)$split(1) learner$param_set$values$classif.xgboost.early_stopping_set <- inner_split$test # 训练模型 learner$train(outer_train_task, row_ids = inner_split$train) # 评估测试集性能 pred <- learner$predict(task$clone()$filter(outer_split$test)) fold_scores <- c(fold_scores, pred$score(msr("classif.acc"))) } mean(fold_scores) } # 构建调优实例 tuning_instance <- TuningInstanceSingleCrit$new( task = tsk("your_spectral_task"), learner = glrn, resampling = outer_resampling, measure = msr("classif.acc"), terminator = trm("evals", n_evals = 20) ) # 运行随机搜索调优 tnr("random_search")$optimize(tuning_instance)
优缺点
- 优势:无需自定义Learner,直接使用mlr3pipelines的组件,快速搭建原型;逻辑直观,容易理解。
- 劣势:手动管理训练循环和验证集传递,代码量较大;调优效率略低于自定义Learner方案。
方案3:给Graph中的XGBoost Learner单独绑定Early Stopping回调
如果不想写太多自定义代码,可以尝试将early stopping回调直接绑定到Graph中的XGBoost Learner上,结合嵌套Resampling传递验证集。
核心代码片段
# 初始化XGBoost Learner并绑定early stopping回调 xgb_learner <- lrn("classif.xgboost", nrounds = 1000) xgb_learner$callbacks <- clbk("early_stopping", early_stopping_rounds = 10, measure = msr("classif.acc") ) # 构建Graph sg_pipe <- po("colapply", applicator = function(x) signal::sgolayfilt(x, p=2, n=5)) graph <- sg_pipe %>>% po("learner", xgb_learner) glrn <- GraphLearner$new(graph) # 调优时通过嵌套Resampling传递验证集(逻辑同方案2)
注意事项
- 此方案需要确保XGBoost Learner能获取到验证集,因此必须结合嵌套Resampling,将内层验证集的行ID传递给
early_stopping_set参数。 - 部分场景下可能存在回调与GraphLearner的兼容性问题,需要测试验证。
内容的提问来源于stack exchange,提问作者franzi-r
相关产品推荐
相关产品推荐

