Дерево решений в R: дерево классификации с примером

⚡ Умное резюме

В R деревья решений разбивают данные на ветви, используя простые правила «да» или «нет», пока каждый лист не будет содержать один доминирующий класс. В этом пошаговом руководстве показано, как построить, построить график, оценить и настроить дерево классификации rpart на наборе данных о выживаемости на «Титанике».

  • ???? Основное определение: Дерево рекурсивно разделяет пространство предикторов, выбирая в каждом узле такое разделение, которое наиболее эффективно уменьшает неоднородность классов.
  • 🔄 Подготовка данных: Перемешайте упорядоченный файл данных о "Титанике" с помощью функции sample(), удалите столбцы с идентификаторами, преобразуйте факторы, а затем удалите строки с отсутствующими значениями (NA).
  • 🇧🇷 Синтаксис модели: Вызовите rpart(survived~., data = data_train, method = 'class') и отобразите результат с помощью rpart.plot(fit, extra = 106).
  • ⚙️ Настройка: Функция rpart.control() предоставляет доступ к параметрам minsplit, minbucket, maxdepth и cp, повышая точность до 79.90 процента.
  • 📊 Оценка: Создайте матрицу ошибок с помощью функции table(), затем разделите диагональ на общую сумму, чтобы получить точность теста 78.47 процента.

Дерево решений в R

Что такое деревья решений?

Деревья решений Деревья решений — это универсальные алгоритмы машинного обучения, способные выполнять как задачи классификации, так и регрессии. Это очень мощные алгоритмы, способные обрабатывать сложные наборы данных. Кроме того, деревья решений являются фундаментальными компонентами случайных лесов, которые входят в число наиболее мощных алгоритмов машинного обучения, доступных сегодня.

Прежде чем создавать дерево в коде, полезно понять, как дерево определяет места разделения.

Как работает дерево решений?

Дерево решений преобразует набор данных в блок-схему вопросов типа «да» или «нет», построенную из трех типов узлов: корень содержит все данные о тренировках. внутренний узел задает вопрос об одном предикторе и разделяет данные на две части, а затем лист Прекращает разделение и возвращает класс большинства.

Рост происходит по жадному сценарию, называемому рекурсивное бинарное разбиение:

  1. Проанализируйте результаты голосования по каждому кандидату. Для каждого предиктора и порогового значения измерьте, насколько «нечистыми» будут две получившиеся группы.
  2. Оставьте лучший. Вопрос, который задается на этом узле, сводится к тому, какое разделение позволяет наиболее эффективно уменьшить количество примесей.
  3. Повторите то же самое для каждого ребенка. до тех пор, пока это не будет остановлено правилом управления: minsplit, minbucket, maxdepth или cp.
  4. Чернослив. Затем cp обрезает ветви, которые не приносят прибыли, что препятствует запоминанию деревом обучающего набора данных.

Поскольку в каждом вопросе сравнивается одна переменная с пороговым значением, алгоритму никогда не требуется масштабирование или ввод фиктивных данных.

Индекс Джини против энтропии в деревьях решений

Эту примесь можно измерить двумя способами, и rpart позволяет выбрать один из них.

Критерии Индекс Джини Энтропия (прирост информации)
Формула 1 – сумма квадратов долей классов - сумма p, умноженная на log2(p)
Диапазон (два класса) 0 - 0.5 0 - 1
Вычисление Быстрее, без логарифмов Более медленный, использует логарифмы.
rpart настройка По умолчанию parms = list(split = “information”)
fit_entropy <- rpart(survived~., data = data_train, method = 'class',
    parms = list(split = "information"))

На практике оба критерия в большинстве случаев выбирают одно и то же разделение, поэтому индекс Джини по умолчанию является безопасным выбором.

Преимущества и недостатки деревьев решений

Компромиссы подскажут, когда достаточно одного дерева, а когда следует перейти к ансамблю.

Преимущества

  • Полностью поддается интерпретации: Представленная модель представляет собой диаграмму, понятную любому заинтересованному лицу.
  • Минимальная предварительная обработка: Масштабирование или нормализация не требуются, и коэффициенты работают изначально.
  • Выполняет обе задачи: Метод `method = 'class'` позволяет построить классификатор, а метод `method = 'anova'` — дерево регрессии.
  • Быстрое обучение: Большие наборы данных обрабатываются за считанные секунды, поэтому деревья решений представляют собой полезный первый базовый уровень.

Недостатки

  • Высокая дисперсия: Небольшое изменение в обучающих данных может привести к построению совершенно другого дерева.
  • Склонность к переобучению: Неограниченное дерево растет до тех пор, пока каждый лист не станет чистым, если только параметры cp и maxdepth не ограничивают его рост.
  • Только разбиения, параллельные осям: Для диагональных границ требуется множество ступенчатых разрезов.

Решением проблемы первых двух недостатков является усреднение количества деревьев, что и является целью данного подхода. случайный лес делает.

Как обучить и визуализировать дерево решений в R

Для построения первого дерева решений в R вам потребуется выполнить семь шагов:

  • Шаг 1. Импортируйте данные
  • Шаг 2. Очистите набор данных
  • Шаг 3. Создайте набор поездов/тестов
  • Шаг 4: Постройте модель
  • Шаг 5: Сделайте прогноз
  • Шаг 6. Измерьте эффективность
  • Шаг 7: Настройте гиперпараметры

Шаг 1) Импортируйте данные

Если вам интересно узнать о судьбе Титаника, вы можете посмотреть это видео на Видео. Цель этого набора данных — предсказать, какие люди с большей вероятностью выживут после столкновения с айсбергом. Набор данных содержит 13 переменных и 1309 наблюдений. Набор данных упорядочен по переменной X.

set.seed(678)
path <- 'https://raw.githubusercontent.com/guru99-edu/R-Programming/master/titanic_data.csv'
titanic <-read.csv(path)
head(titanic)

Выход:

##   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)

Выход:

##         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

Из вывода головы и хвоста вы можете заметить, что данные не перемешиваются. Это большая проблема! Когда вы разделите свои данные между набором поездов и набором тестов, вы выберете Важно пассажир из класса 1 и 2 (ни один пассажир из класса 3 не входит в верхние 80 процентов наблюдений), что означает, что алгоритм никогда не увидит особенности пассажира класса 3. Эта ошибка приведет к плохому прогнозу.

Чтобы решить эту проблему, вы можете использовать функцию sample().

shuffle_index <- sample(1:nrow(titanic))
head(shuffle_index)

Дерево решений R-код Объяснение

  • sample(1:nrow(titanic)): генерирует случайный список индексов от 1 до 1309 (т.е. максимальное количество строк).

Выход:

## [1]  288  874 1078  633  887  992

Вы будете использовать этот индекс для перетасовки набора данных Titanic.

titanic <- titanic[shuffle_index, ]
head(titanic)

Выход:

##         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	

Шаг 2) Очистите набор данных

Некоторые переменные содержат значения NA. Очистка выполняется в три этапа:

  • Удалите переменные home.dest, cabin, name, X и ticket.
  • Создайте факторные переменные для pclass и выжили
  • Отбросьте 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 объяснение

  • select(-c(home.dest, каюта, имя, X, билет)): удалить ненужные переменные.
  • pclass = factor(pclass, levels = c(1,2,3), labels= c('Upper', 'Middle', 'Lower')): Добавляет метку к переменной pclass. 1 становится Upper, 2 становится Middle, а 3 становится Lower
  • factor(survived, levels = c(0,1), labels = c('No', 'Yes')): Добавляет метки к переменной survived. 0 становится No, а 1 становится Yes.
  • na.omit(): удалить наблюдения NA.

Выход:

## 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...		

Шаг 3) Создайте набор поездов/тестов

Прежде чем обучать свою модель, вам необходимо выполнить два шага:

  • Создайте набор поездов и тестов: вы обучаете модель на наборе поездов и проверяете прогноз на тестовом наборе (т. е. невидимые данные).
  • Установите rpart.plot из консоли

Обычной практикой является разделение данных 80/20: 80 процентов данных служат для обучения модели, а 20 процентов — для прогнозирования. Вам необходимо создать два отдельных фрейма данных. Не стоит прикасаться к тестовому набору, пока не закончите построение модели. Вы можете создать функцию create_train_test() с тремя аргументами.

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 объяснение

  • function(data, size=0.8, train = TRUE): добавьте аргументы в функцию.
  • n_row = nrow(data): подсчитать количество строк в наборе данных.
  • total_row = size*n_row: возвращает n-ю строку для построения набора поездов.
  • train_sample <- 1:total_row: выберите строку от первой до n-й.
  • if (train ==TRUE){ } else { }: если условие установлено в true, вернуть набор поездов, иначе — тестовый набор.

Вы можете протестировать свою функцию и проверить размерность.

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)

Выход:

## [1] 836   8
dim(data_test)

Выход:

## [1] 209   8

Обучающий набор данных содержит 836 строк и 8 столбцов, а тестовый набор данных — 209 строк и те же 8 столбцов.

Вы используете функцию prop.table() в сочетании с table(), чтобы проверить правильность процесса рандомизации.

prop.table(table(data_train$survived))

Выход:

##
##        No       Yes 
## 0.5944976 0.4055024
prop.table(table(data_test$survived))

Выход:

## 
##        No       Yes 
## 0.5789474 0.4210526

В обоих наборах данных количество выживших одинаково, около 40 процентов.

Установите rpart.plot

rpart.plot недоступен в библиотеках conda. Установить его можно из консоли:

install.packages("rpart.plot")

Шаг 4) Постройте модель

Вы готовы построить модель. Синтаксис функции дерева решений rpart() следующий:

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	

Вы используете метод класса, потому что прогнозируете класс.

library(rpart)
library(rpart.plot)
fit <- rpart(survived~., data = data_train, method = 'class')
rpart.plot(fit, extra = 106)

Code объяснение

  • rpart(): функция, соответствующая модели. Аргументы:
    • выжил ~.: Формула деревьев решений
    • данные = data_train: Набор данных
    • метод = 'класс': подходит для двоичной модели.
  • rpart.plot(fit, extra= 106): Строит дерево. Аргумент extra установлен на 106, что отображает вероятность второго класса плюс процент наблюдений в каждом узле. Вы можете обратиться к виньетка для получения дополнительной информации о других вариантах.

Выход:

 Построение модели деревьев решений в R

Вы начинаете с корневого узла, в верхней части графа, на глубине 0 из 3:

  1. Вверху — общая вероятность выживания. Он показывает долю пассажиров, выживших в катастрофе. 41 процент пассажиров выжил.
  2. Этот узел запрашивает информацию о том, является ли пассажир мужчиной. Если да, то вы переходите к левому дочернему узлу корня (глубина 1). 63 процента пассажиров — мужчины, вероятность выживания составляет 21 процент.
  3. Во втором узле вы спрашиваете, старше ли пассажир мужского пола 3.5 лет. Если да, то шанс выжить составляет 19 процентов.
  4. Вы продолжаете в том же духе, чтобы понять, какие особенности влияют на вероятность выживания.

Обратите внимание, что одним из многих качеств деревьев решений является то, что они требуют минимальной подготовки данных. В частности, они не требуют масштабирования или центрирования объектов.

По умолчанию функция rpart() использует Джини Для выбора каждого варианта разделения используется показатель неоднородности. Чем выше значение коэффициента Джини, тем более смешаны классы внутри данного узла, поэтому алгоритм всегда выбирает тот вариант разделения, который максимально снижает этот показатель.

Шаг 5) Сделайте прогноз

Вы можете спрогнозировать свой набор тестовых данных. Чтобы сделать прогноз, вы можете использовать функцию Predict(). Основной синтаксис прогнозирования для дерева решений R:

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	

Теперь вам нужно предсказать для каждого из 209 пассажиров в тестовой выборке, ожидает ли модель, что они выживут после столкновения.

predict_unseen <-predict(fit, data_test, type = 'class')

Code объяснение

  • Predict(fit, data_test, type = 'class'): предсказать класс (0/1) набора тестов.

Теперь сравним предсказанные классы с истинными результатами.

table_mat <- table(data_test$survived, predict_unseen)
table_mat

Code объяснение

  • table(data_test$survived, predict_unseen): Создать таблицу сопряженности прогнозируемых классов относительно истинного результата

Выход:

##      predict_unseen
##        No Yes
##   No  106  15
##   Yes  30  58

Строки содержат фактические значения, столбцы — прогнозы. Модель правильно определила 106 умерших и 58 выживших, но отнесла 15 умерших к выжившим, а 30 выживших — к умершим.

Шаг 6) Измерьте производительность

Вы можете вычислить меру точности для задачи классификации с помощью матрица путаницы:

матрица путаницы является лучшим выбором для оценки эффективности классификации. Общая идея состоит в том, чтобы подсчитать, сколько раз истинные экземпляры классифицируются как ложные.

Измерение производительности деревьев решений в R

Каждая строка матрицы ошибок представляет собой фактическую цель, а каждый столбец — прогнозируемую цель. Первая строка этой матрицы рассматривает пассажиров, которые погибли (отрицательный класс): 106 были правильно классифицированы как погибшие (Истинный отрицательный), при этом 15 человек были ошибочно отнесены к числу выживших (Ложно положительныйВо второй строке указаны выжившие: 58 человек были правильно идентифицированы.Истинный положительный), при этом 30 были пропущены (Ложный негатив).

Вы можете вычислить проверка точности из матрицы путаницы:

Измерение производительности деревьев решений в R

Это доля истинно положительных и истинно отрицательных значений в сумме матрицы. С помощью R вы можете написать следующий код:

accuracy_Test <- sum(diag(table_mat)) / sum(table_mat)

Code объяснение

  • sum(diag(table_mat)): Сумма диагонали
  • sum(table_mat): сумма матрицы.

Вы можете распечатать точность тестового набора:

print(paste('Accuracy for test', accuracy_Test))

Выход:

## [1] "Accuracy for test 0.784688995215311"

Точность на тестовом наборе данных составляет 0.7847, то есть 78.47 процента. Повторите упражнение на обучающем наборе данных, чтобы увидеть, насколько модель переобучается.

Шаг 7) Настройте гиперпараметры

Дерево решений в R имеет различные параметры, которые контролируют аспекты соответствия. В библиотеке дерева решений rpart вы можете управлять параметрами с помощью функции rpart.control(). В следующем коде вы вводите параметры, которые будете настраивать. Вы можете обратиться к виньетка для других параметров.

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

Мы будем действовать следующим образом:

  • Создайте функцию для возврата точности
  • Настройте максимальную глубину
  • Настройте минимальное количество выборок, которое должен иметь узел, прежде чем он сможет разделить
  • Настройте минимальное количество выборок, которое должен иметь листовой узел.

Вы можете написать функцию для отображения точности. Вы просто переносите код, который использовали ранее:

  1. предсказать: предсказать_невидимый <- предсказать (подходит, data_test, тип = 'класс')
  2. Создать таблицу: table_mat <- table(data_test$survived, Predict_unseen)
  3. Точность вычислений: Precision_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
}

Теперь настройте параметры и посмотрите, сможете ли вы улучшить результаты по сравнению с моделью по умолчанию. Напоминаем, что вам нужно достичь точности 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)

Выход:

## [1] 0.7990431

Со следующим параметром:

minsplit = 4
minbucket = round(5/3)
maxdepth = 3
cp = 0

Точность повышается с 0.7847 до 0.7990, поэтому оптимизированное дерево превосходит конфигурацию по умолчанию примерно на 1.4 процентных пункта.

Деревья решений в R: Краткий справочник функций

В таблице ниже перечислены все функции, используемые на семи описанных выше этапах, а также пакет, который их предоставляет, и параметры, которые он ожидает. R.

Библиотека Цель Функция Класс Параметры Описание
часть Дерево классификации поездов в R рчасть() класс формула, df, метод
часть Обучить дерево регрессии рчасть() анова формула, df, метод
часть Постройте деревья rpart.plot() подогнанная модель
Использование темпера с изогнутым основанием предсказывать предсказать, () класс встроенная модель, тип
Использование темпера с изогнутым основанием предсказывать предсказать, () проблема встроенная модель, тип
Использование темпера с изогнутым основанием предсказывать предсказать, () вектор встроенная модель, тип
часть Параметры контроля rpart.control() минсплит Установите минимальное количество наблюдений в узле, прежде чем алгоритм выполнит разделение.
минбакет Установите минимальное количество наблюдений в конечном узле, то есть в листовом узле.
Максимальная глубина Установите максимальную глубину любого узла итогового дерева. Корневой узел рассматривается как узел нулевой глубины.
часть Модель поезда с управляющим параметром рчасть() формула, df, метод, контроль

Примечание. Обучите модель на обучающих данных и проверьте производительность на невидимом наборе данных, то есть на тестовом наборе.

Часто задаваемые вопросы (FAQ)

Оба алгоритма подходят как для построения деревьев классификации, так и для регрессии. Функция rpart() реализует алгоритм CART со встроенной перекрестной проверкой и обрезкой с помощью cp, а также используется в паре с rpart.plot для наглядного построения диаграмм, что делает ее более распространенным выбором.

Параметр cp задает минимальное улучшение, которое должно обеспечить разделение, чтобы его можно было сохранить. Большие значения приводят к агрессивной обрезке и образованию более мелких деревьев. Используйте printcp() и plotcp() для поиска значения с наименьшей ошибкой перекрестной проверки.

Да. Функция rpart() использует суррогатное разбиение для направления наблюдений с отсутствующими предикторами по наиболее похожей ветви. В этом руководстве вместо этого используется функция na.omit(), исключительно для упрощения примера набора данных.

Деревья решений лежат в основе объяснимого искусственного интеллекта в кредитовании, страховании и здравоохранении, где регулирующий орган может потребовать точного обоснования принятого решения. Они также являются базовыми алгоритмами обучения в моделях градиентного бустинга и случайного леса.

Да. Искусственный интеллект может переводить правила разделения на простой язык, предлагать значения cp для тестирования и отмечать переобучение в выводе функции printcp(). Перед применением проверяйте каждое предложение на основе собственных результатов перекрестной проверки.

Подведем итог этой публикации следующим образом: