| Function | Works |
|---|---|
tidypredict_fit(), tidypredict_sql(),
parse_model() |
✔ |
tidypredict_to_column() |
✗ |
tidypredict_test() |
✗ |
tidypredict_interval(),
tidypredict_sql_interval() |
✗ |
parsnip |
✔ |
Here is a simple randomForest() model using the
mtcars dataset:
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.82000The 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))))/5From there, the Tidy Eval formula can be used anywhere where it can
be operated. tidypredict provides three paths:
dplyr,
mutate(iris, !! tidypredict_fit(model))tidypredict_to_column(model) to a piped command
settidypredict_to_sql(model) to retrieve the SQL
statementtidypredict 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