ML Atlas

10 · Praktyka · 5 min czytania · aktualizacja

Ile danych potrzebuję, żeby wytrenować model uczenia maszynowego?

W skrócie

Nie ma uniwersalnej liczby. Odpowiada krzywa uczenia: gdy wynik wciąż rośnie z liczbą przykładów, więcej danych pomoże; gdy stoi, zmień model lub cechy.

Co to jest

Potrzebujesz tylu danych, ile trzeba, żeby krzywa uczenia — wynik na walidacji w funkcji liczby przykładów treningowych — przestała istotnie rosnąć, plus tylu przykładów testowych, żeby zmierzyć wynik z wymaganą precyzją. Nie da się tego wyznaczyć z góry jedną regułą; da się to zmierzyć na już zebranych danych i rozsądnie ekstrapolować.

Popularne reguły kciuka dają rząd wielkości, nie odpowiedź. W regresji logistycznej mówi się o co najmniej 10 przypadkach rzadszej klasy na każdą cechę (Peduzzi i in., 1996), ale późniejsze badania pokazały, że potrzebna liczba zależy od siły sygnału i częstości klasy, a nie tylko od liczby cech. Sieci głębokie od zera potrzebują zwykle tysięcy przykładów na klasę; z transfer learningiem — czasem kilkudziesięciu.

Pytanie ma dwie części, które łatwo pomylić: ile danych, żeby model się dobrze nauczył, i ile danych, żeby wiarygodnie ocenić, czy się nauczył. Druga część bywa bardziej wymagająca.

Mechanizm — dlaczego tak działa

Błąd maleje potęgowo. Dla wielu modeli błąd walidacyjny w funkcji liczby przykładów n zachowuje się w przybliżeniu jak a·n^(−b) + c. Składnik c to błąd, którego żadna ilość danych nie usunie: szum etykiet, brakujące cechy, ograniczenia modelu. Na początku każde podwojenie danych daje dużo, później coraz mniej. Ta sama prawidłowość, zmierzona na ogromną skalę, stoi za prawami skalowania modeli językowych (Hestness i in., 2017; Kaplan i in., 2020).

Złożoność modelu przesuwa krzywą. Model o silnych założeniach (liniowy) uczy się szybko, ale wcześnie dochodzi do swojego sufitu. Model elastyczny (las, boosting, sieć) przy małej próbie przegrywa, bo ma za dużo wariancji, ale jego sufit jest wyżej. Dlatego przy małych danych częściej wygrywają proste modele, a krzywe uczenia różnych modeli potrafią się przecinać.

Co zwiększa zapotrzebowanie. Słaby sygnał i dużo szumu w etykietach. Rzadka klasa — liczy się liczba przypadków rzadkiej klasy, nie wszystkich wierszy. Wiele cech i interakcji. Duża różnorodność warunków (różne szpitale, aparaty, akcenty), które model musi pokryć.

Ocena też kosztuje dane. Trafność zmierzona na n przykładach testowych ma błąd standardowy √(p(1 − p)/n). Przy trafności 90% i 100 przykładach testowych 95-procentowy przedział ma szerokość ±5,9 punktu procentowego; przy 1000 — ±1,9; przy 10 000 — ±0,6. Jeśli chcesz odróżnić model 90% od 91%, potrzebujesz tysięcy przykładów testowych, a dla porównania dwóch modeli na tych samych danych — testu par (np. McNemara).

Na przykładzie

Krzywe uczenia: dla każdej liczby przykładów treningowych 20 losowych podziałów z 25% danych odłożonych do testu (StratifiedShuffleSplit, random_state=0), trafność na teście. Regresja logistyczna ze standaryzacją i las losowy (200 drzew).

Przykładów treningowychDigits: regresja log.Digits: lasTitanic: regresja log.Titanic: las
25–300,6570,5880,7400,728
1000,8650,8620,7780,763
2000,9160,9190,7890,788
4000,9450,9510,7880,801
całość (Digits 1347, Titanic 668)0,9690,9760,7890,817

Na Titanicu regresja logistyczna przestaje się poprawiać już przy 200 pasażerach (0,789 przy 200, 0,788 przy 400 i 0,789 przy 668) — więcej danych tego modelu nie poprawi; pomogłyby lepsze cechy albo bardziej elastyczny model. Las losowy wciąż rośnie (0,788 → 0,801 → 0,817), więc dla niego dodatkowe dane miałyby wartość. Na Digits krzywe przecinają się między 100 a 200 przykładami: przy małej próbie wygrywa model liniowy, przy większej las.

