bagger models

Function Works
tidypredict_fit(), tidypredict_sql(), parse_model()
tidypredict_to_column()
tidypredict_test()
tidypredict_interval(), tidypredict_sql_interval()
parsnip

baguette::bagger() fits an ensemble of models on bootstrap samples of the training data. The "CART" base model, which fits rpart::rpart() trees, and the "C5.0" base model, which fits C50::C5.0() trees, are supported. tidypredict_fit() returns one nested case_when() per tree, so the size of the returned expression grows with times.

For regression models the fitted value is the mean of the individual tree predictions. For classification models the class probabilities of each tree are averaged, and the returned expression is the class with the largest average probability.

tidypredict_ functions

set.seed(100)
model <- baguette::bagger(mpg ~ wt + cyl + disp, data = mtcars, times = 5)
#> Registered S3 method overwritten by 'butcher':
#>   method                 from    
#>   as.character.dev_topic generics

Classification

set.seed(100)
model <- baguette::bagger(Species ~ ., data = iris, times = 3)

tidypredict_test(model, iris)
#> tidypredict test results
#> Difference threshold: 0
#> 
#>  All results are within the difference threshold

C5.0 trees are only fit for classification, and are used by passing base_model = "C5.0".

set.seed(100)
model <- baguette::bagger(
  Species ~ .,
  data = iris,
  base_model = "C5.0",
  times = 3
)

tidypredict_test(model, iris)
#> tidypredict test results
#> Difference threshold: 0
#> 
#>  All results are within the difference threshold

parsnip

Models fit with parsnip::bag_tree() and the "rpart" or "C5.0" engine are supported as well.

library(parsnip)

set.seed(100)
model <- bag_tree(mode = "regression") %>%
  set_engine("rpart", times = 5) %>%
  fit(mpg ~ wt + cyl + disp, data = mtcars)

