ML Atlas

10 · Praktyka · 5 min czytania · aktualizacja

Jak wybrać odpowiedni model uczenia maszynowego do swoich danych?

W skrócie

Nie ma modelu najlepszego zawsze. Zacznij od punktu odniesienia i modelu liniowego, dodaj las lub boosting, a wybór rozstrzygnij walidacją krzyżową.

Co to jest

Model wybiera się w trzech krokach: typ danych zawęża rodzinę (tabela — modele liniowe i zespoły drzew; obrazy — sieci splotowe lub gotowe modele wizyjne; tekst — wstępnie wytrenowane transformery), ograniczenia praktyczne (interpretowalność, czas odpowiedzi, ilość danych) zawężają ją dalej, a ostateczny wybór rozstrzyga walidacja krzyżowa kilku kandydatów na Twoich danych. Twierdzenie „no free lunch” mówi, że żaden algorytm nie jest najlepszy na wszystkich problemach — i eksperymenty to potwierdzają.

Najczęstszy błąd to zaczynanie od najmodniejszego modelu. Lepsza kolejność: model trywialny (żeby wiedzieć, ile wart jest jakikolwiek sygnał), prosty model liniowy (żeby wiedzieć, ile daje nieliniowość), a dopiero potem modele złożone. Jeśli las losowy jest lepszy od regresji logistycznej o pół punktu, zwykle wygrywa regresja logistyczna — jest tańsza, stabilniejsza i łatwiejsza do wyjaśnienia.

Wybór modelu jest też mniej ważny, niż się wydaje. Lepsze cechy, czystsze etykiety i więcej danych zwykle dają więcej niż zamiana jednego dobrego algorytmu na inny.

Mechanizm — dlaczego tak działa

Każdy model to założenie o świecie. Regresja liniowa zakłada addytywne, liniowe efekty. kNN — że podobne przykłady mają podobne etykiety w sensie wybranej odległości. Drzewa — że świat da się opisać progami i interakcjami. Sieci splotowe — że ważne wzorce są lokalne i mogą pojawić się w dowolnym miejscu obrazu. Model wygrywa tam, gdzie jego założenia pasują do danych, i to jest treść twierdzenia „no free lunch”.

Rozmiar danych zmienia ranking. Przy setkach przykładów wygrywają modele o silnych założeniach i małej wariancji: liniowe, SVM, naiwny Bayes. Przy dziesiątkach tysięcy wierszy danych tabelarycznych przewagę zyskuje boosting drzew, który może wykorzystać złożone interakcje bez przeuczenia (Grinsztajn i in., 2022). Przy milionach obrazów czy tekstów — sieci głębokie.

Struktura cech. Cechy heterogeniczne (wiek, kategoria, cena w różnych jednostkach), z brakami i progami — to teren drzew. Cechy jednorodne i gęste (piksele, sygnały, osadzenia) — teren modeli liniowych, SVM i sieci. Kategorie o tysiącach wartości — boosting z natywną obsługą kategorii lub osadzenia.

Ograniczenia pozamodelowe. Wymóg wyjaśnienia każdej decyzji (kredyty, medycyna) faworyzuje modele liniowe i płytkie drzewa. Czas odpowiedzi rzędu milisekund na słabym sprzęcie wyklucza duże zespoły. Mało etykiet, a dużo nieoznaczonych danych — kierunek: transfer learning lub uczenie samonadzorowane.

Dlaczego walidacja, a nie intuicja. Różnice między dobrymi modelami są często mniejsze niż szum oszacowania. Porównanie trzeba robić na tych samych podziałach, z powtórzeniami, i patrzeć na rozrzut — inaczej wybieramy model, który miał szczęście na jednym podziale.

Na przykładzie

Dziewięć modeli z ustawieniami domyślnymi scikit-learn (modele wrażliwe na skalę ze standaryzacją w potoku), sześć zbiorów, powtarzana walidacja krzyżowa 5 × 5, trafność:

ModelTitanicIrisPenguinsWineBreast CancerDigitsŚrednie miejsce
Model trywialny0,6160,3330,4380,3990,6270,1019,0
Regresja logistyczna0,7960,9560,9860,9820,9770,9693,2
Naiwny Bayes0,7810,9550,9690,9720,9380,8426,3
kNN0,8000,9530,9840,9620,9670,9764,6
SVM (RBF)0,8250,9590,9790,9830,9740,9811,8
Drzewo decyzyjne0,7850,9480,9650,9030,9280,8567,5
Las losowy0,8160,9490,9770,9790,9600,9774,5
HistGradientBoosting0,8200,9470,9680,9740,9660,9715,3
Sieć MLP0,8100,9530,9870,9800,9750,9802,8

