mlr3中用auto_tune与benchmark构建自定义超参数空间及结果获取问题
在mlr3中构建带参数组合约束的超参数空间并访问调优结果
一、优雅实现特定参数组合的超参数空间
方法1:使用条件约束(Conditional Constraints)
如果参数间存在依赖关系(比如选某参数值时只能搭配另一参数的特定取值),直接用ParamSet的$add_conditional_constraint定义规则:
library(mlr3) library(mlr3tuning) # 定义目标学习器 learner = lrn("classif.ranger") # 创建基础参数空间 ps = ps( num.trees = p_int(lower = 100, upper = 500), mtry = p_int(lower = 2, upper = 5), splitrule = p_fct(c("gini", "extratrees")) ) # 添加约束:当splitrule为"extratrees"时,mtry只能取2或3 ps$add_conditional_constraint( condition = quote(splitrule == "extratrees"), param = "mtry", rhs = quote(mtry %in% c(2,3)) )
方法2:预定义固定参数组合(适用于离散场景)
如果是完全固定的几组参数组合,可将组合打包为因子参数,再通过trafo映射为实际参数,彻底避免无效组合:
# 预定义所有允许的参数组合 allowed_combinations = list( combo1 = list(num.trees = 100, mtry = 2, splitrule = "gini"), combo2 = list(num.trees = 300, mtry = 3, splitrule = "extratrees"), combo3 = list(num.trees = 500, mtry = 4, splitrule = "gini") ) # 创建仅包含组合选择的参数空间 ps = ps( combo = p_fct(names(allowed_combinations)) ) # 添加转换函数,将组合映射为实际参数 ps$trafo = function(x, param_set) { allowed_combinations[[x$combo]] }
方法3:自定义参数筛选(适用于复杂规则)
如果约束逻辑过于复杂,可在调优实例中传入filter_fun过滤无效参数组合:
# 定义筛选函数:仅保留满足num.trees + mtry < 510的组合 filter_fun = function(x) { x$num.trees + x$mtry < 510 } # 创建调优实例时绑定筛选函数 instance = ti( task = tsk("iris"), learner = learner, resampling = rsmp("cv", folds = 3), measure = msr("classif.acc"), search_space = ps, filter_fun = filter_fun )
二、访问调优后的超参数
调优完成后,可通过以下几种方式获取最优参数:
1. 从调优实例结果提取
# 运行调优 tuner = tnr("grid_search") tuner$optimize(instance) # 获取最优参数值 best_params = instance$result$learner_param_vals print(best_params)
2. 从自动调优学习器提取
如果使用auto_tune包装学习器,训练后可直接从包装器中获取:
# 创建自动调优学习器 at_learner = auto_tune( learner = learner, resampling = rsmp("holdout"), measure = msr("classif.acc"), search_space = ps, tuner = tnr("random_search"), term_evals = 20 ) # 训练后提取最优参数 at_learner$train(tsk("iris")) best_params = at_learner$learner$param_set$values print(best_params)
3. 从基准测试结果提取
在benchmark场景中,从BenchmarkResult聚合结果中提取:
# 构建基准测试设计 design = benchmark_grid( tasks = tsk("iris"), learners = list(at_learner), resamplings = rsmp("cv", folds = 3) ) # 运行基准测试 bmr = benchmark(design) # 提取最优参数 best_params = bmr$aggregate()[, .(learner_param_vals)] print(best_params)
三、注意事项
- 使用条件约束时,需确保约束逻辑与学习器参数兼容,避免因无效参数导致训练失败
- 预定义组合的方式适合参数组合较少的场景,可大幅压缩调优搜索空间
learner_param_vals仅包含调优过的参数,而param_set$values包含学习器所有参数(含默认值),注意区分使用
内容的提问来源于stack exchange,提问作者franzi-r
相关产品推荐
相关产品推荐

