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

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

相关产品推荐
方舟 Agent Plan

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

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