tidypredict_fit(model)
#> (case_when(case_when(!is.na(wt) ~ wt < 2.0775, !is.na(disp) ~ 
#>     disp < 101.55, .default = FALSE) ~ case_when(case_when(!is.na(wt) ~ 
#>     wt < 1.674, .default = FALSE) ~ 30.4, !is.na(wt) ~ 33.9, 
#>     .default = 32.15), .default = case_when(case_when(!is.na(wt) ~ 
#>     wt < 2.975, !is.na(disp) ~ disp < 163.8, !is.na(cyl) ~ cyl < 
#>     7, .default = FALSE) ~ case_when(case_when(!is.na(wt) ~ wt < 
#>     2.3925, !is.na(disp) ~ disp < 114.05, .default = FALSE) ~ 
#>     22.8, .default = case_when(case_when(!is.na(cyl) ~ cyl < 
#>     5, !is.na(disp) ~ disp < 133, !is.na(wt) ~ wt < 2.5425, .default = FALSE) ~ 
#>     case_when(case_when(!is.na(wt) ~ wt < 2.6225, .default = TRUE) ~ 
#>         21.5, .default = 21.4), .default = case_when(case_when(!is.na(disp) ~ 
#>     disp < 152.5, !is.na(wt) ~ !wt < 2.695, .default = FALSE) ~ 
#>     19.7, .default = 21))), .default = case_when(case_when(!is.na(disp) ~ 
#>     disp < 450, !is.na(wt) ~ wt < 4.66, .default = TRUE) ~ case_when(case_when(!is.na(disp) ~ 
#>     disp < 355.5, !is.na(wt) ~ wt < 3.7875, .default = TRUE) ~ 
#>     case_when(case_when(!is.na(disp) ~ disp < 288.4, !is.na(wt) ~ 
#>         !wt < 3.65, !is.na(cyl) ~ cyl < 7, .default = FALSE) ~ 
#>         case_when(case_when(!is.na(wt) ~ wt < 3.9, .default = TRUE) ~ 
#>             case_when(case_when(!is.na(wt) ~ wt < 3.595, .default = TRUE) ~ 
#>                 case_when(case_when(!is.na(wt) ~ wt < 3.45, .default = FALSE) ~ 
#>                   17.8, !is.na(wt) ~ 18.1, .default = 17.95), 
#>                 .default = 17.3), .default = 16.4), .default = case_when(case_when(!is.na(disp) ~ 
#>         disp < 311, !is.na(wt) ~ !wt < 3.545, .default = FALSE) ~ 
#>         case_when(case_when(!is.na(disp) ~ disp < 302.5, .default = TRUE) ~ 
#>             15, .default = 15.2), !is.na(disp) | !is.na(wt) ~ 
#>         case_when(case_when(!is.na(wt) ~ wt < 3.345, .default = FALSE) ~ 
#>             15.8, .default = 15.5), .default = 15.3333333333333)), 
#>     .default = case_when(case_when(!is.na(wt) ~ wt < 4.595, .default = TRUE) ~ 
#>         case_when(case_when(!is.na(wt) ~ wt < 3.7075, !is.na(disp) ~ 
#>             disp < 380, .default = FALSE) ~ case_when(case_when(!is.na(wt) ~ 
#>             wt < 3.505, .default = TRUE) ~ 18.7, .default = 14.3), 
#>             !is.na(wt) | !is.na(disp) ~ 19.2, .default = 18.2166666666667), 
#>         .default = 14.7)), .default = 10.4))) + case_when(case_when(!is.na(wt) ~ 
#>     wt < 2.26, !is.na(disp) ~ disp < 101.55, .default = FALSE) ~ 
#>     case_when(case_when(!is.na(wt) ~ wt < 2.17, !is.na(disp) ~ 
#>         !disp < 78.85, .default = TRUE) ~ case_when(case_when(!is.na(wt) ~ 
#>         wt < 1.724, .default = FALSE) ~ 30.4, .default = case_when(case_when(!is.na(wt) ~ 
#>         wt < 2.0375, .default = FALSE) ~ 27.3, !is.na(wt) ~ 26, 
#>         .default = 26.65)), .default = 32.4), .default = case_when(case_when(!is.na(cyl) ~ 
#>     cyl < 7, !is.na(disp) ~ disp < 266.9, !is.na(wt) ~ wt < 3.515, 
#>     .default = FALSE) ~ case_when(case_when(!is.na(wt) ~ wt < 
#>     3.3275, !is.na(disp) ~ disp < 163.8, .default = TRUE) ~ case_when(case_when(!is.na(cyl) ~ 
#>     cyl < 5, !is.na(disp) ~ disp < 142.9, !is.na(wt) ~ wt < 2.5425, 
#>     .default = FALSE) ~ case_when(case_when(!is.na(wt) ~ wt < 
#>     2.3925, .default = FALSE) ~ 22.8, .default = case_when(case_when(!is.na(wt) ~ 
#>     wt < 2.965, .default = TRUE) ~ case_when(case_when(!is.na(wt) ~ 
#>     wt < 2.6225, .default = FALSE) ~ 21.5, .default = 21.4), 
#>     .default = 22.8)), !is.na(cyl) | !is.na(disp) | !is.na(wt) ~ 
#>     case_when(case_when(!is.na(disp) ~ disp < 152.5, .default = FALSE) ~ 
#>         19.7, .default = case_when(case_when(!is.na(wt) ~ wt < 
#>         3.045, .default = TRUE) ~ 21, .default = 21.4)), .default = 21.4), 
#>     .default = case_when(case_when(!is.na(wt) ~ wt < 3.45, .default = FALSE) ~ 
#>         19.2, !is.na(wt) ~ 18.1, .default = 18.65)), .default = case_when(case_when(!is.na(disp) ~ 
#>     disp < 430, !is.na(wt) ~ wt < 4.747, .default = TRUE) ~ case_when(case_when(!is.na(disp) ~ 
#>     disp < 380, !is.na(wt) ~ wt < 3.8425, .default = TRUE) ~ 
#>     case_when(case_when(!is.na(wt) ~ wt < 3.955, .default = TRUE) ~ 
#>         case_when(case_when(!is.na(wt) ~ wt < 3.81, .default = TRUE) ~ 
#>             case_when(case_when(!is.na(disp) ~ disp < 355.5, 
#>                 .default = TRUE) ~ case_when(case_when(!is.na(disp) ~ 
#>                 disp < 326, .default = TRUE) ~ case_when(case_when(!is.na(wt) ~ 
#>                 wt < 3.675, .default = FALSE) ~ 15, .default = 15.2), 
#>                 .default = 15.8), .default = 14.3), .default = 13.3), 
#>         .default = 16.4), .default = 19.2), .default = 10.4))) + 
#>     case_when(case_when(!is.na(wt) ~ wt < 2.26, !is.na(disp) ~ 
#>         disp < 101.55, !is.na(cyl) ~ cyl < 5, .default = FALSE) ~ 
#>         case_when(case_when(!is.na(disp) ~ disp < 78.85, !is.na(wt) ~ 
#>             !wt < 2.0675, .default = TRUE) ~ case_when(case_when(!is.na(wt) ~ 
#>             wt < 1.9075, !is.na(disp) ~ disp < 77.2, .default = FALSE) ~ 
#>             30.4, !is.na(wt) | !is.na(disp) ~ 32.4, .default = 31.4), 
#>             .default = case_when(case_when(!is.na(wt) ~ wt < 
#>                 1.724, .default = FALSE) ~ 30.4, .default = 27.3)), 
#>         .default = case_when(case_when(!is.na(cyl) ~ cyl < 7, 
#>             !is.na(disp) ~ disp < 250.4, !is.na(wt) ~ wt < 3.3125, 
#>             .default = TRUE) ~ case_when(case_when(!is.na(wt) ~ 
#>             wt < 3.315, !is.na(disp) ~ disp < 163.8, .default = TRUE) ~ 
#>             case_when(case_when(!is.na(wt) ~ wt < 3.0325, .default = TRUE) ~ 
#>                 case_when(case_when(!is.na(wt) ~ wt < 2.47, .default = FALSE) ~ 
#>                   22.8, .default = case_when(case_when(!is.na(disp) ~ 
#>                   disp < 152.5, !is.na(wt) ~ !wt < 2.695, .default = FALSE) ~ 
#>                   case_when(case_when(!is.na(wt) ~ wt < 2.775, 
#>                     .default = TRUE) ~ 19.7, .default = 21.4), 
#>                   .default = 21)), .default = 24.4), .default = case_when(case_when(!is.na(wt) ~ 
#>             wt < 3.45, !is.na(disp) ~ disp < 196.3, .default = FALSE) ~ 
#>             18.5, !is.na(wt) | !is.na(disp) ~ 18.1, .default = 18.3)), 
#>             .default = case_when(case_when(!is.na(wt) ~ wt < 
#>                 4.49, !is.na(disp) ~ disp < 410, .default = TRUE) ~ 
#>                 case_when(case_when(!is.na(disp) ~ disp < 339, 
#>                   .default = TRUE) ~ case_when(case_when(!is.na(wt) ~ 
#>                   wt < 3.65, .default = TRUE) ~ case_when(case_when(!is.na(disp) ~ 
#>                   disp < 311, .default = TRUE) ~ case_when(case_when(!is.na(wt) ~ 
#>                   wt < 3.5025, .default = TRUE) ~ 15.2, .default = 15), 
#>                   .default = 15.5), .default = 17.3), .default = case_when(case_when(!is.na(wt) ~ 
#>                   wt < 3.505, .default = TRUE) ~ 18.7, .default = 14.3)), 
#>                 .default = 10.4))) + case_when(case_when(!is.na(disp) ~ 
#>     disp < 163.8, !is.na(wt) ~ wt < 3.3125, !is.na(cyl) ~ cyl < 
#>     5, .default = FALSE) ~ case_when(case_when(!is.na(disp) ~ 
#>     disp < 101.55, !is.na(wt) ~ wt < 2.0375, .default = FALSE) ~ 
#>     case_when(case_when(!is.na(wt) ~ wt < 2.0675, .default = TRUE) ~ 
#>         case_when(case_when(!is.na(wt) ~ wt < 1.775, !is.na(disp) ~ 
#>             !disp < 87.05, .default = TRUE) ~ case_when(case_when(!is.na(wt) ~ 
#>             wt < 1.564, .default = TRUE) ~ 30.4, .default = 30.4), 
#>             .default = 27.3), .default = 32.4), .default = case_when(case_when(!is.na(wt) ~ 
#>     wt < 2.23, .default = FALSE) ~ 26, .default = case_when(case_when(!is.na(wt) ~ 
#>     wt < 3.0325, !is.na(disp) ~ disp < 133.85, .default = TRUE) ~ 
#>     case_when(case_when(!is.na(wt) ~ wt < 2.3925, .default = FALSE) ~ 
#>         22.8, .default = case_when(case_when(!is.na(wt) ~ wt < 
#>         2.8275, !is.na(cyl) ~ cyl < 5, !is.na(disp) ~ disp < 
#>         140.5, .default = FALSE) ~ case_when(case_when(!is.na(wt) ~ 
#>         wt < 2.6225, .default = FALSE) ~ 21.5, !is.na(wt) ~ 21.4, 
#>         .default = 21.45), !is.na(wt) | !is.na(cyl) | !is.na(disp) ~ 
#>         21, .default = 21.225)), .default = 24.4))), .default = case_when(case_when(!is.na(wt) ~ 
#>     wt < 4.66, !is.na(disp) ~ disp < 410, .default = TRUE) ~ 
#>     case_when(case_when(!is.na(disp) ~ disp < 288.4, !is.na(wt) ~ 
#>         !wt < 3.65, !is.na(cyl) ~ cyl < 7, .default = FALSE) ~ 
#>         case_when(case_when(!is.na(wt) ~ wt < 3.755, !is.na(cyl) ~ 
#>             cyl < 7, !is.na(disp) ~ disp < 221.7, .default = TRUE) ~ 
#>             case_when(case_when(!is.na(wt) ~ wt < 3.585, .default = TRUE) ~ 
#>                 17.8, .default = 17.3), .default = case_when(case_when(!is.na(wt) ~ 
#>             wt < 3.925, .default = FALSE) ~ 15.2, !is.na(wt) ~ 
#>             16.4, .default = 15.8)), .default = case_when(case_when(!is.na(wt) ~ 
#>         wt < 3.545, !is.na(disp) ~ disp < 334, .default = TRUE) ~ 
#>         case_when(case_when(!is.na(disp) ~ disp < 311, .default = TRUE) ~ 
#>             15.2, .default = case_when(case_when(!is.na(wt) ~ 
#>             wt < 3.345, .default = FALSE) ~ 15.8, !is.na(wt) ~ 
#>             15.5, .default = 15.65)), .default = case_when(case_when(!is.na(wt) ~ 
#>         wt < 3.705, .default = TRUE) ~ case_when(case_when(!is.na(disp) ~ 
#>         disp < 330.5, .default = FALSE) ~ 15, .default = 14.3), 
#>         .default = 13.3))), .default = 10.4)) + case_when(case_when(!is.na(wt) ~ 
#>     wt < 2.41, !is.na(disp) ~ disp < 120.65, !is.na(cyl) ~ cyl < 
#>     5, .default = FALSE) ~ case_when(case_when(!is.na(disp) ~ 
#>     disp < 107.7, .default = TRUE) ~ case_when(case_when(!is.na(wt) ~ 
#>     wt < 1.9075, !is.na(disp) ~ !disp < 86.9, .default = TRUE) ~ 
#>     30.4, .default = 32.4), .default = 26), .default = case_when(case_when(!is.na(disp) ~ 
#>     disp < 266.9, !is.na(cyl) ~ cyl < 7, !is.na(wt) ~ wt < 3.325, 
#>     .default = FALSE) ~ case_when(case_when(!is.na(cyl) ~ cyl < 
#>     5, !is.na(disp) ~ disp < 153.35, !is.na(wt) ~ wt < 3.2025, 
#>     .default = FALSE) ~ case_when(case_when(!is.na(wt) ~ wt < 
#>     3.17, !is.na(disp) ~ disp < 143.75, .default = TRUE) ~ case_when(case_when(!is.na(wt) ~ 
#>     wt < 2.965, .default = TRUE) ~ 21.4, .default = 22.8), .default = 24.4), 
#>     !is.na(cyl) | !is.na(disp) | !is.na(wt) ~ case_when(case_when(!is.na(wt) ~ 
#>         wt < 3.3275, !is.na(disp) ~ disp < 163.8, .default = TRUE) ~ 
#>         case_when(case_when(!is.na(disp) ~ disp < 152.5, .default = FALSE) ~ 
#>             19.7, .default = case_when(case_when(!is.na(wt) ~ 
#>             wt < 3.045, .default = TRUE) ~ 21, .default = 21.4)), 
#>         .default = 18.5), .default = 21.325), !is.na(disp) | 
#>     !is.na(cyl) | !is.na(wt) ~ case_when(case_when(!is.na(wt) ~ 
#>     wt < 4.66, !is.na(disp) ~ disp < 410, .default = TRUE) ~ 
#>     case_when(case_when(!is.na(wt) ~ wt < 3.505, !is.na(disp) ~ 
#>         !disp < 302.5, .default = FALSE) ~ case_when(case_when(!is.na(wt) ~ 
#>         wt < 3.4375, .default = TRUE) ~ case_when(case_when(!is.na(disp) ~ 
#>         disp < 327.5, .default = TRUE) ~ 15.2, .default = 15.8), 
#>         .default = 18.7), .default = case_when(case_when(!is.na(wt) ~ 
#>         wt < 3.955, !is.na(disp) ~ !disp < 288.4, .default = TRUE) ~ 
#>         case_when(case_when(!is.na(wt) ~ wt < 3.81, .default = TRUE) ~ 
#>             case_when(case_when(!is.na(disp) ~ disp < 330.5, 
#>                 .default = TRUE) ~ case_when(case_when(!is.na(wt) ~ 
#>                 wt < 3.675, .default = FALSE) ~ 15, !is.na(wt) ~ 
#>                 15.2, .default = 15.1), .default = 14.3), .default = 13.3), 
#>         .default = 16.4)), .default = 10.4), .default = 18.0083333333333)))/5L