Odhad výkonu pomocí křížové validace

Modeling with tidymodels in R

David Svancer

Data Scientist

Trénovací a testovací sady

 

Vytvoření trénovací a testovací sady je prvním krokem modelování

  • Chrání před přetrénováním
    • Trénovací data slouží k přizpůsobení modelu
    • Testovací data slouží k vyhodnocení modelu

 

Nevýhoda

  • Pouze jeden odhad výkonu modelu

 

Diagram rozdělení dat na trénovací a testovací sadu

Modeling with tidymodels in R

K-násobná křížová validace

Technika resamplingů pro zkoumání výkonu modelu

  • Poskytuje K odhadů výkonu modelu v průběhu jeho přizpůsobení

 

Rozdělení dat na trénovací a testovací sady

Modeling with tidymodels in R

K-násobná křížová validace

Technika resamplingů pro zkoumání výkonu modelu

  • Poskytuje K odhadů výkonu modelu v průběhu jeho přizpůsobení
  • Trénovací data jsou náhodně rozdělena do K přibližně stejně velkých podmnožin
  • Záložky jsou použity k provedení K iterací přizpůsobení a vyhodnocení modelu

 

Rozdělení trénovacích dat do záložek křížové validace

Modeling with tidymodels in R

Strojové učení s křížovou validací

Provádění 5-násobné křížové validace

  • Pět iterací trénování a vyhodnocení modelu

 

První iterace pětinásobné křížové validace

Modeling with tidymodels in R

Strojové učení s křížovou validací

Provádění 5-násobné křížové validace

  • Pět iterací trénování a vyhodnocení modelu
  • Iterace 1
    • Záložka 1 vyhrazena pro vyhodnocení, záložky 2–5 pro trénování

 

První iterace pětinásobné křížové validace

Modeling with tidymodels in R

Strojové učení s křížovou validací

Provádění 5-násobné křížové validace

  • Pět iterací trénování a vyhodnocení modelu
  • Iterace 1
    • Záložka 1 vyhrazena pro vyhodnocení, záložky 2–5 pro trénování
  • Iterace 2
    • Záložka 2 vyhrazena pro vyhodnocení

 

Druhá iterace pětinásobné křížové validace

Modeling with tidymodels in R

Strojové učení s křížovou validací

Provádění 5-násobné křížové validace

  • Pět iterací trénování a vyhodnocení modelu
  • Iterace 1
    • Záložka 1 vyhrazena pro vyhodnocení, záložky 2–5 pro trénování
  • Iterace 2
    • Záložka 2 vyhrazena pro vyhodnocení

 

Pět odhadů výkonu modelu celkem

 

Pátá iterace pětinásobné křížové validace

Modeling with tidymodels in R

Vytvoření záložek křížové validace

Funkce vfold_cv()

  • Trénovací data
  • Počet záložek, v
  • Stratifikační proměnná, strata
  • Před vfold_cv() spusťte set.seed() pro reprodukovatelnost
  • splits
    • Sloupcový seznam s objekty rozdělení dat pro vytvoření záložek
set.seed(214)
leads_folds <- vfold_cv(leads_training,

v = 10,
strata = purchased)
leads_folds
#  10-fold cross-validation using stratification 
# A tibble: 10 x 2
   splits            id    
   <list>            <chr> 
 1 <split [896/100]> Fold01
 2 <split [896/100]> Fold02
 3 <split [896/100]> Fold03
 . ................  ......
 9 <split [897/99]>  Fold09
10 <split [897/99]>  Fold10
Modeling with tidymodels in R

Trénování modelu s křížovou validací

Funkce fit_resamples()

  • Trénuje model parsnip nebo objekt workflow
  • Zadejte záložky křížové validace, resamples
  • Volitelná vlastní funkce metrik, metrics
    • Výchozí hodnoty jsou přesnost a ROC AUC

 

Každá metrika je odhadnuta 10krát

  • Jeden odhad na záložku
  • Průměrná hodnota ve sloupci mean
leads_rs_fit <- leads_wkfl %>%

fit_resamples(resamples = leads_folds,
metrics = leads_metrics)
leads_rs_fit %>% collect_metrics()
# A tibble: 3 x 5
  .metric .estimator  mean     n std_err
  <chr>   <chr>      <dbl> <int>   <dbl>
1 roc_auc binary     0.823    10  0.0147
2 sens    binary     0.786    10  0.0203
3 spec    binary     0.855    10  0.0159
Modeling with tidymodels in R

Podrobné výsledky křížové validace

Funkce collect_metrics()

  • Předání summarize = FALSE vrátí všechny odhady metrik pro každou záložku křížové validace
  • Celkem 30 kombinací (3 metriky × 10 záložek)
    • Sloupec .metric identifikuje metriku
    • Sloupec .estimate udává odhadovanou hodnotu pro každou záložku
rs_metrics <- leads_rs_fit %>% 
  collect_metrics(summarize = FALSE)

rs_metrics
# A tibble: 30 x 4
   id     .metric .estimator .estimate
   <chr>  <chr>   <chr>          <dbl>
 1 Fold01 sens    binary         0.861
 2 Fold01 spec    binary         0.891
 3 Fold01 roc_auc binary         0.885
 4 Fold02 sens    binary         0.778
 5 Fold02 spec    binary         0.969
 6 Fold02 roc_auc binary         0.885
# ... with 24 more rows
Modeling with tidymodels in R

Shrnutí výsledků křížové validace

Funkce collect_metrics() vrací tibble

  • Výsledky lze shrnout pomocí dplyr
    • Začněte s rs_metrics
    • Seskupte podle hodnot .metric
    • Vypočítejte souhrnné statistiky pomocí summarize()
rs_metrics %>%

group_by(.metric) %>%
summarize(min = min(.estimate), median = median(.estimate), max = max(.estimate), mean = mean(.estimate), sd = sd(.estimate))
# A tibble: 3 x 6
 .metric   min  median   max   mean     sd
  <chr>   <dbl>  <dbl>  <dbl>  <dbl>   <dbl>
1 roc_auc 0.758  0.806  0.885  0.823   0.0466
2 sens    0.667  0.792  0.861  0.786   0.0642
3 spec    0.810  0.843  0.969  0.855   0.0502
Modeling with tidymodels in R

Metodologie křížové validace

Modely trénované pomocí fit_resamples() nejsou schopny poskytovat predikce pro nové zdroje dat

  • Funkce predict() nepřijímá objekty resamplingů

Účel fit_resample()

  • Prozkoumejte a porovnejte výkonnostní profil různých typů modelů
  • Vyberte nejlepší typ modelu a soustřeďte se na jeho přizpůsobení
predict(leads_rs_fit,
        new_data = leads_test)

Error in UseMethod("predict") : 
  no applicable method for 'predict' applied to 
  an object of class 
  "c('resample_results', 
      'tune_results',  
      'tbl_df', 
      'tbl', 'data.frame')"
Modeling with tidymodels in R

Pojďme křížově validovat!

Modeling with tidymodels in R

Preparing Video For Download...