ML Atlas

04 · Ocena · 4 min czytania · Interaktywne · aktualizacja

Jaki rodzaj walidacji krzyżowej wybrać: KFold, Stratified, Group czy TimeSeriesSplit?

W skrócie

Wariant walidacji krzyżowej musi naśladować sposób, w jaki model spotka nowe dane: warstwowanie dla klas, grupy dla pacjentów, porządek czasu dla prognoz.

Co to jest

Warianty walidacji krzyżowej to różne sposoby dzielenia danych na foldy: zwykły podział na k części (KFold), podział warstwowy zachowujący proporcje klas (StratifiedKFold), podział po grupach (GroupKFold), podział chronologiczny (TimeSeriesSplit), leave-one-out i wersje powtarzane. Wybór wariantu nie jest kwestią gustu — decyduje o tym, czy wynik walidacji w ogóle coś znaczy.

Zasada jest jedna: podział walidacyjny ma odtwarzać relację między danymi treningowymi a przyszłymi danymi produkcyjnymi. Jeśli model będzie przewidywał dla nowych pacjentów, w walidacji pacjenci z foldu walidacyjnego nie mogą występować w treningu. Jeśli będzie prognozował jutro, w walidacji nie może się uczyć na pojutrze.

Mechanizm — dlaczego tak działa

KFold dzieli dane na k kolejnych bloków. Bez tasowania bierze je w kolejności z pliku — a pliki bywają posortowane według klasy, daty czy źródła. Wtedy fold walidacyjny może zawierać klasę, której prawie nie ma w treningu. Z tasowaniem (shuffle=True) proporcje są losowo podobne, ale przy małych lub niezbalansowanych danych mogą się wyraźnie rozjechać.

StratifiedKFold rozwiązuje to wprost: w każdym foldzie proporcje klas są (prawie) takie jak w całym zbiorze. Zmniejsza to wariancję oszacowania i chroni przed foldem bez przykładów rzadkiej klasy, w którym precyzja czy recall są nieokreślone. To rozsądny domyślny wybór dla klasyfikacji.

GroupKFold dba, by wszystkie obserwacje z jednej grupy (pacjenta, użytkownika, zdjęcia z tej samej sesji, nagrania tego samego mówcy) trafiły do jednego foldu. Bez tego model może nauczyć się rozpoznawać osobę zamiast zjawiska — a w walidacji ta sama osoba czeka na niego po drugiej stronie, więc wynik jest zawyżony. StratifiedGroupKFold łączy oba wymagania.

TimeSeriesSplit trenuje zawsze na przeszłości i waliduje na kolejnym okresie, z rosnącym oknem treningowym. Losowe foldy na szeregu czasowym pozwalają modelowi „interpolować” między sąsiednimi dniami, co w produkcji jest niemożliwe. Parametr gap dodaje przerwę między treningiem a walidacją, gdy cechy używają opóźnień lub średnich kroczących.

Leave-one-out (k = n) jest prawie nieobciążone, ale kosztowne i ma dużą wariancję; przydaje się głównie przy bardzo małych zbiorach i modelach z tanim wzorem na wynik LOO (np. regresja liniowa). Powtarzana CV (RepeatedStratifiedKFold) uśrednia także losowość samego podziału na foldy.

Na przykładzie

Zbiór Wine (178 win) jest w scikit-learn posortowany według odmiany: najpierw 59 win klasy 0, potem 71 klasy 1, na końcu 48 klasy 2. Regresja logistyczna ze standaryzacją oceniona KFold(3) bez tasowania dała foldy o trafności 1,7%, 71,2% i 16,9%, średnio 29,9% — gorzej niż zgadywanie najczęstszej klasy (39,9%). Przyczyna: pierwszy fold walidacyjny zawierał 59 win klasy 0 i jedno klasy 1, a w treningu nie było ani jednego wina klasy 0. Model nie mógł przewidzieć klasy, której nigdy nie widział.

Ten sam model z StratifiedKFold(3, shuffle=True, random_state=0): foldy 98,3%, 96,6% i 98,3%, średnio 97,8%. Dane, model i k są identyczne; zmienił się tylko sposób podziału. Przy pięciu foldach bez tasowania efekt jest słabszy (96,1% wobec 98,3% dla wersji warstwowej), bo każdy trening zawiera wtedy już wszystkie klasy — ale pierwszy fold walidacyjny nadal składa się wyłącznie z klasy 0.

Ta ilustracja działa w przeglądarce z włączonym JavaScriptem: walidacja krzyżowa kNN na winach ułożonych w pliku klasami: bez tasowania 3 foldy dają średnio 23%, bo całe klasy wypadają z treningu, a po potasowaniu ta sama metoda ma 96%.

Dane: Wine (wina z Piemontu)

W praktyce

  • Klasyfikacja bez struktury: StratifiedKFold(n_splits=5, shuffle=True, random_state=0). Uwaga: przekazanie cv=5 daje warstwowanie, ale bez tasowania.
  • Obserwacje zgrupowane: GroupKFold lub StratifiedGroupKFold z argumentem groups= w cross_val_score; grupy to np. identyfikator pacjenta.
  • Szeregi czasowe: TimeSeriesSplit(n_splits=5, gap=...); nigdy nie tasuj.
  • Regresja z mocno skośnym celem: można warstwować po kwantylach celu (np. pd.qcut) i podać je jako y do StratifiedKFold.split.
  • Małe zbiory: RepeatedStratifiedKFold(n_splits=5, n_repeats=10); leave-one-out tylko, gdy naprawdę brakuje danych.
  • Typowy błąd na Kaggle: lokalna CV losowa, gdy test konkursu to inny okres lub inne grupy — wynik lokalny nie przewiduje rankingu.

Najczęstsze pytania

Kiedy GroupKFold jest konieczny?
Gdy jedna jednostka, o którą pyta zadanie, wnosi do danych wiele wierszy: wiele zdjęć tego samego pacjenta, wiele transakcji tego samego klienta, wiele okien z jednego nagrania. Jeśli w produkcji model zobaczy nowe jednostki, walidacja musi to odtworzyć.
Czy stratyfikacja ma sens przy zbalansowanych klasach?
Tak, choć zysk jest mniejszy. Usuwa część wariancji wynikającej z przypadkowych wahań proporcji klas w foldach i nic nie kosztuje.
Dlaczego wynik TimeSeriesSplit jest zwykle gorszy niż losowej CV?
Bo jest uczciwszy. Losowa CV pozwala modelowi korzystać z informacji z przyszłości i z sąsiednich, bardzo podobnych obserwacji. Prognoza naprawdę przyszłych okresów jest trudniejsza i to ją trzeba mierzyć.

Źródła

  • Hastie T., Tibshirani R., Friedman J. „The Elements of Statistical Learning”, 2nd ed., Springer 2009, rozdz. 7.10 (Cross-Validation).
  • Kohavi R. (1995). A Study of Cross-Validation and Bootstrap for Accuracy Estimation and Model Selection. Proceedings of IJCAI 1995.
  • Roberts D. R. i in. (2017). Cross-validation strategies for data with temporal, spatial, hierarchical, or phylogenetic structure. Ecography, 40(8).
  • Dokumentacja scikit-learn: Cross-validation iterators — https://scikit-learn.org/stable/modules/cross_validation.html#cross-validation-iterators

Zobacz też