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)
#> ifelse(is.na(cyl) | is.na(disp) | is.na(hp) | is.na(drat) | is.na(wt) | 
#>     is.na(qsec) | is.na(vs) | is.na(am) | is.na(gear) | is.na(carb), 
#>     NA_real_, (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 it can be evaluated. 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)
#> ifelse(is.na(cyl) | is.na(disp) | is.na(hp) | is.na(drat) | is.na(wt) | 
#>     is.na(qsec) | is.na(vs) | is.na(am) | is.na(gear) | is.na(carb), 
#>     NA_real_, (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)