Random Forest

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

How it works

Here is a simple randomForest() model using the mtcars dataset:

library(dplyr)
library(tidypredict)
library(randomForest)

model <- randomForest(mpg ~ ., data = mtcars, ntree = 5, proximity = TRUE)

Under the hood

The parser is based on the output from the randomForest::getTree() function. It will return as many decision paths as there are non-NA rows in the prediction field.

getTree(model, labelVar = TRUE) %>%
  head()
#>   left daughter right daughter split var split point status prediction
#> 1             2              3      carb       1.500     -3   20.12813
#> 2             4              5        hp      85.500     -3   28.82222
#> 3             6              7        wt       3.160     -3   16.72609
#> 4             8              9        wt       1.885     -3   31.88571
#> 5             0              0      <NA>       0.000     -1   18.10000
#> 6             0              0      <NA>       0.000     -1   21.82000

The output from parse_model() is transformed into a dplyr, a.k.a Tidy Eval, formula. Each decision tree becomes one dplyr::case_when() statement, which are then combined.

tidypredict_fit(model)
#> (case_when(carb <= 1.5 ~ case_when(hp <= 85.5 ~ case_when(wt <= 
#>     1.885 ~ 33.9, .default = case_when(wt <= 2.0675 ~ 27.3, .default = 32.4)), 
#>     .default = 18.1), .default = case_when(wt <= 3.16 ~ 21.82, 
#>     .default = case_when(carb <= 3.5 ~ case_when(wt <= 3.8125 ~ 
#>         15.68, .default = 17.52), .default = case_when(qsec <= 
#>         18.14 ~ case_when(hp <= 230 ~ 10.4, .default = 14.8), 
#>         .default = 19.2)))) + case_when(vs <= 0.5 ~ case_when(disp <= 
#>     217.9 ~ case_when(drat <= 4.165 ~ case_when(hp <= 142.5 ~ 
#>     21, .default = 19.7), .default = 26), .default = case_when(qsec <= 
#>     17.71 ~ case_when(hp <= 212.5 ~ 17.04, .default = case_when(drat <= 
#>     3.635 ~ 14.65, .default = 13.3)), .default = 10.4)), .default = case_when(wt <= 
#>     2.26 ~ 30.9, .default = case_when(qsec <= 19.72 ~ 21.75, 
#>     .default = 22.95))) + case_when(disp <= 142.9 ~ case_when(hp <= 
#>     65.5 ~ 33.9, .default = case_when(wt <= 2.23 ~ 27.75, .default = 22.8)), 
#>     .default = case_when(drat <= 3.58 ~ case_when(wt <= 3.65 ~ 
#>         case_when(qsec <= 16.355 ~ 14.7666666666667, .default = case_when(disp <= 
#>             339 ~ 15.35, .default = 18.7)), .default = 18.025), 
#>         .default = case_when(disp <= 163.8 ~ 21.16, .default = 17.68))) + 
#>     case_when(disp <= 163.8 ~ case_when(drat <= 4 ~ 22.5, .default = 28.12), 
#>         .default = case_when(carb <= 3.5 ~ case_when(cyl <= 7 ~ 
#>             21.4, .default = case_when(wt <= 3.4375 ~ 15.2, .default = case_when(drat <= 
#>             3.075 ~ case_when(wt <= 3.755 ~ case_when(wt <= 3.625 ~ 
#>             15.5, .default = 17.3), .default = 15.8), .default = 18.95))), 
#>             .default = case_when(vs <= 0.5 ~ case_when(hp <= 
#>                 217.5 ~ 10.4, .default = case_when(drat <= 3.635 ~ 
#>                 14.575, .default = 13.3)), .default = 18.2666666666667))) + 
#>     case_when(carb <= 2.5 ~ case_when(wt <= 2.04 ~ 30.4, .default = case_when(qsec <= 
#>         19.72 ~ case_when(wt <= 3.53 ~ case_when(carb <= 1.5 ~ 
#>         21.4, .default = 21.4), .default = 19.2), .default = 22.9)), 
#>         .default = case_when(hp <= 192.5 ~ case_when(hp <= 116.5 ~ 
#>             21, .default = case_when(vs <= 0.5 ~ 16.85, .default = 18.36)), 
#>             .default = case_when(wt <= 4.545 ~ 13.6333333333333, 
#>                 .default = 10.4))))/5

From there, the Tidy Eval formula can be used anywhere where it can be operated. tidypredict provides three paths:

parsnip

tidypredict also supports randomForest model objects fitted via the parsnip package.

library(parsnip)

parsnip_model <- rand_forest(mode = "regression", trees = 5) %>%
  set_engine("randomForest") %>%
  fit(mpg ~ ., data = mtcars)

tidypredict_fit(parsnip_model)
#> (case_when(wt <= 2.3325 ~ case_when(drat <= 4.325 ~ case_when(qsec <= 
#>     19.185 ~ 29.3666666666667, .default = 32.9), .default = 26), 
#>     .default = case_when(hp <= 116.5 ~ case_when(disp <= 153.35 ~ 
#>         22.6, .default = 20.1666666666667), .default = case_when(disp <= 
#>         221.7 ~ case_when(wt <= 3.105 ~ 19.7, .default = 18.64), 
#>         .default = case_when(drat <= 3.18 ~ case_when(carb <= 
#>             2.5 ~ 17.46, .default = 16.55), .default = 15.05)))) + 
#>     case_when(wt <= 2.26 ~ 33.525, .default = case_when(drat <= 
#>         3.04 ~ case_when(drat <= 2.845 ~ 15.5, .default = 10.4), 
#>         .default = case_when(hp <= 116.5 ~ case_when(hp <= 96 ~ 
#>             23.2, .default = case_when(hp <= 109.5 ~ 21.4333333333333, 
#>             .default = 21.1)), .default = case_when(disp <= 235.8 ~ 
#>             19.22, .default = 16.92)))) + case_when(gear <= 3.5 ~ 
#>     case_when(carb <= 1.5 ~ 20.625, .default = case_when(qsec <= 
#>         17.62 ~ case_when(drat <= 3.48 ~ 15.375, .default = 13.3), 
#>         .default = case_when(carb <= 3.5 ~ 15.2, .default = 10.4))), 
#>     .default = case_when(hp <= 79.5 ~ case_when(qsec <= 19.185 ~ 
#>         29.625, .default = 33.15), .default = case_when(gear <= 
#>         4.5 ~ case_when(wt <= 3.295 ~ case_when(cyl <= 5 ~ 22.45, 
#>         .default = 21), .default = 17.8), .default = 30.4))) + 
#>     case_when(wt <= 2.3025 ~ case_when(hp <= 65.5 ~ 32.15, .default = case_when(hp <= 
#>         102 ~ 26.65, .default = 30.4)), .default = case_when(disp <= 
#>         250.4 ~ case_when(qsec <= 21.56 ~ case_when(disp <= 133 ~ 
#>         21.425, .default = case_when(qsec <= 19.26 ~ 19.5, .default = 18.1)), 
#>         .default = 22.8), .default = case_when(wt <= 4.5475 ~ 
#>         case_when(disp <= 355.5 ~ 15.76, .default = 18.95), .default = 12.55))) + 
#>     case_when(wt <= 3.2025 ~ case_when(wt <= 2.49 ~ case_when(hp <= 
#>         65.5 ~ 33.9, .default = 28.94), .default = 23.4), .default = case_when(hp <= 
#>         197.5 ~ case_when(qsec <= 19.17 ~ case_when(drat <= 2.915 ~ 
#>         15.5, .default = case_when(qsec <= 17.175 ~ 19.075, .default = 16.66)), 
#>         .default = 21.4), .default = 13.54)))/5