Jak zbudować klasyfikator na drzewie decyzyjnym od zera
Drzewo decyzyjne przewiduje wynik, zadając kolejne pytania o wiersz danych. Zbudujemy je w czystym Pythonie na pięciu pacjentach, sprawdzając każdy podział ręcznie.
Dobrym przykładem jest badanie chorób serca z Cleveland, gdzie każdy wiersz to pacjent, a ostatnia kolumna to to, co chcemy przewidzieć:
| age | sex | chest_pain | cholesterol | max_heart_rate | vessels | disease |
|---|---|---|---|---|---|---|
| 63 | 1 | typical | 233 | 150 | 0 | No |
| 67 | 1 | asymptomatic | 286 | 108 | 3 | Yes |
| 37 | 1 | nonanginal | 250 | 187 | 0 | No |
Ostatnia kolumna — podbarwiona powyżej — zawiera zaobserwowane zdarzenie: ten pacjent okazał się mieć chorobę serca, tamten okazał się jej nie mieć. Taka kolumna nazywa się etykietą; każda kolumna przed nią opisuje przypadek, a tylko ta mówi, jak przypadek się skończył. To, czego chcemy się nauczyć, to związek między jednym a drugim — z przypadków, w których etykieta jest już znana, tak by dało się go zastosować do przypadków, w których nie jest. Gdy etykieta jest kategorią, jak tutaj, mamy zadanie klasyfikacji.
Drzewa decyzyjne są też modelami składowymi lasów losowych i gradient boostingu. Zrozumienie sposobu wyboru podziałów ułatwia poznanie tych zespołów.
Na danych tabelarycznych pojawiają się dwa zadania: klasyfikacja, gdzie odpowiedzią jest kategoria, i regresja, gdzie jest nią liczba — pensja, cena. Drzewo decyzyjne obsługuje oba. Ten artykuł buduje klasyfikator; skieruj ten sam kod na cel liczbowy, zastąp miarę nieczystości wariancją tego celu, a zliczenia etykiet w liściu ich średnią, a otrzymasz model regresyjny przewidujący liczbę — nic więcej w kodzie się nie zmienia, choć sporo zmienia się w powstałym modelu, o czym traktuje artykuł towarzyszący.
Boosting to temat późniejszego artykułu. Ten skupia się na tym, co boosting powtarza setki razy: jak z tabeli wierszy powstaje pojedyncze drzewo decyzyjne — skąd biorą się jego pytania, jak jedno z nich zostaje wybrane spośród pozostałych i kiedy dzielenie się kończy. Piszemy to w czystym Pythonie, bez NumPy i bez scikit-learn, na pięciu wierszach na tyle małych, że każdą liczbę można sprawdzić ręcznie.
Pobierz pełny przykład w Pythonie i uruchom python3 heart_tree.py.
Dwa typy węzłów i obszary, które wycinają
Będziemy pracować na pięciu pacjentach z badania chorób serca z Cleveland. Pełna tabela zawiera 303 pacjentów i trzynaście cech (predyktorów), ale użyjemy tylko pięciu wierszy i dwóch cech plus etykiety:
| # | stress_test | vessels | disease |
|---|---|---|---|
| 1 | normal | 0 | No |
| 2 | fixed | 0 | Yes |
| 3 | reversable | 2 | Yes |
| 4 | reversable | 1 | Yes |
| 5 | fixed | 0 | No |
training_data = [
["normal", 0, "No"],
["fixed", 0, "Yes"],
["reversable", 2, "Yes"],
["reversable", 1, "Yes"],
["fixed", 0, "No"],
]Mając wynik testu wysiłkowego pacjenta i liczbę naczyń, drzewo ma przewidzieć, czy ten pacjent ma chorobę serca — warto więc wiedzieć, co te dwie kolumny zapisują.
stress_test (Thal) zapisuje wynik badania z użyciem talu: prawidłowy, ubytek utrwalony lub odwracalny. W kodzie zachowujemy źródłową pisownię reversable. vessels (Ca) to liczba głównych naczyń, od 0 do 3, uwidocznionych we fluoroskopii, a nie liczba chorych naczyń. Definicje pochodzą z dokumentacji zbioru UCI.
Weźmiemy te pięć wierszy i będziemy dzielić je na coraz mniejsze grupy, używając cech i ich wartości do rozstrzygania, do której grupy trafia każdy wiersz. Sekwencja podziałów rysuje się jako rozgałęziający się diagram — stąd nazwa „drzewo”.
Na tym diagramie są dwa typy węzłów.
Rombami oznaczono węzły decyzyjne. Każdy zawiera pytanie tak/nie i kieruje wiersz gałęzią True albo False. Nasze drzewo w stylu CART zawsze tworzy dwoje dzieci; inne algorytmy mogą używać większej liczby gałęzi.
Prostokąty to liście. Każdy przechowuje liczebności etykiet wierszy treningowych, które do niego trafiły: Yes: 2, Yes: 1, No: 1 albo No: 1. Po podzieleniu przez sumę otrzymujemy oszacowania częstości klas: 100% Yes, 50/50 i 100% No. Tak małe próbki nie ustalają prawdopodobieństwa choroby u nowego pacjenta.
Węzły kierują wiersze do liści, dzieląc przestrzeń wejściową. Dla cech liczbowych próg przecina ją wzdłuż jednej osi. Poniższy rysunek używa osobnego zbioru syntetycznego do pokazania kształtów granic: prosta, gładka krzywa w stylu sieci neuronowej i podziały drzewa są schematyczne; panel k najbliższych sąsiadów jest obliczony z przedstawionych punktów.
Czego więc naprawdę trzeba, żeby coś takiego zbudować? Potrzebne są dwie rzeczy.
- Sposób generowania pytań kandydujących z tabeli, bo drzewo musi je skądś wziąć.
- Sposób oceniania tych kandydatów, żeby dało się wybrać najlepszego — węzeł decyzyjny zawiera jedno pytanie i ani jednego więcej.
Potrzebujemy też reguły stopu. Tutaj kończymy, gdy żaden kandydat nie zmniejsza nieczystości ponad małą tolerancję numeryczną. Może to pozostawić mieszany liść, nawet jeśli dalsza sekwencja podziałów rozdzieliłaby wiersze — przykładem jest XOR. Limit głębokości i minimalna liczba wierszy w liściu dodatkowo ograniczają przeuczenie.
Od wierszy do pytań
Kandydatów tworzymy, łącząc każdą cechę z każdą różną wartością w bieżących wierszach. Reguła kategoryczna sprawdza równość, np. stress_test == fixed. Reguła liczbowa sprawdza próg, np. vessels >= 1, prawdziwy dla wartości 1, 2 i 3.
To parowanie stanowi cały generator, więc warto rozpisać je w całości. Dla naszego zbioru mamy dwie kolumny po trzy różne wartości, więc kończymy z listą sześciu par:
| typ | kolumna | wartość | jakie pytanie tworzy |
|---|---|---|---|
| kategoryczna | stress_test | normal | Is stress_test == normal? |
| kategoryczna | stress_test | fixed | Is stress_test == fixed? |
| kategoryczna | stress_test | reversable | Is stress_test == reversable? |
| liczbowa | vessels | 0 | Is vessels >= 0? |
| liczbowa | vessels | 1 | Is vessels >= 1? |
| liczbowa | vessels | 2 | Is vessels >= 2? |
Nasza implementacja sprawdza jedną cechę naraz i używa obserwowanych wartości jako progów liczbowych. To uproszczenie, nie wymóg CART: scikit-learn używa środków między sąsiednimi różnymi wartościami. Oba podejścia dają te same podziały treningowe, lecz mogą inaczej kierować nowe wartości pomiędzy obserwacjami. CART może też dzielić kategorie na podzbiory; tutaj sprawdzamy jedną kategorię względem pozostałych.
Te sześć to wszystko, o co to drzewo może zapytać — i nie wszystkie trafią do gotowego drzewa. Większość kandydatów jest próbowana, oceniana i odrzucana; tutaj tylko dwa dożywają do bycia pytaniami w drzewie, które zbudujemy, a pozostałe cztery zostają ocenione i odrzucone. A lista jest skończona — nigdy nie dłuższa niż liczba różnych wartości w tabeli — dlatego następny krok może po prostu wypróbować je wszystkie.
W kodzie pytania definiuje jedna mała klasa. Rule przechowuje indeks kolumny i wartość, a jego metoda holds decyduje, które z dwóch porównań zastosować, patrząc na typ tej wartości:
FEATURES = ["stress_test", "vessels", "disease"]
class Rule:
"""One yes/no test: a column, and the value it is compared against."""
def __init__(self, column, value):
self.column = column
self.value = value
def holds(self, row):
observed = row[self.column]
if isinstance(self.value, (int, float)):
return observed >= self.value # numeric: threshold
return observed == self.value # categorical: equality
def __repr__(self):
operator = ">=" if isinstance(self.value, (int, float)) else "=="
return f"Is {FEATURES[self.column]} {operator} {self.value}?"Tę logikę dopasowania da się zaimplementować na kilka sposobów. Nasz to sprawdzenie isinstance wewnątrz holds i to właśnie dzięki niej to drzewo obsługuje kolumnę tekstową i liczbową obok siebie zupełnie bez preprocessingu. Biblioteki nie wszystkie idą tą drogą.
Scikit-learn wymaga wejść liczbowych. Dla cechy bez porządku, takiej jak stress_test, kodowanie one-hot zachowuje kategorie bez narzucania im kolejności:
| stress_test | is_normal | is_fixed | is_reversable | |
|---|---|---|---|---|
| normal | → | 1 | 0 | 0 |
| fixed | → | 0 | 1 | 0 |
| reversable | → | 0 | 0 | 1 |
Drzewo pyta wtedy is_fixed >= 0.5 tam, gdzie nasze pyta stress_test == fixed — ten sam podział, rozłożony na trzy kolumny. Samo 0.5 nie ma znaczenia: kolumna zawiera tylko 0 i 1, więc dowolne cięcie między nimi rozdziela te same wiersze, a sklearn stawia progi w połowie odległości między dwiema sąsiednimi wartościami. Kolumna z czterema kategoriami stałaby się po prostu czterema takimi kolumnami 0/1, o każdej wciąż pytaną przy 0.5 — kodowanie rośnie wszerz, a każde pytanie pozostaje testem tak/nie na jednej wartości.
LightGBM dzieli za to po podzbiorach: testuje grupę kategorii naraz, co wciąż jest jednym pytaniem o jedną kolumnę — różnica polega na tym, że testowaną wartością jest zbiór, a nie pojedyncza kategoria:
ours: Is stress_test == fixed?
theirs: Is stress_test in {normal, reversable}?LightGBM i XGBoost obsługują podziały kategorii na grupy, przeszukując porządek wyznaczony przez statystyki kategorii. CatBoost dla wielu cech stosuje uporządkowane statystyki celu, a dla części cech o niewielu kategoriach kodowanie one-hot. Szczegóły zależą od konfiguracji; opisuje je dokumentacja CatBoost.
Mierzenie różnorodności zbioru danych
Aby wybrać podział, najpierw mierzymy wymieszanie etykiet za pomocą nieczystości Giniego. Następnie obliczamy jej spadek z uwzględnieniem wielkości grup. Kod nazywa go przyrostem; ściśle przyrost informacji zwykle oznacza analogiczny spadek entropii.
Spójrzmy najpierw na nieczystość Giniego, czyli miarę tego, jak wymieszany jest zbiór — jedną liczbę mówiącą, czy rzeczy w nim są jednego rodzaju, czy stanowią mieszankę wielu.
Załóżmy, że masz powiedzieć, który z dwóch zbiorów jest bardziej wymieszany. Jedno spojrzenie na obrazek poniżej wystarczy, by stwierdzić, że drugi zbiór jest bardziej różnorodny: cztery rodzaje zamiast dwóch i rozłożone równiej. Oko rozstrzyga to natychmiast.
Załóżmy teraz, że nie znamy składu żadnego ze zbiorów — żadnych zestawień, żadnej listy rodzajów, tylko możliwość sięgnięcia do środka i wyjęcia czegoś. Czy dałoby się ująć różnorodność liczbą?
Losujemy dwa elementy niezależnie ze zwracaniem, zapisujemy, czy są różnych rodzajów, i powtarzamy. Zwrot pierwszego elementu przed drugim losowaniem pozwala mnożyć prawdopodobieństwa w poniższym rachunku.
Rysunek pokazuje przykładową sekwencję dziesięciu par z każdego zbioru:
Cztery pary lewego zbioru zawierają dwa różne rodzaje; siedem par prawego. Podziel przez liczbę losowań i masz oszacowanie — daszek nad oznacza wartość oszacowaną z próby, w odróżnieniu od policzonej z całej populacji:
Większy udział niezgodnych par sugeruje większą różnorodność. Więcej niezależnych par zwykle poprawia oszacowanie, lecz ani dziesięć, ani sto par nie gwarantuje dokładnego wyniku.
Ogólnie jednak nie trzeba wcale próbkować: gdy wiesz, co zbiór zawiera, odrobina rachunku prawdopodobieństwa daje tę dokładną wartość wprost. Policz szansę, że dwa wyciągnięcia się zgadzają, i odejmij ją od 1.
Weźmy lewy zbiór. Siedem z dziesięciu jego elementów to niebieskie kwadraty, więc pojedyncze wyciągnięcie jest kwadratem z prawdopodobieństwem 0.7 — a prawdopodobieństwo wyciągnięcia dwóch kwadratów pod rząd to . Koła dają . To jedyne dwa sposoby, by wyciągnięcia się zgadzały, więc zgadzają się w przypadków. Ale nam chodzi o przeciwieństwo — jak często dwa wyciągnięcia wracają różne — a ponieważ każde losowanie albo się zgadza, albo nie, to jest to jeden minus szansa zgodności: 0.42.
Prawy zbiór to ten sam rachunek z czterema rodzajami zamiast dwóch:
| rodzaj | udział | oba wyciągnięcia tutaj |
|---|---|---|
| kwadrat | 0.4 | 0.16 |
| koło | 0.3 | 0.09 |
| trójkąt | 0.2 | 0.04 |
| gwiazda | 0.1 | 0.01 |
| zgodność 0.30 |
Dwa wyciągnięcia zgadzają się w 30% przypadków, więc różnią się w 0.70 przypadków — co odpowiada siedmiu na dziesięć z próbkowania, bez ani jednego losowania.
Tę samą logikę można pokazać geometrycznie. Rozłóż każdą uporządkowaną parę wyciągnięć jako komórkę siatki — pierwsze wyciągnięcie w poziomie, drugie w pionie. Dziesięć elementów daje sto komórek, a ta siatka to wszystkie możliwe wyniki:
Ta dokładna wartość ma nazwę. Prawdopodobieństwo, że dwa losowo wyciągnięte ze zbioru elementy są różnych rodzajów, to nieczystość Giniego zbioru, a zapisana wygląda tak:
gdzie to ułamek zbioru należący do rodzaju . Dwie połowy to dwa sposoby powiedzenia tego samego: to prawdopodobieństwo, że losowania się zgadzają — dla każdego rodzaju szansa, że oba w nim wylądują, zsumowana — a jeden minus to jest szansa, że się różnią.
Nasze dwa zbiory przepuszczone przez ten wzór to arytmetyka sprzed chwili w skróconej formie:
Statystyka koncentracji występuje również w indeksie Simpsona i indeksie Herfindahla–Hirschmana. Nieczystość Giniego to jeden minus ta suma, a nie ta sama statystyka.
Gini na naszych pięciu wierszach
Jesteśmy więc gotowi policzyć nieczystość Giniego dla naszych pięciu wierszy. Najpierw potrzebujemy zestawienia rodzajów, którymi w naszym przypadku są etykiety: tam gdzie powyższe zbiory zawierały kwadraty, koła, trójkąty i gwiazdę, stos wierszy zawiera Yes i No. Policz je więc — ile których etykiet jest w danym stosie — bo każda wielkość w tym artykule wynika z tego słownika.
def label_counts(rows):
"""Tally the labels in a pile — the label is always the last column."""
counts = {}
for row in rows:
counts[row[-1]] = counts.get(row[-1], 0) + 1
return countsPierwszy przebieg po całym zbiorze, label_counts(training_data), daje {'No': 2, 'Yes': 3} — nasi pięcioro pacjentów, policzeni według rodzaju.
Teraz, mając zestawienie pod ręką, możemy policzyć nieczystość Giniego — pięć linijek Pythona:
def gini(rows):
"""Impurity of a pile: 0 when every row in it carries the same label."""
if not rows:
raise ValueError("Cannot measure impurity of an empty group")
impurity = 1
for count in label_counts(rows).values():
share = count / len(rows)
impurity -= share ** 2
return impurityPętla jest wzorem, po jednym składniku na etykietę — podaj jej stos samych Yes, a zwróci 0.0, podaj jeden Yes i jeden No, a zwróci 0.5. Nasz zbiór treningowy, trzy Yes przeciw dwóm No, startuje z:
gini(training_data) → 0.48Przejdziemy przez każde pytanie kandydujące i zobaczymy, które zostawia po sobie najmniej wymieszania, więc 0.48 to liczba do pobicia. To także wysoki punkt startowy: przy dwóch etykietach Gini osiąga maksimum 0.5, gdy są podzielone po równo, więc trzy Yes przeciw dwóm No zostawiają nas na 0.48 — mniej więcej tak pomieszanie, jak to możliwe dla pięciu wierszy.
Spadek nieczystości Giniego — ocena podziału
Oceń kandydata, odejmując od nieczystości rodzica nieczystość Giniego dzieci ważoną ich liczebnością:
Zapisany, zajmuje jedną linię:
I ma cztery kroki:
- podziel stos pytaniem, dając dwa stosy — , wiersze, które odpowiedziały True, i , wiersze, które odpowiedziały False;
- uruchom
ginina każdym z nich; - połącz te dwie liczby w jedną, ważąc tym, ile wierszy poszło na każdą stronę: daje wagę lewego stosu, a prawego, każda z nich to udział wierszy rodzica, który poszedł w tę stronę;
- odejmij to od nieczystości rodzica, .
Zostaje nieczystość, którą pytanie usunęło — im wyższa, tym lepsze pytanie.
Zerowy przyrost oznacza, że podział nie zmienia ważonej nieczystości. Czyste dzieci dają największy możliwy przyrost, równy całej nieczystości rodzica.
Gini jako strata predykcji
Jeśli każdy wiersz liścia otrzymuje ten sam wektor prawdopodobieństw klas, empiryczne proporcje klas minimalizują średnią sumę błędów kwadratowych po wskaźnikach wszystkich klas. Minimum wynosi , czyli nieczystość Giniego. To konwencja wieloklasowej straty Briera z sumowaniem po klasach; wynik binarny liczony tylko dla klasy dodatniej jest o połowę mniejszy. Zachłanne podziały zmniejszają tę stratę w kolejnych węzłach, nie gwarantując globalnie najlepszego drzewa.
Warto podkreślić, po co w ogóle jest ważenie z kroku 3, bo bez niego ocenę łatwo oszukać. Dwaj nasi kandydaci, Is stress_test == normal? i Is vessels >= 1?, dzielą pięć wierszy każdy na jedno idealnie czyste dziecko o Gini dokładnie 0 i jedno dziecko wciąż wymieszane. Różnią się tym, ile danych niesie to czyste dziecko: jeden odłupuje pojedynczego pacjenta i zostawia za sobą cztery wymieszane wiersze, drugi wyciąga dwóch i zostawia trzy. Tylko ważenie widzi tę różnicę. Sprawia, że czyste dziecko liczy się dokładnie tyle, ile waży, więc dziecko z jednego wiersza ledwie się rejestruje, a wynik ustala pozostawiony bałagan.
Oto obaj rozpisani w całości, każdy ze swoimi dwoma dziećmi połączonymi na dwa sposoby — liczonymi po równo, a potem ważonymi udziałem wierszy, które każde dziecko trzyma:
Teraz to samo dla drugiego kandydata. Is vessels >= 1? również odkrawa idealnie czyste dziecko, ale to dziecko trzyma dwóch pacjentów, a nie jednego, a stos, który zostawia, ma trzy wiersze zamiast czterech — i jest brudniejszy: 0.444 zamiast 0.375:
Liczby na obu rysunkach pokazują więc coś mocniejszego niż przeskalowanie.
Liczone po równo, Is stress_test == normal? daje 0.293, a Is vessels >= 1? daje 0.258, więc wygrywa pierwsze pytanie. Ważone, wychodzą 0.180 i 0.213, i wygrywa drugie. Ważenie nie tylko zmniejsza wyniki — odwraca kolejność, a ponieważ to korzeń, dwie odpowiedzi dają drzewa różniące się aż do samego dołu.
Tak implementujemy podział i jego przyrost informacji, wraz z ważeniem. split_rows wykonuje krok 1, sortując wiersze do dwóch stosów, jakie tworzy pytanie, a split_gain wykonuje kroki 2–4, oceniając to, co wyszło, względem tego, co weszło:
def split_rows(rows, rule):
"""Sort every row into the pile where the rule holds, and the pile where it does not."""
true_pile, false_pile = [], []
for row in rows:
(true_pile if rule.holds(row) else false_pile).append(row)
return true_pile, false_pile
def split_gain(parent_impurity, true_pile, false_pile):
"""What went in, minus the two piles that came out, each weighed by its share."""
share = len(true_pile) / (len(true_pile) + len(false_pile))
return parent_impurity - share * gini(true_pile) - (1 - share) * gini(false_pile)Uruchom je na obu powyższych kandydatach, a zwrócą 0.180 i 0.213 — te same liczby, które rysunki wyliczyły ręcznie, teraz policzone, a nie narysowane.
Mechanizm — podziel, potem rekurencja
Mamy już wszystkie elementy: sposób generowania pytań, sposób mierzenia wymieszania stosu i sposób oceny tego, co pytanie z nim robi. Oto procedura, która je łączy. Drzewo decyzyjne rośnie według jednego przepisu, stosowanego do jednego stosu wierszy treningowych:
- Wypróbuj każde pytanie, na jakie pozwalają dane — każdą cechę, każdą wartość, jaką ta cecha przyjmuje.
- Oceń każde pytanie po tym, jak bardzo rozmieszanie etykiet w stosie maleje — to przyrost informacji, zbudowany na nieczystości Giniego, dokładnie tak, jak przed chwilą wyprowadziliśmy.
- Jeśli żadne pytanie nie pomaga, przerwij: stos staje się liściem, a jego zliczenia etykiet stają się predykcją.
- W przeciwnym razie użyj najlepszego pytania, by podzielić stos na dwa mniejsze.
- Uruchom tę samą procedurę na każdym z dwóch stosów.
Ta procedura ma kanoniczną nazwę — rekurencyjny podział binarny — zstępujący, zachłanny algorytm budowania drzew decyzyjnych przez kolejne dzielenie zbioru danych na dwie grupy. Zaczyna od wszystkich danych w korzeniu, ocenia każdą cechę i punkt podziału, by zminimalizować błąd lub zmaksymalizować czystość, i powtarza proces na każdej nowej podgrupie, aż osiągnięty zostanie limit zatrzymania.
Zachłanność oznacza wybór najlepszego podziału bieżącego węzła bez przewidywania kolejnych kroków ani zmiany wcześniejszych decyzji. Podział bez natychmiastowego przyrostu może otworzyć drogę do użytecznych dalszych podziałów, jak w XOR. Nasza reguła stopu je pominie.
Bycie rekurencyjnym to właśnie to, co wycina prostokąty z rysunku ze wstępu: każde wywołanie posiada jeden obszar przestrzeni cech — wiersze, które przetrwały pytania powyżej — i albo dzieli ten obszar, albo pieczętuje go jako liść. Prostokąty to stosy na dnie rekurencji.
Wybór podziału w korzeniu — i remis
Zanim zbudujemy całą rekurencyjną konstrukcję, przyjrzyjmy się szybko implementacji tej części, która działa w pojedynczym węźle — poszukiwaniu najlepszego pytania dzielącego. W korzeniu ten węzeł trzyma wszystkich pięciu pacjentów, a szuka funkcja choose_split, która próbuje każdej wartości każdej cechy i zachowuje najlepszą.
To dwie zagnieżdżone pętle — każda kolumna na zewnątrz, każda różna wartość, jaką ta kolumna przyjmuje, wewnątrz — a każda wytworzona przez nie para przechodzi cztery kroki:
- zbuduj
Rulez kolumny i wartości; - podać ją do
split_rows, które sortuje wiersze na dwa stosy, jakie tworzy; - oceń te stosy przez
split_gain; - porównaj ten wynik z najlepszym dotąd i zachowaj pytanie, jeśli wygrywa.
Gdy obie pętle się kończą, funkcja zwraca pytanie wciąż trzymające najlepszy wynik.
def choose_split(rows):
parent_impurity = gini(rows)
winning_gain, winning_rule = 0, None
for column in range(len(rows[0]) - 1):
for value in sorted(set(row[column] for row in rows), reverse=True):
rule = Rule(column, value)
true_pile, false_pile = split_rows(rows, rule)
if not true_pile or not false_pile:
continue # this rule doesn't divide the data
gain = split_gain(parent_impurity, true_pile, false_pile)
if gain >= winning_gain:
winning_gain, winning_rule = gain, rule
return winning_gain, winning_ruleJedno wywołanie w korzeniu ocenia każde pytanie wyprodukowane przez generator i wraca z tym:
| pytanie | przyrost |
|---|---|
Is stress_test == reversable? | 0.2133 |
Is vessels >= 1? | 0.2133 |
Is stress_test == normal? | 0.1800 |
Is vessels >= 2? | 0.0800 |
Is stress_test == fixed? | 0.0133 |
Is vessels >= 0? | pominięte — byłoby 0 |
Sześć pytań — dwie kolumny po trzy różne wartości — ale liczbę dostaje tylko pięć z nich. Is vessels >= 0? jest prawdą dla każdego pacjenta, bo 0 to najmniejsza wartość, jaką ta kolumna przyjmuje, więc wysyła wszystkie pięć wierszy gałęzią True i żadnego gałęzią False. Jedno dziecko trzyma cały stos, a drugie nic, co nie jest podziałem, lecz kopią — nic nie zostało podzielone, więc nie ma czego oceniać. Zabezpieczenie not true_pile or not false_pile odrzuca je, zanim split_gain je w ogóle zobaczy. Jego przyrost i tak byłby dokładnie 0 — jedno dziecko nic nieważące, drugie będące rodzicem — ale pominięcie zapobiega także zwróceniu nie-podziału jako najlepszego pytania, gdy nic innego nie przekracza zera.
Teraz spójrz na górę tej tabeli, bo to najciekawsza część.
Dwa różne pytania wróciły z tym samym wynikiem 0.2133.
Gdy tak się zdarzy, późniejszy kandydat nadpisuje wcześniejszego.
To zachowanie jest szczegółem implementacyjnym, a w naszym algorytmie bierze się z dwóch rzeczy: kolumny są skanowane w kolejności indeksów, więc stress_test (kolumna 0) dociera pierwszy, a potem zostaje po cichu wyparty przez równie dobre vessels (kolumna 1); oraz porównanie jest napisane z >= zamiast >, co w ogóle umożliwia to wyparcie:
if gain >= winning_gain:Tutaj remisujące pytania wybierają dokładnie te same dwie grupy treningowe, więc nie zmieniają ich dalszych podziałów. Mogą jednak zmienić predykcję dla nowej kombinacji, np. normal z vessels=3. Różne wartości przeglądamy malejąco, aby remisy były powtarzalne; >= nadal wybiera ostatniego remisującego kandydata.
Rekurencja — budowanie drzewa
Skoro pojedynczy węzeł potrafi znaleźć swoje pytanie, jesteśmy gotowi zbudować całe drzewo — rekurencję, która uruchamia to poszukiwanie na stosie za stosem i zapisuje, co znajdzie. Do zapisu potrzeba po jednej klasie na typ węzła z otwierającego diagramu: Leaf trzyma zliczenia etykiet tych wierszy, które do niego dotarły, a Node trzyma pytanie i dwie gałęzie. W słowniku podręcznikowym pytanie to reguła podziału — jeden predykat na jednej cesze, stąd nazwa klasy Rule — a węzeł decyzyjny to ta reguła wpięta w schemat blokowy, z dwiema gałęziami dającymi jej odpowiedziom tak/nie dokąd pójść. Gotowe drzewo można myśleć jako serię reguł podziału.
Zaczynając od góry drzewa i stosując je w drodze w dół — choose_split uczy się tych reguł, a węzły to miejsce, gdzie mieszkają te wybrane.
class Leaf:
def __init__(self, rows):
self.counts = label_counts(rows)
class Node:
def __init__(self, rule, if_true, if_false):
self.rule = rule
self.if_true = if_true
self.if_false = if_false
def grow_tree(rows):
if not rows:
raise ValueError("Training rows must not be empty")
gain, rule = choose_split(rows)
if gain <= 1e-12 or rule is None:
return Leaf(rows) # base case: no rule helps anymore
true_pile, false_pile = split_rows(rows, rule)
return Node(rule, grow_tree(true_pile), grow_tree(false_pile))Gdy uruchomimy to na pięciu pacjentach, otrzymamy takie drzewo — narysowane ze stosami widocznymi przy każdej gałęzi:
Drzewo ma głębokość 2 i trzy liście. Korzeń kieruje dwa wiersze treningowe z vessels >= 1 do czystego liścia Yes. To opis tych dwóch wierszy, nie reguła medyczna.
Warto się przy tym zatrzymać: cała grupa wypadła z danych zupełnie bez domieszki — Gini 0, uzyskane jednym pytaniem. Poziom niżej samotny pacjent normal robi to samo — i zauważ, że stress_test ma trzy wartości, a drzewo pyta tylko o jedną z nich. stress_test == fixed? odłupuje pacjentów z ubytkiem fixed, a wszystko, co nie jest fixed, jedzie gałęzią False razem, nierozróżnione. Tutaj akurat jest to pojedynczy pacjent normal, bo obaj pacjenci reversable odeszli już w korzeniu.
Z trzech liści drzewa dwa są czyste; każdy wiersz treningowy poza kolidującą parą trafia do grupy o zerowym wymieszaniu, a rekurencja zatrzymuje się w każdej z nich właśnie dlatego, że nie ma już nieczystości do usunięcia.
Zostaje trzeci liść, trzymający jedną etykietę Yes i jedną No dla identycznego zestawu wartości cech. Żadne pytanie nie mogłoby rozdzielić tych dwóch pacjentów — i żaden inny model też nie, bo to, co by ich odróżniało, w danych w ogóle nie występuje. Dałoby się to rozwiązać, gdybyśmy użyli więcej predyktorów, bo jedenaście odrzuconych kolumn może dobrze zawierać coś, co tych dwóch pacjentów rozdziela.
Uczenie wybiera strukturę drzewa, progi i liczebności w liściach. Po zakończeniu wywołań rekurencyjnych model jest gotowy do predykcji; ta implementacja nie ma aktualizacji gradientowych ani epok.
Klasyfikacja — odczytywanie prawdopodobieństwa z liścia
Predykcja to znowu rekurencja, i krótsza niż kod treningowy:
def descend(row, node):
if isinstance(node, Leaf):
return node.counts
branch = node.if_true if node.rule.holds(row) else node.if_false
return descend(row, branch)
def as_percentages(counts):
total = sum(counts.values())
return {label: f"{count / total:.0%}" for label, count in counts.items()}Każdy Node przechowuje jedno Rule — indeks kolumny plus wartość — a holds porównuje z nim wpis wiersza w tej kolumnie, zwracając zwykłe True albo False. Ta wartość logiczna to jedyne, czego potrzebuje descend: True wysyła wiersz w if_true, False w if_false, a rekurencja zatrzymuje się, gdy tylko wyląduje na Leaf.
Weźmy dane jednego pacjenta, ['fixed', 0, 'Yes'], i zobaczmy, jak drzewo przewiduje, czy ma on chorobę serca:
- korzeń pyta
Is vessels >= 1?;holdsodczytuje wpisvesselspacjenta — to0— a ponieważ wartość jest liczbowa, oblicza0 >= 1, dając False, więc wiersz idzie gałęzią fałszywą; - ten węzeł pyta
Is stress_test == fixed?;holdsodczytuje wpisstress_testpacjenta —'fixed'— a ponieważ wartość jest napisem, oblicza'fixed' == 'fixed', dając True, więc wiersz idzie gałęzią prawdziwą; - ta gałąź to
Leaf, więcdescendzwraca zapisane tam zliczenia: jedenYesi jedenNo.
Te zwrócone przez descend zliczenia są predykcją w surowej postaci. Można je odczytać jako jedną etykietę, biorąc tę częstszą w liściu — tak robi predict w bibliotece takiej jak sklearn, i jest to jednoznaczne w czystych liściach, gdzie {'Yes': 2} znaczy Yes. Można je też odczytać jako prawdopodobieństwo, dzieląc każde zliczenie przez sumę, czyli predict_proba, i to właśnie robi tu as_percentages.
Mieszany liść zwraca 50/50, bo zawiera po jednym przykładzie każdej klasy. To empiryczne oszacowanie z dwóch wierszy, a nie dowód na 50% ryzyka w populacji ani pełny opis niepewności. Sprzeczne etykiety przy identycznych wejściach wyznaczają dolną granicę błędu treningowego deterministycznego modelu używającego tych cech. Biblioteka zwracająca jedną klasę musi dodatkowo rozstrzygać remisy liczebności.
Widżet poniżej to nieco większa zabawka — dwie cechy liczbowe, progi zamiast naszych mieszanych typów — ale mechanizm jest identyczny i pokazuje dwa widoki drzewa naraz. Lewy panel to podział; prawy to spacer. Dwa suwaki to wartości cech i — przeciąganie ich układa nowy wiersz i przesuwa go po przestrzeni cech. Moment, w którym ten punkt przecina przerywaną linię, to dokładnie moment, w którym zmienia się ścieżka przez drzewo, bo obszar i liść to ten sam obiekt w innym stroju.
Zauważ też, że to drzewo używa cechy oznaczonej dwa razy — raz w korzeniu i znowu dwa poziomy niżej przy innym progu. To dwa różne pytania o jedną kolumnę — ta sama cecha, inna wartość — bo cecha nie zużywa się przez to, że po niej podzielono: pierwsze cięcie rozdziela, co może, a pozostałe wiersze mogą wciąż być rozdzielalne wzdłuż tej samej osi.
Ta sama cecha może wystąpić w kilku węzłach. Każdy węzeł tworzy kandydatów z wierszy, które do niego dotarły, więc brakująca w nich wartość nie jest już tam kandydatem. Na przykład dziecko False naszego korzenia nie zawiera wiersza z vessels=2.
Teraz podajmy mu pięciu pacjentów, których nigdy nie widziało, wszystkich będących prawdziwymi wierszami z tego samego pliku:
patient 2 ['normal', 3] Actual: Yes. Predicted: {'Yes': '100%'}
patient 9 ['reversable', 1] Actual: Yes. Predicted: {'Yes': '100%'}
patient 266 ['fixed', 0] Actual: Yes. Predicted: {'Yes': '50%', 'No': '50%'}
patient 88 ['NA', 0] Actual: No. Predicted: {'No': '100%'}
patient 267 ['NA', 0] Actual: Yes. Predicted: {'No': '100%'}Pierwszy wiersz testowy to ['normal', 3, 'Yes'], a liczba vessels równa 3 nigdy nie pojawiła się w treningu — nasi pięcioro pacjentów pokazywali tylko 0, 1 i 2. Mimo to wpada do liścia, bo vessels >= 1 jest progiem, a nie wyszukiwaniem: 3 go przekracza i idzie tą samą ścieżką co 1. Progi uogólniają poza swoje wartości treningowe za darmo.
Z brakującymi wartościami jest inaczej: model po prostu nie obsługuje ich poprawnie. Pacjenci 88 i 267 nie mają zapisanego testu wysiłkowego — w pliku widnieje NA — a ponieważ NA nigdy nie pojawiło się w treningu, test stress_test == fixed? zawodzi i wiersz zsuwa się gałęzią False, przez co pacjent 267 wraca jako {'No': '100%'}, choć chorobę ma.
W praktyce zajęłbyś się tym przed treningiem — usunął niekompletne wiersze, uzupełnił braki albo uczynił z "missing" osobną kategorię.
Nieznana kategoria trafia do tej samej gałęzi False co każda wartość niespełniająca testu równości. Model nie wykrywa nowości ani nie obniża wyświetlanego procentu. W przykładzie powyżej 100% No wynika z liścia zawierającego tylko jeden wiersz treningowy.
Obsługa braków zależy od implementacji. Niektóre biblioteki uczą się domyślnej gałęzi, a oryginalny CART może korzystać z podziałów zastępczych. Nie oznacza to, że dowolny napis, np. "NA", zostanie rozpoznany jako brak — należy użyć formatu oczekiwanego przez bibliotekę.
Dlaczego jedno drzewo to nie koniec historii
To niewielka implementacja w stylu CART. Jej główne ograniczenia to nieograniczony wzrost, wrażliwość na próbkę treningową, uproszczona obsługa kategorii i brak jawnej polityki brakujących wartości.
Głębokie drzewo może się przeuczyć. Tworząc małe liście, poprawia dopasowanie treningowe, ale może odzwierciedlać szum próbki. Zatrzymanie przy znikomym przyroście zapobiega bezużytecznym bieżącym podziałom, lecz nie kontroluje uogólniania. Służą temu limity głębokości, minimalny rozmiar liścia i przycinanie.
W większym przykładzie użyjemy DecisionTreeClassifier ze scikit-learn. Stosuje tę samą zachłanną zasadę, ale progi w środkach przedziałów, rozstrzyganie remisów i szczegóły zatrzymania mogą dać inne drzewo niż nasz kod.
Teraz skierujmy go na zbiór o raku piersi (398 wierszy treningowych, 171 testowych, 30 cech):
from sklearn.datasets import load_breast_cancer
from sklearn.model_selection import train_test_split
from sklearn.tree import DecisionTreeClassifier
data = load_breast_cancer()
X_tr, X_te, y_tr, y_te = train_test_split(
data.data, data.target, test_size=0.3, random_state=0
)
unpruned = DecisionTreeClassifier(criterion="gini", random_state=0).fit(X_tr, y_tr)
print(unpruned.get_depth(), unpruned.get_n_leaves())
print(unpruned.score(X_tr, y_tr), unpruned.score(X_te, y_te))
for depth in (1, 2, 3, 5, 7):
pruned = DecisionTreeClassifier(
criterion="gini", max_depth=depth, random_state=0
).fit(X_tr, y_tr)
print(depth, pruned.score(X_tr, y_tr), pruned.score(X_te, y_te))Niepohamowane drzewo — to, którego nic nie powstrzymuje — wychodzi tak:
UNPRUNED (scikit-learn)
depth 7
leaves 19
train acc 1.000
test acc 0.912Dokładność treningowa 1.000 oznacza odtworzenie wszystkich 398 etykiet treningowych. Nie wymaga to liścia dla każdego wiersza: drzewo ma tylko 19 liści. Dokładność na 171 odłożonych wierszach wynosi 0.912.
Dodajmy teraz jedno pokrętło, którego nasza wersja nie ma — max_depth, twardy limit tego, na ile pytań w głąb drzewo może pójść, zatrzymujący dzielenie niezależnie od tego, czy został jeszcze przyrost do zebrania. Gdy uruchomimy to samo drzewo ograniczone serią głębokości i zbierzemy wyniki, wychodzi to:
| max_depth | trening | test | różnica |
|---|---|---|---|
| 1 | 0.930 | 0.895 | +0.035 |
| 2 | 0.960 | 0.947 | +0.012 |
| 3 | 0.967 | 0.947 | +0.020 |
| 5 | 0.987 | 0.936 | +0.052 |
| 7 (bez przycinania) | 1.000 | 0.912 | +0.088 |
W tym podziale głębokość 2 daje 0.947 na danych odłożonych, a głębokość 7 — 0.912, mimo wyższej dokładności treningowej. To przykład przeuczenia, nie dowód na ogólną optymalność głębokości 2. Jeśli te wyniki służą do wyboru głębokości, odłożony zbiór pełni rolę walidacyjną. Wybrany model należy ocenić na osobnym zbiorze testowym.
Pojedyncze drzewo może być też niestabilne. Małe zmiany wierszy treningowych mogą zmienić zwycięzcę wśród podobnie ocenionych kandydatów, a wraz z nim całe poddrzewo. Ta wrażliwość na próbkę różni się od możliwego do uniknięcia problemu implementacji, takiego jak iterowanie nieuporządkowanego zbioru.
Sortowanie kandydatów usuwa różnice między uruchomieniami wynikające z kolejności zbioru w Pythonie. Nie usuwa niestabilności statystycznej: inna próbka treningowa nadal może dać inne drzewo.
Te dwa zachowania to dwie strony kompromisu obciążenia i wariancji. W uczeniu statystycznym obciążenie (bias) to błąd wprowadzony przez przybliżanie skomplikowanej rzeczywistości prostszym modelem — model zbyt sztywny, by odwzorować wzorzec, będzie się mylił niezależnie od tego, ile danych mu wręczysz. Wariancja to wielkość, o jaką zmieniłby się dopasowany model, gdybyś oszacował go z innego zbioru treningowego: przetrenuj na innej próbie pacjentów, a metoda o wysokiej wariancji da ci zauważalnie inny model, popełniający zauważalnie inne błędy.
Jak obciążenie i wariancja sumują się w błąd predykcji
W regresji z błędem kwadratowym, dla ustalonego wejścia, oczekiwany błąd testowy rozkłada się na kwadrat obciążenia, wariancję predykcji między próbkami treningowymi i warunkową wariancję szumu. Wartość oczekiwana obejmuje różne próbki treningowe i nowe wyniki:
Obciążenie występuje w kwadracie, bo jest wielkością ze znakiem — jak daleko średnia predykcja modelu leży od prawdy — która inaczej znosiłaby się, zamiast się kumulować. Dla błędu klasyfikacji działają te same trzy źródła, ale nie sumują się tak schludnie; intuicja się przenosi, arytmetyka nie.
Obciążenie mierzy różnicę między średnią predykcją a prawdziwą średnią warunkową. Wariancja mierzy zmienność predykcji między próbkami treningowymi. Metoda może mieć oba składniki; nie są to osobne części każdego pojedynczego błędu.
Zwiększanie głębokości często zmniejsza obciążenie i zwiększa wariancję, lecz jest to tendencja, nie gwarancja dla każdego zbioru lub wyniku dokładności. Szum warunkowy względem dostępnych cech ogranicza osiągalną jakość oczekiwaną; sprzeczna para w małej próbce nie wyznacza liczbowo tej granicy w populacji.
Istnieją ugruntowane metody radzenia sobie z tymi problemami — limit głębokości, minimalna liczba wierszy na liść, minimalny przyrost warty podziału oraz obcinanie gałęzi po fakcie. W praktyce rzadko stosuje się je do samotnego drzewa; to pokrętła, które kręcisz wewnątrz zespołu — modelu zbudowanego z wielu drzew, których odpowiedzi łączy się w jedną, a właśnie tym są las losowy i gradient boosting.
Przycinanie wstępne ogranicza wzrost przez parametry takie jak max_depth i min_samples_leaf. Przycinanie po uczeniu najpierw buduje większe drzewo, potem usuwa gałęzie. Przycinanie koszt–złożoność w CART równoważy nieczystość treningową i karę za każdy liść; scikit-learn udostępnia tę karę jako ccp_alpha. Wybierz ją przez walidację lub walidację krzyżową, zachowując osobny zbiór testowy.
Las losowy uśrednia wiele drzew uczonych z losowaniem wierszy i cech, zmniejszając wariancję, jeśli ich błędy nie są doskonale skorelowane. Gradient boosting dopasowuje drzewa kolejno, poprawiając stratę całego zespołu. Dalej przeczytaj o drzewie regresyjnym.
Gdzie nasza wersja jest wolniejsza od prawdziwej
Nasze wyszukiwanie podziałów wielokrotnie przegląda te same wiersze. Implementacje produkcyjne mogą ponownie wykorzystywać liczebności lub histogramy. Różnią się też progami, obsługą kategorii i braków oraz rozstrzyganiem remisów — szybkość nie jest jedyną różnicą.
Pętla przebiega po każdej cesze i każdej różnej wartości, jaką ta cecha przyjmuje, więc liczba kandydatów to cechy × wartości.
Każdy kandydat kosztuje potem pełne przejście po danych: split_rows obchodzi każdy wiersz, by posortować go do dwóch stosów, a split_gain wywołuje gini na każdym stosie, które liczy jego etykiety od zera. To . Na pięciu wierszach z trzema wartościami na kolumnę — niewidoczne. Na cesze ciągłej — powiedzmy cholesterolu — niemal każdy wiersz niesie inną wartość, więc liczba kandydatów rośnie razem z danymi, a każdy kandydat wciąż kosztuje pełny skan: kwadratowo względem liczby wierszy i beznadziejnie przy stu tysiącach.
Weź pięć wierszy kolumny cholesterol — 210 (No), 233 (No), 250 (Yes), 286 (Yes), 300 (Yes). Generator zamienia je w pięć pytań kandydujących, po jednym na zaobserwowaną wartość, z których cztery faktycznie dzielą stos:
| kandydat | poniżej progu | na progu lub powyżej |
|---|---|---|
>= 210 | puste | wszystkie pięć |
>= 233 | 210 | 233, 250, 286, 300 |
>= 250 | 210, 233 | 250, 286, 300 |
>= 286 | 210, 233, 250 | 286, 300 |
>= 300 | 210, 233, 250, 286 | 300 |
Prześledźmy dwóch z nich, >= 233 i >= 250, przez nasz kod.
Dla >= 233 split_rows obchodzi wszystkie pięć wierszy i wrzuca 210 na listę False, a pozostałe cztery na listę True. split_gain wywołuje potem gini na każdej, a gini obchodzi jednowierszowy stos, licząc etykiety, a potem czterowierszowy stos, licząc etykiety. Pięć odwiedzin na podział, pięć na liczenie. Dla >= 250 zaczyna od nowa od tych samych pięciu wierszy, i tak dalej w dół listy:
>= 233: split_rows 5 rows → gini({210}) + gini({233,250,286,300}) = 10 visits
>= 250: split_rows 5 rows → gini({210,233}) + gini({250,286,300}) = 10 visits
>= 286: split_rows 5 rows → gini({210,233,250}) + gini({286,300}) = 10 visits
>= 300: split_rows 5 rows → gini({210,233,250,286}) + gini({300}) = 10 visitsCzterdzieści odwiedzin wierszy, i nic nie jest przenoszone między liniami — mimo że każda para stosów różni się od pary powyżej dokładnie jednym wierszem.
Oto te same pięć wierszy, posortowane po cholesterolu i niosące etykietę disease, jaką każdy pacjent okazał się mieć — tak trzymałaby je prawdziwa implementacja:
| cholesterol | disease |
|---|---|
| 210 | No |
| 233 | No |
| 250 | Yes |
| 286 | Yes |
| 300 | Yes |
Prawdziwe implementacje wyciągają obie odpowiedzi z pierwszego przejścia. Posortuj wiersze po cesze, a potem, oceniając pierwszego kandydata, przejdź przez nie raz, prowadząc bieżące zestawienie widzianych dotąd etykiet — do końca tego jednego przejścia każdy późniejszy kandydat też ma już odpowiedź:
| cholesterol, disease | bieżące zestawienie |
|---|---|
| 210, No | {No: 1} |
| 233, No | {No: 2} |
| 250, Yes | {No: 2, Yes: 1} |
| 286, Yes | {No: 2, Yes: 2} |
Każda linia to linia powyżej plus etykieta właśnie minionego wiersza: 210 to No, więc zestawienie otwiera się na {No: 1}; 233 to kolejne No, doprowadzające do {No: 2}; 250 to Yes, dodające pierwsze Yes; i tak dalej. Jeden przyrost na wiersz, pięć wierszy, jedno przejście.
Ponieważ wiersze przychodzą w porządku rosnącym, każde z tych zestawień jest też grupą wpadającą poniżej określonego progu — i to właśnie zamienia je w odpowiedź:
| zestawienie | obsługuje pytanie | grupa poniżej | grupa na progu lub powyżej |
|---|---|---|---|
{No: 1} | >= 233 | {No: 1} | {No: 1, Yes: 3} |
{No: 2} | >= 250 | {No: 2} | {Yes: 3} |
{No: 2, Yes: 1} | >= 286 | {No: 2, Yes: 1} | {Yes: 2} |
{No: 2, Yes: 2} | >= 300 | {No: 2, Yes: 2} | {Yes: 1} |
Ponieważ wzór na przyrost waży oboje dzieci, potrzebne są obie grupy — ale w trakcie przejścia śledzona jest tylko pierwsza. Grupy na progu lub powyżej nigdy nie trzeba liczyć; da się ją wyprowadzić. Własne sumy węzła zostały policzone przy jego tworzeniu — tutaj {No: 2, Yes: 3} — więc cokolwiek nie jest poniżej progu, jest powyżej. Żadne pytanie niczego nie odczytuje ponownie — po dwa odczyty i odejmowanie na każde, a Gini wynika z czterech liczb całkowitych.
Dla wierszy i cech naiwny algorytm może wymagać pracy w jednym węźle. Sortowanie każdej cechy liczbowej i jednokrotny przegląd zmniejszają koszt do około z sortowaniem, przy stałej liczbie klas. Sam przegląd jest liniowy. To porównanie dla jednego węzła; koszt całego drzewa zależy też od jego kształtu.