Oblique Random Forest, using aorsf

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

Only regression models with numeric predictors are supported. Classification requires a voting mechanism that cannot be expressed as a single formula, and oblique splits on categorical predictors operate on an internal encoding that is not reproduced here.

How it works

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

library(dplyr)
library(tidypredict)
library(aorsf)

model <- orsf(mtcars, mpg ~ ., n_tree = 5)

Under the hood

Unlike axis-aligned forests, each split in an oblique random forest is a linear combination of (standardized) predictors compared against a cutpoint. tidypredict reads the trees stored in the fitted forest, folds the forest’s centering and scaling into the split coefficients, and turns each tree into a nested dplyr::case_when() statement. The trees are then averaged.

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(2.06536636741246 * drat + -0.00730901496547223 * 
#>         hp + -1.60923770108693 * cyl + -0.0116412947123201 * 
#>         disp <= -0.976685770752394 ~ case_when(1.33600153764993 * 
#>         gear + -0.877675539168235 * wt + -0.01637170336932 * 
#>         disp + -0.778241609587267 * carb <= -4.96165809189402 ~ 
#>         case_when(-0.847698481375021 * carb + 2.87017225467603 * 
#>             am + -0.0152116696806259 * disp <= -8.49012817019253 ~ 
#>             13.3, .default = 16.2166666666667), .default = case_when(-4.67779090266086 * 
#>         vs + 0.918943275372421 * qsec <= 13.9032421253695 ~ 19.05, 
#>         .default = 20.74)), .default = 27.3333333333333) + case_when(-0.619435390359988 * 
#>         qsec + 0.485193348925287 * gear + -3.43370586093924 * 
#>         cyl + -0.400649386990832 * vs <= -34.0254933031073 ~ 
#>         case_when(-0.00663936492314984 * hp + -2.23454809400203 * 
#>             wt + 0.908222187542599 * drat <= -10.4313562932566 ~ 
#>             11.26, .default = case_when(-0.00575446001389749 * 
#>             hp + 0.0178945235954444 * disp <= 4.57676617093046 ~ 
#>             15.16, .default = 15.96)), .default = case_when(-2.80881466405904 * 
#>         vs + -2.8426126824166 * carb + 5.08377087539071 * gear + 
#>         -0.027788874912095 * disp <= 8.4785896083071 ~ case_when(-0.0282375336950912 * 
#>         disp + -3.7706864647061 * drat <= -19.2236826035684 ~ 
#>         19.8666666666667, .default = 21.9), .default = 29.76)) + 
#>         case_when(1.83283067644656 * vs + 7.09327393625157 * 
#>             am + 2.14987922789576 * drat + 1.81127586660398 * 
#>             qsec <= 44.4934711286132 ~ case_when(-0.0198612578403386 * 
#>             hp + -1.95757007384445 * wt + 0.569267232831812 * 
#>             vs + 0.270535049271622 * carb <= -9.32876523744002 ~ 
#>             case_when(3.62806688200883 * am + -0.934771723499277 * 
#>                 carb + -1.28873797086969 * drat <= -7.5150891486453 ~ 
#>                 13.54, .default = 16.0444444444444), .default = 18.58), 
#>             .default = case_when(4.51911439486263 * gear + -0.0899527281617354 * 
#>                 hp + -1.88846514891657 * carb <= 4.49467991198821 ~ 
#>                 21.28, .default = 29.2625)) + case_when(-0.0197303575282214 * 
#>         hp + -0.482615605145762 * carb + 2.03062149221965 * drat + 
#>         -3.27930365367313 * wt <= -2.10763058674639 ~ case_when(-0.801577475550684 * 
#>         cyl + -0.0249207173057865 * hp + 1.20155162574692 * am + 
#>         0.75489603295991 * drat <= -5.34262711944234 ~ case_when(-2.6855500871004 * 
#>         drat + -0.282315855955557 * vs + -2.59293686166068 * 
#>         wt <= -18.763631904464 ~ 13.84, .default = case_when(0.0272865582343068 * 
#>         disp + -10.6314879455189 * drat <= -25.1130352317213 ~ 
#>         16.04, .default = 18.66)), .default = 21.5444444444444), 
#>         .default = 28.825) + case_when(-0.445623850825133 * cyl + 
#>         1.00871006237179 * am + -0.0195950630540864 * hp + -2.61351296011026 * 
#>         wt <= -14.0744204433827 ~ case_when(-1.49283057960904 * 
#>         am + -1.64180931695495 * cyl + -0.517642126196677 * qsec <= 
#>         -22.1518003739857 ~ 13.7, .default = case_when(-1.20108088369544 * 
#>         am + 0.616395837137112 * wt <= 2.12040167975166 ~ 17.58, 
#>         .default = 16.7833333333333)), .default = case_when(0.00161167151624021 * 
#>         disp + -4.82977886840442 * wt <= -15.1119278107302 ~ 
#>         19.75, .default = 23.025)))/5L)

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

A note on split boundaries

aorsf uses observed linear-combination values from the training data as split cutpoints. A training row can therefore land exactly on a split boundary, where floating-point differences between aorsf’s internal traversal and the generated formula may send it down a different branch. This affects only rows that coincide with a training cutpoint; on new data the formula reproduces predict() exactly.

parsnip

tidypredict also supports aorsf model objects fitted via the parsnip package (using the bonsai extension).

library(parsnip)
library(bonsai)

parsnip_model <- rand_forest(mode = "regression", trees = 5) %>%
  set_engine("aorsf") %>%
  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(-1.2734192843392 * qsec + -0.0598951391406583 * 
#>         hp + 3.29342977695824 * vs + 0.0442820913297539 * carb <= 
#>         -34.6681421880752 ~ 12.6888888888889, .default = case_when(-0.473191279714246 * 
#>         am + 2.5544574454982 * gear + -0.774975982263831 * cyl + 
#>         -0.828562630865865 * carb <= 2.67786426899857 ~ case_when(-3.14088167064683 * 
#>         wt + 1.89280081676347 * vs + 0.0132372629374619 * hp <= 
#>         -8.80333909805259 ~ 16.6, .default = 19.5), .default = 23.3)) + 
#>         case_when(-0.91281431138112 * cyl + 5.82725172803301 * 
#>             gear + -2.48128733833168 * carb + -0.0283305994075317 * 
#>             qsec <= 14.1682258409641 ~ case_when(-0.219475547908853 * 
#>             carb + 0.438643840817042 * am + -1.03516346973444 * 
#>             cyl + -2.0457751746568 * wt <= -13.508838470628 ~ 
#>             case_when(6.2321983033304 * am + -0.0427664085398967 * 
#>                 hp + 1.08902964909628 * drat <= -4.23938359973417 ~ 
#>                 13.38, .default = case_when(-0.751666666666669 * 
#>                 cyl <= -6.01333333333335 ~ 16.4166666666667, 
#>                 .default = 17.92)), .default = 21.5571428571429), 
#>             .default = 31.0555555555556) + case_when(-0.0155088133947272 * 
#>         disp + -2.90168726517077 * cyl + -0.711151871779383 * 
#>         gear + -2.92682093212002 * vs <= -22.7361412212985 ~ 
#>         case_when(-2.17476075440768 * wt + 0.441602303596645 * 
#>             gear + -0.0193547995986316 * hp + -0.00952700629782295 * 
#>             vs <= -11.7682002878003 ~ 12.9875, .default = case_when(0.527882037533511 * 
#>             vs + -0.696782841823054 * gear + 0.87479892761394 * 
#>             carb <= 0.534048257372655 ~ case_when(0.329788758178246 * 
#>             qsec + 0.0569174931151015 * cyl <= 6.06834460911455 ~ 
#>             16.2, .default = 18.18), .default = 19.5857142857143)), 
#>         .default = 28.2285714285714) + case_when(-2.39048448845207 * 
#>         cyl + 3.06294289643119 * vs + -2.44369092908812 * qsec + 
#>         -1.7710266638723 * carb <= -60.5563423596266 ~ case_when(-3.50907701490441 * 
#>         wt + 0.475524182845409 * vs + -2.02290238067963 * drat + 
#>         -0.918598037002364 * cyl <= -27.8410380553663 ~ 13.3666666666667, 
#>         .default = case_when(-2.73093950112935 * cyl + -1.39929397089746 * 
#>             carb + -1.65459347053921 * qsec <= -51.2408109519763 ~ 
#>             case_when(0.0056509586974207 * disp + 0.625561971219155 * 
#>                 carb <= 3.43522032240609 ~ 16.48, .default = 19.2), 
#>             .default = 20.9)), .default = case_when(0.19665100265338 * 
#>         qsec + -9.47942645524413 * wt <= -17.001900869911 ~ 24.0666666666667, 
#>         .default = 29.16)) + case_when(1.4597768014324 * qsec + 
#>         2.93923658504198 * am + -3.53347907993741 * wt + 0.919950754641844 * 
#>         drat <= 19.1082616876211 ~ case_when(-0.831846220453981 * 
#>         wt + -0.230794186368748 * am + -0.0229438050936533 * 
#>         disp <= -11.2294608407359 ~ 12.74, .default = case_when(3.64398894726172 * 
#>         am + -0.0285141125146731 * hp <= -5.13254025264116 ~ 
#>         15.5, .default = 16.7666666666667)), .default = case_when(0.0102135489883771 * 
#>         disp + -5.08579851421966 * wt + -1.96332947140481 * cyl <= 
#>         -19.1631639896666 ~ case_when(0.00159589201249759 * disp + 
#>         1.90019508078217 * drat <= 7.24583720631959 ~ 21.36, 
#>         .default = 21.36), .default = 28.74)))/5L)