ML Atlas

04 · Ocena · 4 min czytania · Interaktywne · aktualizacja

Co pokazuje krzywa walidacyjna i jak dobrać na niej złożoność modelu?

W skrócie

Krzywa walidacyjna pokazuje wynik treningowy i walidacyjny w zależności od jednego hiperparametru. Widać, gdzie model jest za prosty, a gdzie przeuczony.

Co to jest

Krzywa walidacyjna (validation curve) to wykres wyniku modelu na danych treningowych i walidacyjnych w funkcji wartości jednego hiperparametru — np. głębokości drzewa, siły regularyzacji, liczby sąsiadów w kNN albo szerokości jądra w SVM. Dla każdej wartości trenuje się model (zwykle w walidacji krzyżowej) i zapisuje oba wyniki. Wykres pokazuje, przy jakiej złożoności model działa najlepiej i jak szybko się psuje po obu stronach optimum.

Jest to bezpośrednia wizualizacja kompromisu między niedouczeniem a przeuczeniem. Po stronie „prostej” obie krzywe są nisko. Po stronie „złożonej” krzywa treningowa idzie w górę, a walidacyjna spada. Optimum leży tam, gdzie krzywa walidacyjna osiąga szczyt.

Mechanizm — dlaczego tak działa

Większość hiperparametrów steruje efektywną pojemnością modelu: jak skomplikowane funkcje potrafi on wyrazić przy danych treningowych. Gdy pojemności przybywa, model dopasowuje trening coraz lepiej — krzywa treningowa rośnie mniej więcej monotonicznie. Wynik walidacyjny najpierw rośnie, bo model zaczyna łapać prawdziwą strukturę, a potem spada, bo zaczyna łapać szum, który w nowych danych się nie powtarza. Różnica między krzywymi to miara przeuczenia.

Kierunek osi bywa mylący. Dla max_depth czy liczby neuronów większa wartość = większa złożoność. Dla siły regularyzacji alpha czy liczby sąsiadów k jest odwrotnie: większa wartość = prostszy model. Parametr C w SVM i regresji logistycznej jest odwrotnością siły regularyzacji, więc duże C = złożony model. Przy gamma w jądrze RBF duże wartości oznaczają wąskie, lokalne jądra, które otaczają pojedyncze punkty.

Wiele hiperparametrów działa w skali multiplikatywnej, dlatego wartości sprawdza się na siatce logarytmicznej (0,001, 0,01, 0,1…), a nie liniowej. Liniowa siatka marnuje większość punktów w nieciekawym obszarze.

Ograniczenia: krzywa walidacyjna bada jeden hiperparametr przy ustalonych pozostałych, a optimum jednego często zależy od drugiego (w SVM optymalne gamma zależy od C). Do strojenia wielu parametrów naraz służy przeszukiwanie siatki lub losowe. Ponadto szczyt krzywej walidacyjnej jest optymistyczny — wybierając maksimum spośród wielu wartości, wybieramy też szczęście — więc do raportowania potrzebny jest osobny test lub zagnieżdżona walidacja krzyżowa. Gdy kilka wartości daje wynik w granicach szumu, rozsądnie jest wybrać najprostszy model z tej grupy.

Na przykładzie

Na zbiorze Digits (1797 cyfr 8 × 8) policzyłem krzywą walidacyjną SVM z jądrem RBF dla parametru gamma od 0,00001 do 0,1, z 5-krotną warstwową walidacją krzyżową (validation_curve, StratifiedKFold(5, shuffle=True, random_state=0), C = 1). Przy gamma = 0,00001 jądro jest tak szerokie, że model jest niemal liniowy: 91,5% na treningu, 90,7% na walidacji — obie krzywe nisko, czyli niedouczenie. Przy gamma = 0,001 walidacja osiąga szczyt 99,0% (trening 99,9%).

Dalej krzywe się rozchodzą. Przy gamma = 0,01 trening ma już 100%, a walidacja spada do 83,0%; przy 0,03 — do 21,9%, przy 0,1 — do 10,7%, czyli poziomu zgadywania jednej z dziesięciu cyfr. Każdy przykład treningowy otacza wtedy własne wąskie jądro: model idealnie pamięta trening, ale dla nowej cyfry, odległej od wszystkich zapamiętanych, nie ma żadnej informacji. Odchylenie standardowe walidacji w foldach przy optimum to zaledwie 0,4 punktu, a gamma = 0,0003 i 0,003 dają 98,5% i 98,6% — szczyt jest szeroki.

Ta ilustracja działa w przeglądarce z włączonym JavaScriptem: wielomiany stopnia 0–15 dopasowane do 20 zaszumionych punktów sinusoidy: błąd treningowy stale spada, a błąd na 200 nowych punktach najpierw maleje, potem rośnie lawinowo.

Dane: Digits (ręcznie pisane cyfry 8×8)

W praktyce

  • validation_curve(model, X, y, param_name='svc__gamma', param_range=np.logspace(-5, -1, 9), cv=5); w Pipeline nazwa parametru ma prefiks kroku.
  • Wykres: ValidationCurveDisplay.from_estimator (scikit-learn 1.3+); dorysuj pasma odchylenia z foldów.
  • Typowe parametry do zbadania: max_depth i min_samples_leaf drzew, C i gamma SVM, alpha w Ridge/Lasso, n_neighbors w kNN, liczba epok lub weight_decay w sieciach.
  • Wybieraj najprostszy model, którego wynik mieści się w jednym odchyleniu standardowym od najlepszego (reguła jednego błędu standardowego).
  • Jeśli szczyt wypada na krańcu przeszukiwanego zakresu, poszerz zakres — prawdziwe optimum może leżeć dalej.
  • Typowy błąd: odczytywanie wartości walidacji w szczycie jako oczekiwanej jakości modelu.

Najczęstsze pytania

Czym krzywa walidacyjna różni się od przeszukiwania siatki?
Krzywa walidacyjna to przeszukiwanie jednego wymiaru z wizualizacją także wyniku treningowego, co pozwala zobaczyć, *dlaczego* dana wartość wygrywa. Przeszukiwanie siatki obejmuje wiele hiperparametrów naraz, ale zwykle zwraca tylko najlepszą kombinację.
Dlaczego krzywa treningowa nie zawsze rośnie monotonicznie?
Bo trening też zawiera losowość (inicjalizacja, kolejność przykładów, losowanie cech w lasach) i bo niektóre parametry nie sterują pojemnością w prosty sposób. Niewielkie wahania są normalne; liczy się ogólny trend.
Co oznacza krzywa walidacyjna płaska w szerokim zakresie?
Że model nie jest wrażliwy na ten parametr w tym zakresie. To dobra wiadomość: wybór jest bezpieczny, a czas strojenia lepiej poświęcić innym parametrom lub cechom.

Źródła

  • Hastie T., Tibshirani R., Friedman J. „The Elements of Statistical Learning”, 2nd ed., Springer 2009, rozdz. 7.10 (Cross-Validation — reguła jednego błędu standardowego).
  • James G., Witten D., Hastie T., Tibshirani R. „An Introduction to Statistical Learning”, 2nd ed., Springer 2021, rozdz. 9.3 (Support Vector Machines) i 5.1 (Cross-Validation).
  • Géron A. „Hands-On Machine Learning with Scikit-Learn, Keras, and TensorFlow”, 3rd ed., O’Reilly 2022, rozdz. 5 (Support Vector Machines).
  • Dokumentacja scikit-learn: Validation curves and learning curves — https://scikit-learn.org/stable/modules/learning_curve.html

Zobacz też