04 · Ocena · 4 min czytania · Interaktywne · aktualizacja
Po co dzielić dane na zbiór treningowy, walidacyjny i testowy?
W skrócie
Trening służy do uczenia, walidacja do wyboru modelu i ustawień, a test do jednej uczciwej oceny na końcu. Każde podejrzenie testu psuje jego wartość.
Co to jest
Podział na trzy zbiory to podstawowa procedura oceny modeli: na zbiorze treningowym model uczy się parametrów, na walidacyjnym porównujemy warianty i dobieramy ustawienia (hiperparametry, cechy, architekturę), a zbiór testowy odkładamy do jednorazowej, końcowej oceny. Wynik testowy ma odpowiadać na pytanie: jak model poradzi sobie z danymi, których nikt — ani algorytm, ani badacz — wcześniej nie oglądał.
Dlaczego aż trzy, a nie dwa? Bo w praktyce „uczymy się” na dwa sposoby. Algorytm dopasowuje wagi do treningu. My sami dopasowujemy decyzje — wybieramy model, który wypadł najlepiej. Jeśli te decyzje podejmujemy na zbiorze testowym, przestaje on być niezależny i jego wynik staje się zawyżony.
Typowe proporcje to 60/20/20 lub 70/15/15, a przy milionach przykładów nawet 98/1/1 — liczy się bezwzględna liczba przykładów w walidacji i teście, nie procent.
Mechanizm — dlaczego tak działa
Wynik na zbiorze treningowym jest optymistyczny, bo model widział te przykłady i mógł dopasować się do ich szumu. To oczywiste. Mniej oczywiste jest to, że wynik zwycięzcy na walidacji też jest optymistyczny. Jeśli porównujemy dziesięć wariantów, a każdy wynik walidacyjny to „prawdziwa jakość + losowy błąd próbki”, to wygrywa ten, który ma dobrą jakość albo szczęśliwy błąd. Wybierając maksimum, systematycznie wybieramy też szczęście. To zjawisko nazywa się klątwą zwycięzcy albo obciążeniem selekcji.
Zbiór testowy przerywa ten mechanizm: model wybrany na walidacji oceniamy na nowych danych, gdzie szczęście z selekcji już nie działa. Dlatego test musi być użyty raz. Gdy po zobaczeniu wyniku testowego wracamy do poprawek i testujemy znowu, test staje się drugą walidacją, a jego wynik zaczyna rosnąć z powodów, które nie przeniosą się na prawdziwe dane.
Drugim warunkiem jest niezależność i reprezentatywność. Wszystkie trzy zbiory muszą pochodzić z tego samego rozkładu co dane, na których model będzie pracował, a żadna informacja z walidacji ani testu nie może przeciec do treningu. Typowe przecieki: skalowanie lub imputacja dopasowane na całych danych przed podziałem, ten sam pacjent lub klient w treningu i teście, przyszłość w treningu przy prognozowaniu szeregów czasowych.
Ograniczenie: pojedynczy podział jest zaszumiony, zwłaszcza przy małych danych. Przy kilkuset przykładach wynik walidacyjny potrafi zmienić się o kilka punktów procentowych od samego losowania podziału. Wtedy walidację zastępuje się walidacją krzyżową na części treningowej, a test i tak trzyma się osobno.
Na przykładzie
Zbiór Wine ma 178 win z trzech odmian i 13 cech chemicznych. Podzieliłem go warstwowo na 106 przykładów treningowych, 36 walidacyjnych i 36 testowych (train_test_split dwa razy, random_state=0) i dobierałem liczbę sąsiadów k w klasyfikatorze kNN na surowych, nieskalowanych cechach. Najlepszy na walidacji okazał się k = 1 z trafnością 83,3%. Ten sam model na teście: 63,9%. Spadek o prawie 20 punktów to klątwa zwycięzcy w czystej postaci — przy 36 przykładach walidacyjnych jeden przykład to 2,8 punktu, więc spośród ośmiu kandydatów łatwo wygrywa ten z najszczęśliwszym losowaniem.
Dla porównania kNN po standaryzacji cech (k = 5) uzyskał 97,2% na walidacji i 97,2% na teście. Tu dobry wynik walidacyjny wynikał z prawdziwej jakości, a nie z selekcji, więc test go potwierdził. Morał: wynik walidacyjny służy do wyboru, wynik testowy do raportowania.
Dane: Wine (wina z Piemontu)
W praktyce
- Dziel dane przed jakąkolwiek eksploracją:
train_test_split(X, y, test_size=0.2, stratify=y, random_state=0), a potem drugi podział albo walidacja krzyżowa na części treningowej. - Przy klasyfikacji używaj
stratify=y, by proporcje klas były takie same w każdym zbiorze. - Gdy obserwacje są zgrupowane (pacjent, użytkownik, sesja), dziel po grupach:
GroupShuffleSplit. Dla szeregów czasowych dziel chronologicznie, nigdy losowo. - Całe przetwarzanie (skalowanie, imputacja, kodowanie) zamykaj w
Pipeline, żeby dopasowywało się wyłącznie na treningu. - W PyTorch walidację licz w trybie
model.eval()i podtorch.no_grad(); wczesne zatrzymanie opieraj na walidacji, nie na teście. - Typowy błąd: wielokrotne „zaglądanie” do testu i poprawianie modelu. Jeśli to się stało, uczciwie nazwij ten zbiór walidacyjnym i zdobądź nowy test.
Najczęstsze pytania
- Czy zbiór walidacyjny jest potrzebny, jeśli robię walidację krzyżową?
- Walidacja krzyżowa zastępuje pojedynczy zbiór walidacyjny — każdy fold po kolei pełni jego rolę. Nadal jednak potrzebny jest osobny test, jeśli chcesz raportować wynik modelu wybranego dzięki tej walidacji. Alternatywą jest zagnieżdżona walidacja krzyżowa.
- Czy po ocenie na teście mogę dotrenować model na wszystkich danych?
- Tak, to częsta praktyka: ostateczny model trenuje się na treningu, walidacji i teście razem z wybranymi ustawieniami. Raportowany wynik pochodzi wtedy z wcześniejszej oceny i jest szacunkiem dla tej procedury, a nie dla dokładnie tego egzemplarza modelu.
- Jak duży powinien być zbiór testowy?
- Na tyle duży, by niepewność wyniku była mniejsza niż różnice, które Cię interesują. Przy trafności około 90% i 100 przykładach testowych błąd standardowy to około 3 punkty procentowe, przy 1000 przykładach około 1 punkt.
Źródła
- James G., Witten D., Hastie T., Tibshirani R. „An Introduction to Statistical Learning”, 2nd ed., Springer 2021, rozdz. 5.1 (Cross-Validation).
- Hastie T., Tibshirani R., Friedman J. „The Elements of Statistical Learning”, 2nd ed., Springer 2009, rozdz. 7.2 (Bias, Variance and Model Complexity).
- Cawley G. C., Talbot N. L. C. (2010). On Over-fitting in Model Selection and Subsequent Selection Bias in Performance Evaluation. Journal of Machine Learning Research, 11.
- Kaufman S., Rosset S., Perlich C., Stitelman O. (2012). Leakage in Data Mining: Formulation, Detection, and Avoidance. ACM Transactions on Knowledge Discovery from Data, 6(4).
- Dokumentacja scikit-learn: Cross-validation: evaluating estimator performance — https://scikit-learn.org/stable/modules/cross_validation.html