ML Atlas

02 · Dane · 4 min czytania · aktualizacja

Czym są wagi klas i kiedy użyć class_weight='balanced'?

W skrócie

Wagi klas mnożą stratę przykładu przez współczynnik zależny od klasy, zwykle odwrotność jej częstości. Pomyłka na klasie rzadkiej kosztuje w treningu więcej.

Co to jest

Wagi klas (class weights) to współczynniki, przez które mnoży się stratę każdego przykładu w zależności od jego klasy. Najczęstszy wybór — „balanced” — to waga proporcjonalna do odwrotności częstości klasy, tak by każda klasa miała w sumie ten sam udział w stracie.

Stosuje się je w regresji logistycznej, sieciach, SVM, drzewach i boostingu, gdy klasy są niezbalansowane, a metryka (np. balanced accuracy) traktuje je równo albo gdy pomyłki na różnych klasach mają różne koszty.

Mechanizm — dlaczego tak działa

Strata treningowa to suma po przykładach. Przy klasie A z 900 przykładami i klasie B ze 100 gradient straty jest w 90% „głosem” klasy A; model minimalizujący stratę uczy się przede wszystkim nie mylić A i przy niejasnych cechach wybiera A. Waga wₖ = N / (K · Nₖ) — wzór scikit-learn dla "balanced" — daje A wagę 0,56, a B wagę 5,0. Każdy przykład B liczy się wtedy dziewięć razy bardziej, a obie klasy wnoszą po połowie straty i gradientu. Dla modelu to tak, jakby dane miały równe proporcje klas — efekt równoważny nadpróbkowaniu rzadkiej klasy, ale bez powielania wierszy.

Skutek: model przesuwa prawdopodobieństwa w stronę rzadkich klas już w treningu. W ujęciu bayesowskim zmienia prior z częstości treningowych na jednostajny, więc jego prawdopodobieństwa przestają odpowiadać prawdziwym proporcjom — są „skalibrowane do świata, w którym klasy są równoliczne”. Przy metryce BA to pożądane; przy log-loss albo gdy potrzebne są prawdziwe prawdopodobieństwa — nie.

Stąd ważne zastrzeżenie: wagi klas i korekta progu decyzji po treningu robią to samo w dwóch różnych miejscach. Zastosowane naraz przesuwają decyzje za daleko: model już faworyzuje klasę rzadką, a próg faworyzuje ją po raz drugi, więc rzadka klasa zalewa przykłady z klas częstych. Wybierz jedno albo strój oba razem na walidacji.

Wagi mogą też wyrażać koszty pomyłek (cost-sensitive learning): jeśli przeoczenie klasy B kosztuje dziesięć razy więcej niż fałszywy alarm, waga 10 na B jest właściwym wyborem niezależnie od częstości.

Na przykładzie

Z Breast Cancer Wisconsin zbudowaliśmy zbiór niezbalansowany: wszystkie 357 guzów łagodnych i losowe 40 złośliwych (seed 0), czyli 10% klasy rzadkiej. Wagi „balanced” wynoszą tu 0,56 dla łagodnych i 4,96 dla złośliwych. Regresja logistyczna na dwóch słabych cechach (średnia tekstura i gładkość) była oceniana 5-krotną walidacją krzyżową.

Bez wag model rozpoznał 10% guzów złośliwych (BA 0,54). Z wagami — 73% złośliwych i 73% łagodnych (BA 0,73). Dokładnie ten sam wynik dało trenowanie bez wag i obniżenie progu do 0,1. Oba zabiegi naraz okazały się szkodliwe: recall złośliwych 100%, ale łagodnych tylko 10%, BA spadła do 0,55. Wagi zmieniły też prawdopodobieństwa: średnie przewidywane ryzyko złośliwości wzrosło z 0,10 (zgodnie z rzeczywistym udziałem) do 0,38.

Dane: Breast Cancer Wisconsin (diagnostyka raka piersi)

W praktyce

  • scikit-learn: class_weight="balanced" w LogisticRegression, SVC, DecisionTreeClassifier, RandomForestClassifier; słownik {0: 1, 1: 5} dla własnych kosztów. MLPClassifier nie ma tego parametru; wiele estymatorów przyjmuje sample_weight w fit.
  • PyTorch: nn.CrossEntropyLoss(weight=...) z wagą na klasę; pos_weight w BCEWithLogitsLoss dla dwóch klas.
  • XGBoost: scale_pos_weight (stosunek liczności klas) dla dwóch klas, sample_weight wieloklasowo; LightGBM: class_weight="balanced" lub is_unbalance=True.
  • Zacznij od wag „balanced”, sprawdź BA i macierz pomyłek; jeśli rzadka klasa zalewa resztę, złagodź wagi (np. pierwiastek z odwrotności częstości).
  • Typowy błąd: class_weight="balanced" plus obniżony próg „dla pewności” — podwójna korekta.

Najczęstsze pytania

Jak działa class_weight='balanced' w scikit-learn?
Każda klasa dostaje wagę N / (K · Nₖ): liczba wszystkich przykładów podzielona przez liczbę klas i liczność klasy. W zadaniu dwuklasowym klasa z 10% danych dostaje wagę 5, klasa z 90% — 0,56. W sumie obie klasy ważą w stracie tyle samo.
Wagi klas czy oversampling (SMOTE)?
Wagi „balanced” są w oczekiwaniu równoważne powieleniu wierszy rzadkiej klasy, a są tańsze i nie tworzą duplikatów. SMOTE generuje syntetyczne punkty między istniejącymi, co czasem pomaga, ale może tworzyć nierealne przykłady. Zacznij od wag, SMOTE sprawdź na walidacji.
Czy wagi klas psują prawdopodobieństwa modelu?
Tak: model zachowuje się, jakby klasy były równoliczne, więc zawyża prawdopodobieństwo klasy rzadkiej, jak w przykładzie (0,38 zamiast 0,10). Jeśli potrzebujesz prawdziwych prawdopodobieństw, trenuj bez wag i koryguj próg albo skalibruj model po treningu.

Źródła

  • He, H., Garcia, E. A. (2009). "Learning from imbalanced data". IEEE Transactions on Knowledge and Data Engineering 21(9), 1263–1284. doi:10.1109/TKDE.2008.239
  • Elkan, C. (2001). "The foundations of cost-sensitive learning". IJCAI, 973–978.
  • King, G., Zeng, L. (2001). "Logistic regression in rare events data". Political Analysis 9(2), 137–163.
  • Buda, M., Maki, A., Mazurowski, M. A. (2018). "A systematic study of the class imbalance problem in convolutional neural networks". Neural Networks 106, 249–259. arXiv:1710.05381
  • scikit-learn: sklearn.utils.class_weight.compute_class_weight. https://scikit-learn.org/stable/modules/generated/sklearn.utils.class_weight.compute_class_weight.html

Zobacz też