【问题标题】:How to use different engine parameters for each fold with tidymodels during crossvalidation?在交叉验证期间,如何使用 tidymodels 为每个折叠使用不同的引擎参数?
【发布时间】:2021-06-25 13:39:04
【问题描述】:

我想使用 tidymodels 在交叉验证设置中调整游侠随机森林。 我的数据集不平衡。因此,我想使用游侠参数 class.weights。

但是,每个折叠可以有不同的权重。 如何将折叠特定重量传递给引擎?

MWE:

library(tidyverse)
library(tidymodels)
set.seed(111)

iris_cut <- iris[30:110,] # Dummy unbalanced dataset

# Create folds  
folds <- vfold_cv(iris_cut, v = 3)

# Calculate class weights: 
calc_weights <- function(df) {
  weights <- df %>% 
    group_by(Species) %>% 
    mutate(n_total = n()) %>% 
    ungroup() %>% 
    mutate(weight = max(n_total)/n_total) %>% 
    distinct(n_total, .keep_all = TRUE) %>% 
    as.data.frame() %>% 
    .$weight 
  return(weights)
}

# Use during training of fold 1:
weights_fold1 <- folds$splits[[1]]$data[folds$splits[[1]]$in_id,] %>% calc_weights() 
# Use during training of fold 2:
weights_fold2 <- folds$splits[[2]]$data[folds$splits[[2]]$in_id,] %>% calc_weights()
# Use during training of fold 3:
weights_fold3 <- folds$splits[[3]]$data[folds$splits[[3]]$in_id,] %>% calc_weights()


# Defining a recipe
rec <- recipe(Species~ ., data = iris_cut) 
  
# Create Model Specification
rf_mod <- rand_forest(
  mtry = tune(), 
  trees = 1000, 
  min_n = 1
) %>% 
  set_mode("classification") %>% 
  set_engine("ranger",
             class.weights=!!weights_fold1 # Wrong! Here for each fold another weight vector should be passed
             )  

rf_grid <- crossing(
  mtry = c(1,2,3)
)
  
# Setup workflow 
tune_wf <- workflow() %>% 
  add_recipe(rec) %>% 
  add_model(rf_mod)

# Start rf tuning
tune_res <- tune_grid( 
  tune_wf,
  resamples = folds,
  grid = rf_grid
)

【问题讨论】:

    标签: r tidymodels


    【解决方案1】:

    tidymodels 生态系统不支持将这些权重传递给折叠。相反,我们鼓励人们使用themis 包来处理特征工程期间的类不平衡。上采样和下采样有多种选择。

    例如SMOTE算法:

    library(tidyverse)
    library(tidymodels)
    #> Registered S3 method overwritten by 'tune':
    #>   method                   from   
    #>   required_pkgs.model_spec parsnip
    library(themis)
    #> Registered S3 methods overwritten by 'themis':
    #>   method                  from   
    #>   bake.step_downsample    recipes
    #>   bake.step_upsample      recipes
    #>   prep.step_downsample    recipes
    #>   prep.step_upsample      recipes
    #>   tidy.step_downsample    recipes
    #>   tidy.step_upsample      recipes
    #>   tunable.step_downsample recipes
    #>   tunable.step_upsample   recipes
    #> 
    #> Attaching package: 'themis'
    #> The following objects are masked from 'package:recipes':
    #> 
    #>     step_downsample, step_upsample
    set.seed(111)
    
    iris_cut <- iris[30:110,] # Dummy unbalanced dataset
    
    iris_rec <- recipe(Species ~ ., data = iris_cut) %>%
      step_smote(Species)
      
    iris_rec %>% prep() %>% bake(new_data = NULL)
    #> # A tibble: 150 x 5
    #>    Sepal.Length Sepal.Width Petal.Length Petal.Width Species
    #>           <dbl>       <dbl>        <dbl>       <dbl> <fct>  
    #>  1          4.7         3.2          1.6         0.2 setosa 
    #>  2          4.8         3.1          1.6         0.2 setosa 
    #>  3          5.4         3.4          1.5         0.4 setosa 
    #>  4          5.2         4.1          1.5         0.1 setosa 
    #>  5          5.5         4.2          1.4         0.2 setosa 
    #>  6          4.9         3.1          1.5         0.2 setosa 
    #>  7          5           3.2          1.2         0.2 setosa 
    #>  8          5.5         3.5          1.3         0.2 setosa 
    #>  9          4.9         3.6          1.4         0.1 setosa 
    #> 10          4.4         3            1.3         0.2 setosa 
    #> # … with 140 more rows
    

    由reprex package (v2.0.0) 于 2021 年 6 月 25 日创建

    您可以阅读更多关于建模过程中的二次抽样here 和here。

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 2019-08-20
      • 2019-06-21
      • 2015-07-23
      • 2018-09-10
      • 2017-12-09
      • 1970-01-01
      • 2019-03-18
      • 1970-01-01
      相关资源
      最近更新 更多