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

使用mlr3中Hyperband调参cv_glmnet时遇向量过长错误

Hyperband调优surv.cv_glmnet/classif.cv_glmnet时触发向量过长错误

问题详情

调优surv.cv_glmnet或classif.cv_glmnet超参数时,出现如下错误:

Error in 0:s_max : result would be too long a vector

复现代码

library(mlr3)
library(mlr3proba)
library(mlr3hyperband)
library(mlr3extralearners)
library(mlr3learners)

instance = ti(
  # task = tsk("rats"),
  # learner = lrn("surv.cv_glmnet",
  #               id = "class_semipar_cvglmnet",
  #               alpha = to_tune(p_dbl(0, 1, tags = "budget"))),
  task = tsk("sonar"),
  learner = lrn("classif.cv_glmnet",
                id = "class_semipar_cvglmnet",
                alpha = to_tune(p_dbl(0, 1, tags = "budget"))),
  resampling = rsmp("holdout"),
  #measure = msr("surv.cindex"),
  measures = msr("classif.ce"),
  terminator = trm("none")
)

tuner = tnr("hyperband", eta = 2, repetitions = 1) #same error with eta=5,10 
#hyperband_schedule(r_min = 0.01, r_max = 1, eta = 2)
tuner$optimize(instance)

环境信息

> sessionInfo()
R version 4.3.1 (2023-06-16 ucrt)
Platform: x86_64-w64-mingw32/x64 (64-bit)
Running under: Windows 10 x64 (build 19045)

Matrix products: default


locale:
[1] LC_COLLATE=English_United States.utf8  LC_CTYPE=English_United States.utf8   
[3] LC_MONETARY=English_United States.utf8 LC_NUMERIC=C                           
[5] LC_TIME=English_United States.utf8    

time zone: Asia/Dubai
tzcode source: internal

attached base packages:
[1] stats     graphics  grDevices utils     datasets  methods   base      

other attached packages:
[1] mlr3proba_0.5.2              mlr3learners_0.7.0           mlr3extralearners_0.7.0-9000
[4] mlr3hyperband_0.6.0          mlr3tuning_1.0.0             paradox_1.0.1               
[7] mlr3_0.20.2                  

loaded via a namespace (and not attached):
 [1] utf8_1.2.3           future_1.33.0        generics_0.1.3       distr6_1.8.0         lattice_0.21-8      
 [6] listenv_0.9.0        digest_0.6.33        magrittr_2.0.3       grid_4.3.1           ooplah_0.2.0        
[11] xgboost_1.7.8.1      jsonlite_1.8.7       Matrix_1.6-1.1       backports_1.4.1      survival_3.5-7      
[16] param6_0.2.4         fansi_1.0.4          scales_1.3.0         codetools_0.2-19     mlr3measures_0.6.0  
[21] palmerpenguins_0.1.1 cli_3.6.1            rlang_1.1.1          crayon_1.5.2         parallelly_1.36.0   
[26] mlr3viz_0.9.0        future.apply_1.11.0  splines_4.3.1        munsell_0.5.0        withr_2.5.1         
[31] mlr3pipelines_0.6.0  tools_4.3.1          parallel_4.3.1       uuid_1.1-1           set6_0.2.6          
[36] checkmate_2.2.0      dplyr_1.1.3          colorspace_2.1-0     ggplot2_3.5.0        globals_0.16.2      
[41] bbotk_1.0.1          vctrs_0.6.3          R6_2.5.1             lifecycle_1.0.3      dictionar6_0.1.3    
[46] mlr3misc_0.15.1      pkgconfig_2.0.3      pillar_1.9.0         gtable_0.3.4         data.table_1.15.4   
[51] glue_1.6.2           Rcpp_1.0.11          lgr_0.4.4            tibble_3.2.1         tidyselect_1.2.0    
[56] rstudioapi_0.15.0    compiler_4.3.1 

解决方案

错误原因是将连续型超参数alpha标记为budget参数,但Hyperband要求budget参数必须是离散的、代表训练资源量级的参数(比如迭代次数、训练样本占比等),alpha作为弹性网的混合系数,属于普通超参数,不能充当budget角色。

修正步骤:

  • 移除alpha的tags = "budget"标记,将其作为普通超参数调优
  • 为cv_glmnet系列 learner 指定合法的budget参数,比如nlambda(lambda序列长度,直接控制训练资源消耗,且为离散值)

修正后的代码:

library(mlr3)
library(mlr3proba)
library(mlr3hyperband)
library(mlr3extralearners)
library(mlr3learners)

instance = ti(
  task = tsk("sonar"),
  learner = lrn("classif.cv_glmnet",
                id = "class_semipar_cvglmnet",
                alpha = to_tune(p_dbl(0, 1)),  # 移除budget标签
                nlambda = to_tune(p_int(10, 100, tags = "budget"))  # 用nlambda作为budget参数
  ),
  resampling = rsmp("holdout"),
  measures = msr("classif.ce"),
  terminator = trm("evals", n_evals = 100)  # 替换无限制终止条件
)

tuner = tnr("hyperband", eta = 2, repetitions = 1)
tuner$optimize(instance)

内容的提问来源于stack exchange,提问作者A. Suliman

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 08:50:53