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ść:
| Model | Titanic | Iris | Penguins | Wine | Breast Cancer | Digits | Średnie miejsce |
|---|---|---|---|---|---|---|---|
| Model trywialny | 0,616 | 0,333 | 0,438 | 0,399 | 0,627 | 0,101 | 9,0 |
| Regresja logistyczna | 0,796 | 0,956 | 0,986 | 0,982 | 0,977 | 0,969 | 3,2 |
| Naiwny Bayes | 0,781 | 0,955 | 0,969 | 0,972 | 0,938 | 0,842 | 6,3 |
| kNN | 0,800 | 0,953 | 0,984 | 0,962 | 0,967 | 0,976 | 4,6 |
| SVM (RBF) | 0,825 | 0,959 | 0,979 | 0,983 | 0,974 | 0,981 | 1,8 |
| Drzewo decyzyjne | 0,785 | 0,948 | 0,965 | 0,903 | 0,928 | 0,856 | 7,5 |
| Las losowy | 0,816 | 0,949 | 0,977 | 0,979 | 0,960 | 0,977 | 4,5 |
| HistGradientBoosting | 0,820 | 0,947 | 0,968 | 0,974 | 0,966 | 0,971 | 5,3 |
| Sieć MLP | 0,810 | 0,953 | 0,987 | 0,980 | 0,975 | 0,980 | 2,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()orazmake_pipeline(StandardScaler(), LogisticRegression())lubRidgeCV(). - Dane tabelaryczne, mieszane typy cech, braki →
HistGradientBoostingClassifier()iRandomForestClassifier(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)icross_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