Skip to content

Instantly share code, notes, and snippets.

@DavisVaughan
Created November 28, 2018 13:52
Show Gist options
  • Select an option

  • Save DavisVaughan/38e8f06144e357d04c592373fd77a271 to your computer and use it in GitHub Desktop.

Select an option

Save DavisVaughan/38e8f06144e357d04c592373fd77a271 to your computer and use it in GitHub Desktop.
rap with models
suppressPackageStartupMessages({
library(rap)
library(AmesHousing)
library(rsample)
library(parsnip)
library(dplyr)
})
ames <- make_ames()
set.seed(4595)
data_split <- initial_split(ames, strata = "Sale_Price", prop = .75)
ames_train <- training(data_split)
ames_test <- testing(data_split)
spec_lin_reg <- linear_reg()
spec_lm <- set_engine(spec_lin_reg, "lm")
ames_train_log <- ames_train %>%
mutate(Sale_Price_Log = log10(Sale_Price))
set.seed(2453)
cv_splits <- vfold_cv(
data = ames_train_log,
v = 10,
strata = "Sale_Price"
)
# for each split, fit a model
cv_splits
#> # 10-fold cross-validation using stratification
#> # A tibble: 10 x 2
#> splits id
#> <list> <chr>
#> 1 <split [2K/222]> Fold01
#> 2 <split [2K/222]> Fold02
#> 3 <split [2K/222]> Fold03
#> 4 <split [2K/222]> Fold04
#> 5 <split [2K/222]> Fold05
#> 6 <split [2K/219]> Fold06
#> 7 <split [2K/219]> Fold07
#> 8 <split [2K/217]> Fold08
#> 9 <split [2K/217]> Fold09
#> 10 <split [2K/217]> Fold10
cv_splits %>%
rap(
models = ~fit(spec_lm, Sale_Price_Log ~ Latitude + Longitude, analysis(splits)),
assess = ~assessment(splits),
pred = ~predict(models, new_data = assess),
rmse = double() ~ rmse_vec(assess$Sale_Price_Log, pred$.pred)
)
#> # 10-fold cross-validation using stratification
#> # A tibble: 10 x 6
#> splits id models assess pred rmse
#> <list> <chr> <list> <list> <list> <dbl>
#> 1 <split [2K/222… Fold01 <fit[+]> <tibble [222 × 8… <tibble [222 × … 0.150
#> 2 <split [2K/222… Fold02 <fit[+]> <tibble [222 × 8… <tibble [222 × … 0.164
#> 3 <split [2K/222… Fold03 <fit[+]> <tibble [222 × 8… <tibble [222 × … 0.160
#> 4 <split [2K/222… Fold04 <fit[+]> <tibble [222 × 8… <tibble [222 × … 0.158
#> 5 <split [2K/222… Fold05 <fit[+]> <tibble [222 × 8… <tibble [222 × … 0.161
#> 6 <split [2K/219… Fold06 <fit[+]> <tibble [219 × 8… <tibble [219 × … 0.156
#> 7 <split [2K/219… Fold07 <fit[+]> <tibble [219 × 8… <tibble [219 × … 0.179
#> 8 <split [2K/217… Fold08 <fit[+]> <tibble [217 × 8… <tibble [217 × … 0.159
#> 9 <split [2K/217… Fold09 <fit[+]> <tibble [217 × 8… <tibble [217 × … 0.168
#> 10 <split [2K/217… Fold10 <fit[+]> <tibble [217 × 8… <tibble [217 × … 0.158
# compare against
cv_splits %>%
mutate(
models = map(splits, ~fit(spec_lm, Sale_Price_Log ~ Latitude + Longitude, analysis(.x))),
assess = map(splits, assessment),
pred = map2(models, assess, ~predict(.x, .y)),
rmse = map2_dbl(assess, pred, ~rmse_vec(.x$Sale_Price_Log, .y$.pred))
)
#> # 10-fold cross-validation using stratification
#> # A tibble: 10 x 6
#> splits id models assess pred rmse
#> * <list> <chr> <list> <list> <list> <dbl>
#> 1 <split [2K/222… Fold01 <fit[+]> <tibble [222 × 8… <tibble [222 × … 0.150
#> 2 <split [2K/222… Fold02 <fit[+]> <tibble [222 × 8… <tibble [222 × … 0.164
#> 3 <split [2K/222… Fold03 <fit[+]> <tibble [222 × 8… <tibble [222 × … 0.160
#> 4 <split [2K/222… Fold04 <fit[+]> <tibble [222 × 8… <tibble [222 × … 0.158
#> 5 <split [2K/222… Fold05 <fit[+]> <tibble [222 × 8… <tibble [222 × … 0.161
#> 6 <split [2K/219… Fold06 <fit[+]> <tibble [219 × 8… <tibble [219 × … 0.156
#> 7 <split [2K/219… Fold07 <fit[+]> <tibble [219 × 8… <tibble [219 × … 0.179
#> 8 <split [2K/217… Fold08 <fit[+]> <tibble [217 × 8… <tibble [217 × … 0.159
#> 9 <split [2K/217… Fold09 <fit[+]> <tibble [217 × 8… <tibble [217 × … 0.168
#> 10 <split [2K/217… Fold10 <fit[+]> <tibble [217 × 8… <tibble [217 × … 0.158
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment