Stablo odlučivanja u R: Klasifikacijsko stablo s primjerom
⚡ Pametni sažetak
Stabla odlučivanja u R-u dijele podatke na grane koristeći jednostavna pravila da ili ne sve dok svaki list ne sadrži jednu dominantnu klasu. Ovaj vodič gradi, iscrtava, evaluira i podešava rpart klasifikacijsko stablo na skupu podataka o preživljavanju Titanica.
Što su stabla odlučivanja?
Stabla odlučivanja su svestrani algoritmi strojnog učenja koji mogu obavljati i klasifikacijske i regresijske zadatke. To su vrlo moćni algoritmi, sposobni za prilagođavanje složenih skupova podataka. Osim toga, stabla odlučivanja su temeljne komponente slučajnih šuma, koje su među najmoćnijim algoritmima strojnog učenja dostupnim danas.
Prije nego što ga izgradite u kodu, korisno je znati kako stablo odlučuje gdje će se podijeliti.
Kako funkcionira stablo odlučivanja?
Stablo odlučivanja pretvara skup podataka u dijagram toka pitanja s odgovorima da ili ne izgrađen od tri vrste čvorova: korijen održava svako promatranje treninga, unutarnji čvor postavlja pitanje o jednom prediktoru i dijeli podatke na dva dijela, i list zaustavlja cijepanje i vraća većinsku klasu.
Rast slijedi pohlepni postupak koji se zove rekurzivno binarno particioniranje:
- Procijenite svaku podjelu kandidata. Za svaki prediktor i graničnu vrijednost, izmjerite koliko bi nečiste bile dvije rezultirajuće skupine.
- Zadrži najbolji. Podjela koja najviše smanjuje nečistoću postaje pitanje koje se postavlja na tom čvoru.
- Ponovite za svako dijete dok ga ne zaustavi neko kontrolno pravilo: minsplit, minbucket, maxdepth ili cp.
- Orezati. cp zatim orezuje grane koje se ne isplate same, što sprječava stablo da pamti skup za učenje.
Budući da svako pitanje uspoređuje jednu varijablu s pragom, algoritam nikada ne treba skaliranje ili lažno kodiranje.
Ginijev indeks u odnosu na entropiju u stablima odlučivanja
Ta se nečistoća može mjeriti na dva načina, a rpart vam omogućuje da odaberete.
| Kriteriji | Ginijev indeks | Entropija (prikupljanje informacija) |
|---|---|---|
| Formula | 1 – zbroj kvadrata proporcija klasa | -zbroj p puta log²(p) |
| Raspon (dvije klase) | 0 0.5 se | 0 1 se |
| računanje | Brže, bez logaritma | Sporije, koristi logaritme |
| postavka rparta | Zadano | parms = popis(razdvajanje = “informacije”) |
fit_entropy <- rpart(survived~., data = data_train, method = 'class', parms = list(split = "information"))
U praksi oba kriterija većinu vremena odabiru istu podjelu, pa je zadani Ginijev indeks siguran izbor.
Prednosti i nedostaci stabala odlučivanja
Kompromisi vam govore kada je jedno stablo dovoljno, a kada prijeći na ansambl.
Prednosti
- Potpuno interpretirano: Prilagođeni model je dijagram koji bilo koja zainteresirana strana može pročitati.
- Minimalna predobrada: Nije potrebno skaliranje ili normalizacija, a faktori rade izvorno.
- Obavlja oba zadatka: method = 'class' odgovara klasifikatoru, a method = 'anova' odgovara regresijskom stablu.
- Brzo za treniranje: Veliki skupovi podataka stanu u sekunde, pa stabla čine korisnu prvu osnovu.
Nedostaci
- Visoka varijanca: Mala promjena u podacima za obuku može proizvesti potpuno drugačije stablo.
- Sklon prekomjernom opremanju: Neograničeno stablo raste sve dok svaki list ne postane čist, osim ako ga cp i maxdepth ne ograniče.
- Samo paralelne podjele po osima: dijagonalne granice trebaju mnogo rezova u obliku stepenica.
Lijek za prve dvije slabosti je usrednjavanje većeg broja stabala, što je ono što slučajna šuma radi.
Kako trenirati i vizualizirati stablo odlučivanja u R-u
Za izgradnju vašeg prvog stabla odlučivanja u R-u, proći ćete kroz sedam koraka:
- Korak 1: Uvezite podatke
- Korak 2: Očistite skup podataka
- Korak 3: Stvorite set za treniranje/testiranje
- Korak 4: Izgradite model
- Korak 5: Napravite predviđanje
- Korak 6: Izmjerite učinak
- Korak 7: Podesite hiperparametre
Korak 1) Uvezite podatke
Ako vas zanima kakva je sudbina Titanica, ovaj video možete pogledati na Youtube. Svrha ovog skupa podataka je predvidjeti koji ljudi imaju veću vjerojatnost da će preživjeti nakon sudara s santom leda. Skup podataka sadrži 13 varijabli i 1309 opažanja. Skup podataka je poredan po varijabli X.
set.seed(678) path <- 'https://raw.githubusercontent.com/guru99-edu/R-Programming/master/titanic_data.csv' titanic <-read.csv(path) head(titanic)
Izlaz:
## X pclass survived name sex ## 1 1 1 1 Allen, Miss. Elisabeth Walton female ## 2 2 1 1 Allison, Master. Hudson Trevor male ## 3 3 1 0 Allison, Miss. Helen Loraine female ## 4 4 1 0 Allison, Mr. Hudson Joshua Creighton male ## 5 5 1 0 Allison, Mrs. Hudson J C (Bessie Waldo Daniels) female ## 6 6 1 1 Anderson, Mr. Harry male ## age sibsp parch ticket fare cabin embarked ## 1 29.0000 0 0 24160 211.3375 B5 S ## 2 0.9167 1 2 113781 151.5500 C22 C26 S ## 3 2.0000 1 2 113781 151.5500 C22 C26 S ## 4 30.0000 1 2 113781 151.5500 C22 C26 S ## 5 25.0000 1 2 113781 151.5500 C22 C26 S ## 6 48.0000 0 0 19952 26.5500 E12 S ## home.dest ## 1 St Louis, MO ## 2 Montreal, PQ / Chesterville, ON ## 3 Montreal, PQ / Chesterville, ON ## 4 Montreal, PQ / Chesterville, ON ## 5 Montreal, PQ / Chesterville, ON ## 6 New York, NY
tail(titanic)
Izlaz:
## X pclass survived name sex age sibsp ## 1304 1304 3 0 Yousseff, Mr. Gerious male NA 0 ## 1305 1305 3 0 Zabour, Miss. Hileni female 14.5 1 ## 1306 1306 3 0 Zabour, Miss. Thamine female NA 1 ## 1307 1307 3 0 Zakarian, Mr. Mapriededer male 26.5 0 ## 1308 1308 3 0 Zakarian, Mr. Ortin male 27.0 0 ## 1309 1309 3 0 Zimmerman, Mr. Leo male 29.0 0 ## parch ticket fare cabin embarked home.dest ## 1304 0 2627 14.4583 C ## 1305 0 2665 14.4542 C ## 1306 0 2665 14.4542 C ## 1307 0 2656 7.2250 C ## 1308 0 2670 7.2250 C ## 1309 0 315082 7.8750 S
Iz ispisa glave i repa možete primijetiti da podaci nisu pomiješani. Ovo je veliki problem! Kada podijelite podatke između skupa vlakova i skupa za testiranje, odabrat ćete samo putnik iz klase 1 i 2 (Nijedan putnik iz klase 3 nije u prvih 80 posto opažanja), što znači da algoritam nikada neće vidjeti značajke putnika iz klase 3. Ova će pogreška dovesti do lošeg predviđanja.
Da biste prevladali ovaj problem, možete koristiti funkciju sample().
shuffle_index <- sample(1:nrow(titanic)) head(shuffle_index)
R kod stabla odlučivanja Objašnjenje
- sample(1:nrow(titanic)): Generirajte nasumični popis indeksa od 1 do 1309 (tj. najveći broj redaka).
Izlaz:
## [1] 288 874 1078 633 887 992
Koristit ćete ovaj indeks za miješanje titanskog skupa podataka.
titanic <- titanic[shuffle_index, ]
head(titanic)
Izlaz:
## X pclass survived ## 288 288 1 0 ## 874 874 3 0 ## 1078 1078 3 1 ## 633 633 3 0 ## 887 887 3 1 ## 992 992 3 1 ## name sex age ## 288 Sutton, Mr. Frederick male 61 ## 874 Humblen, Mr. Adolf Mathias Nicolai Olsen male 42 ## 1078 O'Driscoll, Miss. Bridget female NA ## 633 Andersson, Mrs. Anders Johan (Alfrida Konstantia Brogren) female 39 ## 887 Jermyn, Miss. Annie female NA ## 992 Mamee, Mr. Hanna male NA ## sibsp parch ticket fare cabin embarked home.dest## 288 0 0 36963 32.3208 D50 S Haddenfield, NJ ## 874 0 0 348121 7.6500 F G63 S ## 1078 0 0 14311 7.7500 Q ## 633 1 5 347082 31.2750 S Sweden Winnipeg, MN ## 887 0 0 14313 7.7500 Q ## 992 0 0 2677 7.2292 C
Korak 2) Očistite skup podataka
Nekoliko varijabli sadrži NA vrijednosti. Čišćenje se provodi u tri dijela:
- Izbacite varijable home.dest, cabin, name, X i ticket
- Stvorite faktorske varijable za pclass i survived
- Odbaci NA
library(dplyr) # Drop variables clean_titanic <- titanic %>% select(-c(home.dest, cabin, name, X, ticket)) %>% #Convert to factor level mutate(pclass = factor(pclass, levels = c(1, 2, 3), labels = c('Upper', 'Middle', 'Lower')), survived = factor(survived, levels = c(0, 1), labels = c('No', 'Yes'))) %>% na.omit() glimpse(clean_titanic)
Code Objašnjenje
- select(-c(home.dest, cabin, name, X, ticket)): Ispustite nepotrebne varijable
- pclass = factor(pclass, levels = c(1,2,3), labels= c('Gornji', 'Srednji', 'Donji')): Dodaj oznaku varijabli pclass. 1 postaje Gornji, 2 postaje Srednji, a 3 postaje Donji
- faktor(preživio, razine = c(0,1), oznake = c('Ne', 'Da')): Dodaj oznake varijabli preživjelo. 0 postaje Ne, a 1 postaje Da
- na.omit(): Ukloni NA opažanja
Izlaz:
## Observations: 1,045 ## Variables: 8 ## $ pclass <fctr> Upper, Lower, Lower, Upper, Middle, Upper, Middle, U... ## $ survived <fctr> No, No, No, Yes, No, Yes, Yes, No, No, No, No, No, Y... ## $ sex <fctr> male, male, female, female, male, male, female, male... ## $ age <dbl> 61.0, 42.0, 39.0, 49.0, 29.0, 37.0, 20.0, 54.0, 2.0, ... ## $ sibsp <int> 0, 0, 1, 0, 0, 1, 0, 0, 4, 0, 0, 1, 1, 0, 0, 0, 1, 1,... ## $ parch <int> 0, 0, 5, 0, 0, 1, 0, 1, 1, 0, 0, 1, 1, 0, 2, 0, 4, 0,... ## $ fare <dbl> 32.3208, 7.6500, 31.2750, 25.9292, 10.5000, 52.5542, ... ## $ embarked <fctr> S, S, S, S, S, S, S, S, S, C, S, S, S, Q, C, S, S, C...
Korak 3) Stvorite set za treniranje/testiranje
Prije nego što uvježbate svoj model, morate izvršiti dva koraka:
- Stvorite vlak i skup testova: Vi trenirate model na skupu vlakova i testirate predviđanje na test skupu (tj. nevidljivi podaci)
- Instalirajte rpart.plot s konzole
Uobičajena praksa je podijeliti podatke 80/20, 80 posto podataka služi za treniranje modela, a 20 posto za predviđanje. Morate stvoriti dva odvojena podatkovna okvira. Ne želite dirati testni set dok ne završite izradu svog modela. Možete stvoriti ime funkcije create_train_test() koja uzima tri argumenta.
create_train_test(df, size = 0.8, train = TRUE) arguments: -df: Dataset used to train the model. -size: Size of the split. By default, 0.8. Numerical value -train: If set to `TRUE`, the function creates the train set, otherwise the test set. Default value sets to `TRUE`. Boolean value.You need to add a Boolean parameter because R does not allow to return two data frames simultaneously.
create_train_test <- function(data, size = 0.8, train = TRUE) { n_row = nrow(data) total_row = size * n_row train_sample <- 1: total_row if (train == TRUE) { return (data[train_sample, ]) } else { return (data[-train_sample, ]) } }
Code Objašnjenje
- funkcija (podaci, veličina=0.8, vlak = TRUE): Dodajte argumente u funkciju
- n_row = nrow(data): Broj redaka u skupu podataka
- total_row = size*n_row: Vratite n-ti red za konstrukciju vlaka
- train_sample <- 1:total_row: Odaberite prvi red do n-tog reda
- if (train ==TRUE){ } else { }: Ako se uvjet postavi na istinito, vrati skup niza, inače ispitni skup.
Možete testirati svoju funkciju i provjeriti dimenziju.
data_train <- create_train_test(clean_titanic, 0.8, train = TRUE) data_test <- create_train_test(clean_titanic, 0.8, train = FALSE) dim(data_train)
Izlaz:
## [1] 836 8
dim(data_test)
Izlaz:
## [1] 209 8
Skup podataka o vlaku ima 836 redaka i 8 stupaca, dok testni skup podataka ima 209 redaka i istih 8 stupaca.
Koristite funkciju prop.table() u kombinaciji s table() da biste provjerili je li proces nasumičnog odabira točan.
prop.table(table(data_train$survived))
Izlaz:
## ## No Yes ## 0.5944976 0.4055024
prop.table(table(data_test$survived))
Izlaz:
## ## No Yes ## 0.5789474 0.4210526
U oba skupa podataka, broj preživjelih je isti, oko 40 posto.
Instalirajte rpart.plot
rpart.plot nije dostupan iz conda biblioteka. Možete ga instalirati s konzole:
install.packages("rpart.plot")
Korak 4) Izgradite model
Spremni ste za izgradnju modela. Sintaksa za funkciju stabla odlučivanja rpart() je:
rpart(formula, data=, method='') arguments: - formula: The function to predict - data: Specifies the data frame - method: - "class" for a classification tree - "anova" for a regression tree
Koristite metodu klase jer predviđate klasu.
library(rpart) library(rpart.plot) fit <- rpart(survived~., data = data_train, method = 'class') rpart.plot(fit, extra = 106)
Code Objašnjenje
- rpart(): Funkcija za uklapanje u model. Argumenti su:
- preživio ~.: Formula of the Decision Trees
- data = data_train: Skup podataka
- method = 'class': Fit binarni model
- rpart.plot(fit, extra= 106): Nacrtajte stablo. Dodatni argument postavljen je na 106, što prikazuje vjerojatnost druge klase plus postotak opažanja u svakom čvoru. Možete se pozvati na vinjeta za više informacija o drugim izborima.
Izlaz:
Počinjete od korijenskog čvora, na vrhu grafa i na dubini 0 od 3:
- Na vrhu, to je ukupna vjerojatnost preživljavanja. Prikazuje udio putnika koji su preživjeli nesreću. Preživjelo je 41 posto putnika.
- Ovaj čvor pita je li spol putnika muški. Ako jest, onda idete do lijevog djeteta korijena (dubina 1). 63 posto su muškarci s vjerojatnošću preživljavanja od 21 posto.
- U drugom čvoru pitate je li muški putnik stariji od 3.5 godine. Ako da, onda je šansa za preživljavanje 19 posto.
- Nastavite tako da shvatite koje značajke utječu na vjerojatnost preživljavanja.
Imajte na umu da je jedna od mnogih kvaliteta stabala odlučivanja to što zahtijevaju vrlo malo pripreme podataka. Konkretno, ne zahtijevaju skaliranje ili centriranje značajki.
Prema zadanim postavkama, funkcija rpart() koristi Gini mjera nečistoće za odabir svake podjele. Što je veća Ginijeva vrijednost, to su klase unutar tog čvora miješanije, pa algoritam uvijek odabire podjelu koja je najviše snižava.
Korak 5) Napravite predviđanje
Možete predvidjeti svoj testni skup podataka. Da biste napravili predviđanje, možete koristiti funkciju predict(). Osnovna sintaksa predviđanja za R stablo odlučivanja je:
predict(fitted_model, df, type = 'class') arguments: - fitted_model: This is the object stored after model estimation. - df: Data frame used to make the prediction - type: Type of prediction - 'class': for classification - 'prob': to compute the probability of each class - 'vector': Predict the mean response at the node level
Sada predviđate, za svakog od 209 putnika u testnom skupu, očekuje li model da će preživjeti sudar.
predict_unseen <-predict(fit, data_test, type = 'class')
Code Objašnjenje
- predict(fit, data_test, type = 'class'): Predvidite klasu (0/1) skupa testova
Sada usporedite predviđene klase sa stvarnim rezultatima.
table_mat <- table(data_test$survived, predict_unseen)
table_mat
Code Objašnjenje
- table(data_test$survived, predict_unseen): Izradi tablicu kontingencije predviđenih klasa u odnosu na stvarni ishod
Izlaz:
## predict_unseen ## No Yes ## No 106 15 ## Yes 30 58
Redci predstavljaju stvarne vrijednosti, stupci su predviđanja. Model je ispravno identificirao 106 nepreživjelih i 58 preživjelih, ali je 15 nepreživjelih označio kao preživjele, a 30 preživjelih kao umrle.
Korak 6) Mjerite učinak
Možete izračunati mjeru točnosti za zadatak klasifikacije s matrica zabune:
The matrica zabune je bolji izbor za procjenu učinkovitosti klasifikacije. Opća ideja je brojati koliko su puta True instance klasificirane kao False.
Svaki redak u matrici konfuzije predstavlja stvarnu metu, dok svaki stupac predstavlja predviđenu metu. Prvi redak ove matrice uzima u obzir putnike koji su umrli (negativna klasa): 106 je ispravno klasificirano kao mrtvo (Istinski negativan), dok je 15 pogrešno klasificirano kao preživjeli (Lažno pozitivno). Drugi red uzima u obzir preživjele: 58 ih je ispravno identificirano (Prava pozitiva), dok je 30 propušteno (Lažno negativan).
Možete izračunati test točnosti iz matrice zabune:
To je udio stvarnog pozitivnog i pravog negativnog u zbroju matrice. S R možete kodirati na sljedeći način:
accuracy_Test <- sum(diag(table_mat)) / sum(table_mat)
Code Objašnjenje
- sum(diag(table_mat)): Zbroj dijagonale
- sum(table_mat): Zbroj matrice.
Možete ispisati točnost testnog skupa:
print(paste('Accuracy for test', accuracy_Test))
Izlaz:
## [1] "Accuracy for test 0.784688995215311"
Točnost na testnom skupu je 0.7847, što je 78.47 posto. Ponovite vježbu na skupu za učenje kako biste vidjeli koliko model previše prilagođava se.
Korak 7) Podesite hiper-parametre
Stablo odlučivanja u R ima različite parametre koji kontroliraju aspekte prilagodbe. U biblioteci stabla odlučivanja rpart možete kontrolirati parametre pomoću funkcije rpart.control(). U sljedećem kodu uvodite parametre koje ćete podešavati. Možete se obratiti na vinjeta za ostale parametre.
rpart.control(minsplit = 20, minbucket = round(minsplit/3), maxdepth = 30) Arguments: -minsplit: Set the minimum number of observations in the node before the algorithm perform a split -minbucket: Set the minimum number of observations in a terminal node, i.e. the leaf -maxdepth: Set the maximum depth of any node of the final tree. The root node is treated as depth 0
Postupit ćemo na sljedeći način:
- Konstruirajte funkciju za vraćanje točnosti
- Podesite maksimalnu dubinu
- Podesite minimalni broj uzoraka koji čvor mora imati prije nego što se može podijeliti
- Podesite minimalni broj uzoraka koji listni čvor mora imati
Možete napisati funkciju za prikaz točnosti. Jednostavno omotate kod koji ste prije koristili:
- predvidi: predict_unseen <- predict(fit, data_test, type = 'class')
- Izradi tablicu: table_mat <- table(data_test$survived, predict_unseen)
- Izračunaj točnost: accuracy_Test <- sum(diag(table_mat))/sum(table_mat)
accuracy_tune <- function(fit) { predict_unseen <- predict(fit, data_test, type = 'class') table_mat <- table(data_test$survived, predict_unseen) accuracy_Test <- sum(diag(table_mat)) / sum(table_mat) accuracy_Test }
Sada podesite parametre i provjerite možete li poboljšati zadani model. Podsjećamo, morate nadmašiti točnost od 0.7847.
control <- rpart.control(minsplit = 4,
minbucket = round(5 / 3),
maxdepth = 3,
cp = 0)
tune_fit <- rpart(survived~., data = data_train, method = 'class', control = control)
accuracy_tune(tune_fit)
Izlaz:
## [1] 0.7990431
Sa sljedećim parametrom:
minsplit = 4
minbucket = round(5/3)
maxdepth = 3
cp = 0
Točnost raste s 0.7847 na 0.7990, tako da podešeno stablo nadmašuje zadanu konfiguraciju za oko 1.4 postotna boda.
Stabla odlučivanja u R-u: Kratki pregled funkcija
Donja tablica navodi svaku funkciju korištenu u gore navedenih sedam koraka, zajedno s paketom koji je pruža i parametrima koje očekuje. R.
| Knjižnica | Cilj | funkcija | Klasa | Parametri | Detaljnije |
|---|---|---|---|---|---|
| rpart | Stablo klasifikacije vlakova u R | rpart() | razred | formula, df, metoda | |
| rpart | Uvježbavanje regresijskog stabla | rpart() | anova | formula, df, metoda | |
| rpart | Nacrtajte drveće | rpart.plot() | ugrađeni model | ||
| baza | predvidjeti | predvidjeti() | razred | opremljeni model, tip | |
| baza | predvidjeti | predvidjeti() | prob | opremljeni model, tip | |
| baza | predvidjeti | predvidjeti() | vektor | opremljeni model, tip | |
| rpart | Kontrolni parametri | rpart.control() | minsplit | Postavite minimalni broj promatranja u čvoru prije nego što algoritam izvrši dijeljenje | |
| minbucket | Postavite minimalni broj opažanja u terminalnom čvoru, tj. listu | ||||
| najveća dubina | Postavite maksimalnu dubinu bilo kojeg čvora konačnog stabla. Korijenski čvor se tretira kao dubina 0. | ||||
| rpart | Model vlaka s kontrolnim parametrom | rpart() | formula, df, metoda, kontrola |
Napomena: obučite model na podacima za obuku i testirajte izvedbu na nevidljivom skupu podataka, tj. testnom skupu.




