Árbol de decisión en R: Árbol de clasificación con ejemplo

⚡ Resumen inteligente

Los árboles de decisión en R dividen los datos en ramas mediante reglas simples de sí o no hasta que cada hoja contiene una única clase dominante. Este tutorial muestra cómo construir, graficar, evaluar y ajustar un árbol de clasificación rpart en el conjunto de datos de supervivencia del Titanic.

  • ???? Definición básica: Un árbol divide recursivamente el espacio de predicción, eligiendo en cada nodo la división que más reduce la impureza de clase.
  • 🔄 Preparación de datos: Mezcla el archivo Titanic ordenado con sample(), elimina las columnas de identificadores, convierte los factores y luego elimina las filas NA.
  • 🛠️ Sintaxis del modelo: Llama a rpart(survived~., data = data_train, method = 'class') y renderiza el resultado con rpart.plot(fit, extra = 106).
  • ⚙️ Sintonización: rpart.control() expone minsplit, minbucket, maxdepth y cp, lo que eleva la precisión al 79.90 por ciento.
  • 📊 Evaluación: Construye una matriz de confusión con table(), luego divide la diagonal por el total para obtener una precisión de prueba del 78.47 por ciento.

Árbol de decisión en R

¿Qué son los árboles de decisión?

Árboles de decisión Los árboles de decisión son un algoritmo de aprendizaje automático versátil que puede realizar tareas de clasificación y regresión. Son algoritmos muy potentes, capaces de ajustarse a conjuntos de datos complejos. Además, los árboles de decisión son componentes fundamentales de los bosques aleatorios, que se encuentran entre los algoritmos de aprendizaje automático más potentes disponibles en la actualidad.

Antes de construir uno mediante código, resulta útil saber cómo decide un árbol dónde dividirse.

¿Cómo funciona un árbol de decisiones?

Un árbol de decisión convierte un conjunto de datos en un diagrama de flujo de preguntas de sí o no construido a partir de tres tipos de nodos: el raíz sostiene cada observación de entrenamiento, un nodo interno formula una pregunta sobre un predictor y divide los datos en dos, y un hoja Deja de dividir y devuelve la clase mayoritaria.

El crecimiento sigue un procedimiento voraz llamado particionamiento binario recursivo:

  1. Evaluar cada división de candidatos. Para cada predictor y punto de corte, mida el grado de impureza de los dos grupos resultantes.
  2. Quédate con el mejor. La división que reduce más las impurezas se convierte en la pregunta que se plantea en ese nodo.
  3. Repita el procedimiento con cada niño. hasta que una regla de control lo detenga: minsplit, minbucket, maxdepth o cp.
  4. Ciruela pasa. A continuación, cp poda las ramas que no generan ingresos, lo que impide que el árbol memorice el conjunto de entrenamiento.

Dado que cada pregunta compara una variable con un umbral, el algoritmo nunca necesita escalado ni codificación ficticia.

Índice de Gini frente a entropía en árboles de decisión

Esa impureza se puede medir de dos maneras, y rpart te permite elegir.

Criterios Índice de Gini Entropía (ganancia de información)
Fórmula 1 – suma de las proporciones de clase al cuadrado -suma de p veces log2(p)
Gama (dos clases) 0 a 0.5 0 a 1
Cálculo Más rápido, sin logaritmo Más lento, utiliza logaritmos.
Configuración de rpart Predeterminado parámetros = lista(split = “información”)
fit_entropy <- rpart(survived~., data = data_train, method = 'class',
    parms = list(split = "information"))

En la práctica, ambos criterios suelen coincidir en la misma división, por lo que el índice de Gini predeterminado es una opción segura.

Ventajas y desventajas de los árboles de decisión

Las ventajas y desventajas te indican cuándo un solo árbol es suficiente y cuándo conviene optar por un conjunto de árboles.

Ventajas

  • Totalmente interpretable: El modelo ajustado es un diagrama que cualquier interesado puede interpretar.
  • Preprocesamiento mínimo: No se requiere escalado ni normalización, y los factores funcionan de forma nativa.
  • Realiza ambas tareas: El método = 'class' ajusta un clasificador y el método = 'anova' ajusta un árbol de regresión.
  • Rápido para entrenar: Los conjuntos de datos grandes caben en segundos, por lo que los árboles constituyen una primera línea de base útil.

Desventajas

  • Alta varianza: Un pequeño cambio en los datos de entrenamiento puede producir un árbol completamente diferente.
  • Propenso al sobreajuste: Un árbol sin restricciones crece hasta que cada hoja es pura, a menos que cp y maxdepth lo restrinjan.
  • Solo divisiones paralelas a los ejes: Los límites diagonales requieren muchos cortes en forma de escalera.

El remedio para las dos primeras debilidades es promediar muchos árboles, que es lo que un bosque al azar hace.

Cómo entrenar y visualizar un árbol de decisión en R

Para construir tu primer árbol de decisión en R, deberás seguir siete pasos:

  • Paso 1: importar los datos
  • Paso 2: limpiar el conjunto de datos
  • Paso 3: Crear tren/conjunto de prueba
  • Paso 4: construye el modelo
  • Paso 5: haz una predicción
  • Paso 6: medir el rendimiento
  • Paso 7: ajuste los hiperparámetros

Paso 1) Importar los datos

Si tienes curiosidad sobre el destino del Titanic, puedes ver este vídeo en YouTube. El propósito de este conjunto de datos es predecir qué personas tienen más probabilidades de sobrevivir después de la colisión con el iceberg. El conjunto de datos contiene 13 variables y 1309 observaciones. El conjunto de datos está ordenado por la variable X.

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

Salida:

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

Salida:

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

Desde la salida inicial y final, puede notar que los datos no están mezclados. ¡Este es un gran problema! Cuando divida sus datos entre un conjunto de tren y un conjunto de prueba, seleccionará único el pasajero de las clases 1 y 2 (ningún pasajero de la clase 3 está en el 80 por ciento superior de las observaciones), lo que significa que el algoritmo nunca verá las características del pasajero de la clase 3. Este error conducirá a una mala predicción.

Para superar este problema, puede utilizar la función sample().

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

Árbol de decisión Código R Explicación

  • muestra (1: nrow (titanic)): genera una lista aleatoria de índices del 1 al 1309 (es decir, el número máximo de filas).

Salida:

## [1]  288  874 1078  633  887  992

Utilizará este índice para mezclar el conjunto de datos del Titanic.

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

Salida:

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

Paso 2) Limpiar el conjunto de datos

Varias variables contienen valores NA. La limpieza se realiza en tres partes:

  • Elimine las variables home.dest, cabin, name, X y ticket
  • Crear variables de factor para pclass y sobrevivió
  • Deja la 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 Explicación

  • select(-c(home.dest, cabina, nombre, X, boleto)): Elimina variables innecesarias
  • pclass = factor(pclass, levels = c(1,2,3), labels= c('Superior', 'Media', 'Inferior')): Agrega una etiqueta a la variable pclass. 1 se convierte en Superior, 2 en Media y 3 en Inferior
  • factor(sobrevivió, niveles = c(0,1), etiquetas = c('No', 'Sí')): Añade etiquetas a la variable sobrevivió. 0 se convierte en No y 1 en Sí
  • na.omit(): Elimina las observaciones de NA

Salida:

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

Paso 3) Crear tren/conjunto de prueba

Antes de entrenar su modelo, debe realizar dos pasos:

  • Cree un tren y un conjunto de prueba: entrene el modelo en el conjunto de trenes y pruebe la predicción en el conjunto de prueba (es decir, datos no vistos).
  • Instale rpart.plot desde la consola

La práctica común es dividir los datos 80/20, el 80 por ciento de los datos sirve para entrenar el modelo y el 20 por ciento para hacer predicciones. Necesita crear dos marcos de datos separados. No querrás tocar el conjunto de prueba hasta que termines de construir tu modelo. Puede crear un nombre de función create_train_test() que tome tres argumentos.

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 Explicación

  • función (datos, tamaño = 0.8, tren = VERDADERO): agrega los argumentos en la función
  • n_row = nrow(data): cuenta el número de filas en el conjunto de datos
  • total_row = size*n_row: Devuelve la enésima fila para construir el conjunto de trenes
  • train_sample <- 1:total_row: seleccione la primera fila hasta la enésima fila
  • if (train ==TRUE){ } else { }: si la condición se establece en verdadera, devuelve el conjunto de tren; de lo contrario, el conjunto de prueba.

Puede probar su función y verificar la dimensión.

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)

Salida:

## [1] 836   8
dim(data_test)

Salida:

## [1] 209   8

El conjunto de datos de entrenamiento tiene 836 filas y 8 columnas, mientras que el conjunto de datos de prueba tiene 209 filas y las mismas 8 columnas.

Utiliza la función prop.table() combinada con table() para verificar si el proceso de aleatorización es correcto.

prop.table(table(data_train$survived))

Salida:

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

Salida:

## 
##        No       Yes 
## 0.5789474 0.4210526

En ambos conjuntos de datos, la cantidad de supervivientes es la misma, alrededor del 40 por ciento.

Instalar rpart.plot

rpart.plot no está disponible en las bibliotecas conda. Puedes instalarlo desde la consola:

install.packages("rpart.plot")

Paso 4) Construye el modelo

Ya puedes construir el modelo. La sintaxis para la función de árbol de decisión rpart() es:

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	

Utiliza el método de clase porque predice una clase.

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

Code Explicación

  • rpart(): Función para ajustar el modelo. Los argumentos son:
    • sobrevivió ~.: Fórmula de los árboles de decisión
    • datos = data_train: conjunto de datos
    • método = 'clase': Ajustar un modelo binario
  • rpart.plot(fit, extra= 106): Grafica el árbol. El argumento extra se establece en 106, que muestra la probabilidad de la segunda clase más el porcentaje de observaciones en cada nodo. Puede consultar el viñeta para obtener más información sobre las otras opciones.

Salida:

 Construya un modelo de árboles de decisión en R

Empiezas en el nodo raíz, en la parte superior del gráfico y en la profundidad 0 de 3:

  1. En la parte superior está la probabilidad global de supervivencia. Muestra la proporción de pasajeros que sobrevivieron al accidente. El 41 por ciento de los pasajeros sobrevivió.
  2. Este nodo pregunta si el género del pasajero es masculino. Si es así, se baja al hijo izquierdo de la raíz (profundidad 1). El 63 por ciento son hombres con una probabilidad de supervivencia del 21 por ciento.
  3. En el segundo nodo se pregunta si el pasajero varón tiene más de 3.5 años. En caso afirmativo, la probabilidad de supervivencia es del 19 por ciento.
  4. Continúe así para comprender qué características afectan la probabilidad de supervivencia.

Tenga en cuenta que una de las muchas cualidades de los árboles de decisión es que requieren muy poca preparación de datos. En particular, no requieren escalado ni centrado de funciones.

Por defecto, la función rpart() utiliza el Gini Medida de impureza para elegir cada división. Cuanto mayor sea el valor de Gini, más mezcladas estarán las clases dentro de ese nodo, por lo que el algoritmo siempre elige la división que lo reduce más.

Paso 5) Haz una predicción

Puede predecir su conjunto de datos de prueba. Para hacer una predicción, puede utilizar la función predict(). La sintaxis básica de predicción para el árbol de decisión R es:

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	

Ahora debes predecir, para cada uno de los 209 pasajeros del conjunto de prueba, si el modelo prevé que sobrevivirán a la colisión.

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

Code Explicación

  • predecir (ajuste, prueba_datos, tipo = 'clase'): predice la clase (0/1) del conjunto de prueba

Ahora compare las clases pronosticadas con los resultados reales.

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

Code Explicación

  • tabla(data_test$survived, predict_unseen): Construye una tabla de contingencia de las clases predichas frente al resultado real.

Salida:

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

Las filas representan los valores reales y las columnas, las predicciones. El modelo identificó correctamente a 106 personas que no sobrevivieron y a 58 que sí lo hicieron, pero etiquetó a 15 personas que no sobrevivieron como supervivientes y a 30 supervivientes como fallecidos.

Paso 6) Medir el desempeño

Puede calcular una medida de precisión para la tarea de clasificación con el matriz de confusión:

El matriz de confusión es una mejor opción para evaluar el rendimiento de la clasificación. La idea general es contar el número de veces que las instancias Verdaderas se clasifican como Falsas.

Medir el rendimiento de los árboles de decisión en R

Cada fila de una matriz de confusión representa un objetivo real, mientras que cada columna representa un objetivo predicho. La primera fila de esta matriz considera a los pasajeros que murieron (la clase negativa): 106 fueron clasificados correctamente como muertos (Verdadero-negativo), mientras que 15 fueron clasificados erróneamente como supervivientes (Falso positivo). La segunda fila considera a los supervivientes: 58 fueron identificados correctamente (Verdadero positivo), mientras que 30 quedaron sin registrar (Falso negativo).

Puede calcular el prueba de precisión de la matriz de confusión:

Medir el rendimiento de los árboles de decisión en R

Es la proporción de verdaderos positivos y verdaderos negativos sobre la suma de la matriz. Con R, puedes codificar de la siguiente manera:

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

Code Explicación

  • sum(diag(table_mat)): Suma de la diagonal
  • sum(table_mat): Suma de la matriz.

Puede imprimir la precisión del conjunto de prueba:

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

Salida:

## [1] "Accuracy for test 0.784688995215311"

La precisión en el conjunto de prueba es de 0.7847, es decir, el 78.47 por ciento. Repita el ejercicio en el conjunto de entrenamiento para ver cuánto sobreajusta el modelo.

Paso 7) Ajusta los hiperparámetros

El árbol de decisión en R tiene varios parámetros que controlan aspectos del ajuste. En la biblioteca de árboles de decisión de rpart, puede controlar los parámetros mediante la función rpart.control(). En el siguiente código, introduce los parámetros que ajustará. Puede consultar la viñeta para otros parámetros.

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

Procederemos de la siguiente manera:

  • Construir función para devolver precisión
  • Ajusta la profundidad máxima
  • Ajuste la cantidad mínima de muestra que debe tener un nodo antes de poder dividirse
  • Ajuste el número mínimo de muestras que debe tener un nodo hoja

Puede escribir una función para mostrar la precisión. Simplemente envuelve el código que usaste antes:

  1. predecir: predecir_unseen <- predecir (ajuste, prueba_datos, tipo = 'clase')
  2. Producir tabla: table_mat <- table(data_test$survived, predict_unseen)
  3. Precisión de cálculo: exactitud_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
}

Ahora ajusta los parámetros y comprueba si puedes mejorar el modelo predeterminado. Recuerda que debes superar una precisión de 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)

Salida:

## [1] 0.7990431

Con el siguiente parámetro:

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

La precisión aumenta de 0.7847 a 0.7990, por lo que el árbol optimizado supera a la configuración predeterminada en aproximadamente 1.4 puntos porcentuales.

Árboles de decisión en R: Guía rápida de funciones

La tabla a continuación enumera todas las funciones utilizadas en los siete pasos anteriores, junto con el paquete que las proporciona y los parámetros que espera en R.

Biblioteca Objetivo Función Clase Parámetros Detalles
parte Árbol de clasificación de trenes en R parte() clase fórmula, df, método
parte Tren de árbol de regresión parte() anova fórmula, df, método
parte Trazar los árboles rpart.plot() modelo ajustado
bases predecir predecir() clase modelo equipado, tipo
bases predecir predecir() problema modelo equipado, tipo
bases predecir predecir() vector modelo equipado, tipo
parte Parámetros de control rpart.control() división mínima Establezca el número mínimo de observaciones en el nodo antes de que el algoritmo realice una división
minbucket Establezca el número mínimo de observaciones en un nodo terminal, es decir, la hoja.
máxima profundidad Establece la profundidad máxima de cualquier nodo del árbol final. El nodo raíz se trata como profundidad 0.
parte Modelo de tren con parámetro de control. parte() fórmula, df, método, control

Nota: Entrene el modelo con datos de entrenamiento y pruebe el rendimiento en un conjunto de datos invisible, es decir, un conjunto de prueba.

Preguntas Frecuentes

Ambos ajustan árboles de clasificación y regresión. rpart() implementa CART con poda integrada mediante validación cruzada a través de cp y se combina con rpart.plot para diagramas claros, lo que lo convierte en la opción más común.

cp establece la mejora mínima que debe ofrecer una división para conservarse. Valores mayores podan agresivamente y producen árboles más pequeños. Utilice printcp() y plotcp() para encontrar el valor con el menor error de validación cruzada.

Sí. rpart() utiliza divisiones sustitutas para dirigir las observaciones con predictores faltantes hacia la rama más similar. En este tutorial, se utiliza na.omit() simplemente para simplificar el conjunto de datos de ejemplo.

Los árboles de decisión impulsan la IA explicable en los sectores de crédito, seguros y sanidad, donde un organismo regulador puede exigir el razonamiento exacto detrás de una decisión. También constituyen la base de los modelos de aprendizaje automático basados ​​en el aumento del gradiente y los bosques aleatorios.

Sí. Los asistentes de IA pueden traducir las reglas de división a lenguaje sencillo, sugerir valores de cp para probar e indicar el sobreajuste en la salida de printcp(). Verifica cada sugerencia comparándola con tus propios resultados de validación cruzada antes de implementarla.

Resumir este post con: