使用mlr3包基于多公式调优GAM时的ParamUty错误解决方法
如何通过测试多个公式调优广义可加模型(GAM)
问题背景
我尝试用网格搜索,基于不同k值(平滑项的基维度)对应的多个公式调优广义可加模型(GAM),但运行代码时触发错误:
Error: 'TunerGridSearch' does not support param types: 'ParamUty'
以下是可复现代码:
## Install the learner "classif.gam" remotes::install_github("mlr-org/mlr3extralearners@*release") install_learners("classif.gam") ## Task task_sonar = tsk("sonar") ## Search space(原错误写法) search_space <- paradox::ps(formula = paradox::p_uty(c("Class ~ s(V15, k = 1)", "Class ~ s(V15, k = 2)"))) ## Learner learner <- mlr3extralearners::lrn("classif.gam", predict_type = "prob") ## Performance measure measure <- mlr3::msr("classif.auc") ## Terminator terminator <- mlr3tuning::trm("none") ## Tuner tuner <- mlr3tuning::tnr("grid_search") ## Resampling resampling <- rsmp ("cv", folds = 5) ## Run an automatic tuning process at = mlr3tuning::auto_tuner(tuner = tuner, learner = learner, resampling = resampling, measure = measure, search_space = search_space, terminator = terminator) at$train(task_sonar)
问题原因
网格搜索不支持ParamUty(任意类型参数),因为它无法自动生成参数网格。直接把完整公式作为p_uty参数传入,不符合网格搜索对参数类型的要求。
解决方法
不需要直接传递完整公式,而是把k值作为独立的可调参数,再动态生成对应的GAM公式。具体实现如下:
修改后的完整代码
## 安装依赖包 remotes::install_github("mlr-org/mlr3extralearners@*release") install_learners("classif.gam") ## 加载所需库 library(mlr3) library(mlr3tuning) library(paradox) library(mlr3extralearners) ## 定义任务 task_sonar <- tsk("sonar") ### 1. 定义k值的搜索空间(整数类型,符合网格搜索要求) search_space <- ps(k = p_int(lower = 1, upper = 2)) # 可根据需求扩展k的范围,比如1到5 ### 2. 自定义learner,动态生成公式 learner <- lrn("classif.gam", predict_type = "prob") # 克隆learner并添加k参数 learner <- learner$clone() learner$param_set$add(ParamInt$new("k", lower = 1, upper = 10)) # 重写训练方法,根据k值动态设置公式 learner$train_internal <- function(task) { formula_str <- sprintf("Class ~ s(V15, k = %d)", self$param_set$values$k) self$param_set$values$formula <- as.formula(formula_str) super$train_internal(task) } ### 3. 设置调优组件并运行 measure <- msr("classif.auc") terminator <- trm("none") tuner <- tnr("grid_search") resampling <- rsmp("cv", folds = 5) at <- auto_tuner( tuner = tuner, learner = learner, resampling = resampling, measure = measure, search_space = search_space, terminator = terminator ) at$train(task_sonar) # 查看调优结果 print(at$tuning_result)
代码说明
- 把k定义为整数参数,网格搜索可以正常遍历所有指定的k值
- 通过重写learner的
train_internal方法,在每次训练前根据当前k值动态生成GAM公式 - 既满足了网格搜索对参数类型的要求,又实现了测试不同k值对应公式的需求
内容的提问来源于stack exchange,提问作者Pierre
相关产品推荐
相关产品推荐