Na tych małych, czystych, liczbowych zbiorach najlepiej wypada SVM ze standaryzacją (pierwsze miejsce na czterech z sześciu), a regresja logistyczna jest trzecia w średnim rankingu. Boosting, który dominuje w konkursach na dużych danych tabelarycznych, jest tu dopiero szósty. Nie przeczy to benchmarkom — pokazuje, że ranking zależy od rodzaju i rozmiaru danych.

Drugi wniosek: różnice w czołówce są małe. Na Breast Cancer pierwsze pięć modeli mieści się w przedziale 0,966–0,977, a odchylenie między częściami walidacji wynosi ok. 0,015. Największe różnice dzieli model trywialny od reszty i pojedyncze drzewo od zespołów.

Dane: Titanic Iris (irysy Fishera) Breast Cancer Wisconsin (diagnostyka raka piersi) Wine (wina z Piemontu) Digits (ręcznie pisane cyfry 8×8)

W praktyce

Reguła wyboru:

  • Zawsze najpierw punkt odniesienia: DummyClassifier() / DummyRegressor() oraz make_pipeline(StandardScaler(), LogisticRegression()) lub RidgeCV().
  • Dane tabelaryczne, mieszane typy cech, braki → HistGradientBoostingClassifier() i RandomForestClassifier(n_estimators=500); dla setek wierszy dorzuć SVC() ze standaryzacją.
  • Obrazy, dźwięk, tekst → nie trenuj od zera; weź wytrenowany model (np. torchvision.models.resnet18(weights="DEFAULT") albo transformer z Hugging Face) i dostrój go.
  • Wymagana interpretowalność → regresja logistyczna z regularyzacją albo płytkie drzewo; złożony model tylko wtedy, gdy zysk jest wyraźnie większy niż szum.
  • Porównuj na tych samych podziałach: cv = RepeatedStratifiedKFold(n_splits=5, n_repeats=5, random_state=0) i cross_val_score(m, X, y, cv=cv) dla każdego kandydata; patrz na średnią i odchylenie.
  • Wybrany model sprawdź raz na odłożonym zbiorze testowym. Strojenie i wybór na tym samym zbiorze zawyżają wynik.

Najczęstsze pytania

Czy jest jeden model, od którego zawsze warto zacząć?
Dla danych tabelarycznych bezpiecznym startem jest para: regresja logistyczna (lub grzbietowa) i las losowy albo boosting. Pierwszy model mówi, ile da się osiągnąć liniowo, drugi — ile dodają nieliniowości i interakcje. Różnica między nimi podpowiada, gdzie szukać dalej.
Kiedy sieć neuronowa jest lepsza od drzew?
Gdy dane są jednorodne i mają strukturę (obrazy, dźwięk, tekst, sekwencje) albo gdy jest ich bardzo dużo. Na typowych danych tabelarycznych średniej wielkości zespoły drzew wciąż zwykle wygrywają lub remisują przy znacznie mniejszym nakładzie pracy.
Ile modeli warto porównać?
Kilka sensownie różnych wystarczy: liniowy, zespół drzew, ewentualnie SVM lub kNN. Porównywanie dziesiątek modeli i setek konfiguracji na małym zbiorze zwiększa ryzyko, że wygra ten, który miał szczęście na walidacji.

Źródła

  • Wolpert D. H. „The Lack of A Priori Distinctions Between Learning Algorithms”, Neural Computation 8(7), 1996, s. 1341–1390.
  • Fernández-Delgado M., Cernadas E., Barro S., Amorim D. „Do we Need Hundreds of Classifiers to Solve Real World Classification Problems?”, Journal of Machine Learning Research 15, 2014, s. 3133–3181.
  • Grinsztajn L., Oyallon E., Varoquaux G. „Why do tree-based models still outperform deep learning on typical tabular data?”, NeurIPS 2022 (Datasets and Benchmarks Track).
  • Kuhn M., Johnson K. „Applied Predictive Modeling”, Springer 2013, rozdz. 2 i 4.
  • Dokumentacja scikit-learn, „Choosing the right estimator”: https://scikit-learn.org/stable/machine_learning_map.html

Zobacz też