You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.13 00:07:05