Ekstrapolacja działa zaskakująco dobrze. Dopasowałem krzywą a·n^(−b) + c do błędu lasu na Digits tylko dla 30–400 przykładów. Przewidywana trafność przy 1347 przykładach: 0,976 — dokładnie tyle, ile wyszło. Dla regresji logistycznej prognoza (0,963) była nieco zbyt pesymistyczna wobec rzeczywistego 0,969.

Breast Cancer pokazuje, jak mało czasem trzeba: regresja logistyczna osiąga 0,933 już na 20 guzach i 0,979 na 426 — sygnał jest tak silny, że kilkadziesiąt przykładów daje prawie pełną jakość.

Dane: Titanic Breast Cancer Wisconsin (diagnostyka raka piersi) Digits (ręcznie pisane cyfry 8×8)

W praktyce

Reguła postępowania:

  • Narysuj krzywą uczenia na tym, co masz: learning_curve(model, X, y, train_sizes=np.linspace(0.1, 1.0, 8), cv=StratifiedShuffleSplit(n_splits=20, test_size=0.25, random_state=0)).
  • Krzywa wciąż rośnie → zbieraj dane; dopasuj a n*(-b) + c (scipy.optimize.curve_fit) do błędu i oszacuj, ile przykładów da wymagany wynik.
  • Krzywa płaska, a wynik za słaby → więcej danych nie pomoże; zmień cechy, model albo popraw etykiety.
  • Rzadka klasa → licz przypadki tej klasy; w regresji logistycznej do predykcji klinicznej użyj wzorów Riley i in. (2020) zamiast reguły 10 na cechę.
  • Planuj zbiór testowy osobno: szerokość przedziału ≈ 1.96 np.sqrt(p (1 - p) / n_test); dla precyzji ±2 punkty przy p ≈ 0,9 to ok. 900 przykładów.
  • Mało etykiet → transfer learning z wytrenowanego modelu, augmentacja danych albo uczenie półnadzorowane, zanim zaczniesz drogie etykietowanie.

Najczęstsze pytania

Czy 1000 przykładów wystarczy?
Dla prostego problemu tabelarycznego z wyraźnym sygnałem często tak, dla rozpoznawania 100 klas obrazów od zera — zdecydowanie nie. Sama liczba nic nie mówi bez siły sygnału, liczby klas, częstości rzadkiej klasy i złożoności modelu. Krzywa uczenia na pierwszych kilkuset przykładach powie więcej niż jakakolwiek reguła.
Czy lepiej mieć więcej danych, czy lepszy model?
Zależy od kształtu krzywej. Gdy krzywa rośnie, dane zwykle dają więcej niż strojenie. Gdy jest płaska, nowe dane są marnotrawstwem, a zysk przyniesie lepszy model albo lepsze cechy. W wielu projektach największą poprawę daje poprawienie etykiet, nie ich dokładanie.
Ile danych potrzeba do fine-tuningu modelu językowego?
Znacznie mniej niż do treningu od zera, bo model już zna język. Do nauczenia formatu lub stylu wystarczają często setki dobrych przykładów, do nowej wiedzy dziedzinowej — tysiące, choć wiedzę zwykle lepiej dostarczać przez wyszukiwanie (RAG).

Źródła

  • Peduzzi P., Concato J., Kemper E., Holford T. R., Feinstein A. R. „A simulation study of the number of events per variable in logistic regression analysis”, Journal of Clinical Epidemiology 49(12), 1996, s. 1373–1379.
  • Riley R. D. i in. „Calculating the sample size required for developing a clinical prediction model”, BMJ 368, 2020, m441.
  • Hestness J. i in. „Deep Learning Scaling is Predictable, Empirically”, arXiv:1712.00409, 2017.
  • Kaplan J. i in. „Scaling Laws for Neural Language Models”, arXiv:2001.08361, 2020.
  • Perlich C., Provost F., Simonoff J. S. „Tree Induction vs. Logistic Regression: A Learning-Curve Analysis”, Journal of Machine Learning Research 4, 2003, s. 211–255.
  • Dokumentacja scikit-learn, „Validation curves: plotting scores to evaluate models”: https://scikit-learn.org/stable/modules/learning_curve.html

Zobacz też