Tutorial de Random Forest en R: Algoritmo con ejemplo

โšก Resumen inteligente

El algoritmo Random Forest en R construye cientos de รกrboles de decisiรณn a partir de muestras de arranque y promedia sus votos para obtener una predicciรณn robusta. Este tutorial ajusta mtry, maxnodes y ntree con caret y evalรบa el modelo final con los datos de supervivencia del Titanic.

  • ???? Principio bรกsico: El mรฉtodo Bagging entrena cada รกrbol con una muestra bootstrap y un subconjunto aleatorio de predictores, de modo que los errores individuales se cancelan durante la votaciรณn mayoritaria.
  • ๐Ÿงช Configuraciรณn de validaciรณn: trainControl(method = โ€œcvโ€, number = 10, search = โ€œgridโ€) fija una bรบsqueda en cuadrรญcula de diez pliegues que se reutiliza en cada paso de ajuste.
  • ๐ŸŽฏ Orden de afinaciรณn: Optimiza primero mtry, luego maxnodes y despuรฉs ntree, llevando el valor ganador en cada etapa.
  • ๐Ÿ“ˆ Mejores configuraciones: Los valores mtry = 4, maxnodes = 24 y ntree = 800 produjeron la mayor precisiรณn de validaciรณn cruzada en este conjunto de datos.
  • ๐Ÿงฎ Resultado de la prueba: La funciรณn confusionMatrix() informa una precisiรณn del 79.43 por ciento, una sensibilidad de 0.9091 y una especificidad de 0.6364 en los datos reservados.

Bosque aleatorio en R

ยฟQuรฉ es el bosque aleatorio en R?

Los bosques aleatorios se basan en una idea simple: "la sabidurรญa de la multitud". La suma de los resultados de mรบltiples predictores da una mejor predicciรณn que el mejor predictor individual. Un grupo de predictores se llama junto. Por eso esta tรฉcnica se llama Ensemble Learning.

En un tutorial anterior, aprendiste a usar รกrboles de decisiรณn para hacer una predicciรณn binaria. Para mejorar nuestra tรฉcnica, podemos entrenar a un grupo de Clasificadores de รกrbol de decisiรณn, cada uno en un subconjunto aleatorio diferente del conjunto de entrenamiento. Para hacer una predicciรณn, se recopilan las predicciones de todos los รกrboles individuales y se devuelve la clase que recibe la mayor cantidad de votos. Esta tรฉcnica se llama Bosque al azar.

Antes de escribir cualquier cรณdigo, resulta รบtil ver exactamente cรณmo se forma el bosque a partir de esos รกrboles individuales.

ยฟCรณmo funciona el algoritmo Random Forest en R?

Comprender la mecรกnica facilita el razonamiento sobre cada parรกmetro de ajuste. Un bosque aleatorio se construye en cuatro etapas.

  1. Bootstrap toma de muestras. El algoritmo extrae una muestra aleatoria de filas. con reemplazo del conjunto de entrenamiento para cada รกrbol. Aproximadamente un tercio de las filas se dejan fuera de cada muestra; estas son las observaciones fuera de la bolsa (OOB).
  2. Selecciรณn aleatoria de caracterรญsticas. En cada divisiรณn, solo se considera un subconjunto aleatorio de predictores. El tamaรฑo de ese subconjunto es el mtry parรกmetro. Restringir la elecciรณn es lo que evita que todos los รกrboles se parezcan.
  3. Crecimiento de รกrboles sin podar. Cada รกrbol crece hasta que llega a un punto muerto.ping regla como maxnodos or tamaรฑo del nodoSe permite deliberadamente que los รกrboles individuales se sobreajusten, porque sus errores no estรกn correlacionados.
  4. Agregaciรณn. Para la clasificaciรณn, el bosque devuelve la clase con la mayor cantidad de votos; para la regresiรณn, devuelve la predicciรณn promedio. Este paso de promediado es lo que el tรฉrmino harpillera (agregaciรณn bootstrap) describe.

La clave estรก en combinar el muestreo por filas con el muestreo por columnas. Un รบnico รกrbol profundo presenta un sesgo bajo y una varianza muy alta; promediar cientos de ellos mantiene el sesgo bajo a la vez que reduce la varianza.

Error de salida de la bolsa. Debido a que cada observaciรณn se excluye de aproximadamente un tercio de los รกrboles, R puede puntuar cada fila utilizando solo los รกrboles que nunca la vieron. El error OOB resultante es una estimaciรณn de validaciรณn integrada y gratuita que imprime randomForest():

rf_oob <- randomForest(survived~., data = data_train, ntree = 800, mtry = 4)
print(rf_oob)   # reports the OOB estimate of error rate

El error OOB es una comprobaciรณn rรกpida y prรกctica, pero este tutorial utiliza validaciรณn cruzada de diez pliegues mediante caret, de modo que cada cuadrรญcula de ajuste se compara en pliegues idรฉnticos.

Bosque aleatorio frente a รกrbol de decisiรณn en R

Un bosque aleatorio es un conjunto de los mismos รกrboles de decisiรณn Como ya se explicรณ en el tutorial anterior, conviene aclarar las diferencias antes de elegir entre ellos.

Criterios รrbol de decisiรณn Bosque al azar
Estructura Un รกrbol Cientos de รกrboles unidos por votaciรณn
Diferencia Alto, muy sensible a la muestra de entrenamiento Bajo, el promedio cancela los errores individuales
Riesgo de sobreajuste Alto a menos que se pode Bajo incluso con รกrboles sin podar
Interpretabilidad Completamente legible como diagrama de flujo. Solo se puede leer la importancia de las variables.
Costo de formaciรณn Muy rรกpido Proporcional a ntree
Validaciรณn incorporada Ninguna Estimaciรณn del error fuera de la muestra
Funciรณn R parte() Bosque aleatorio()

Elija un solo รกrbol cuando deba explicar el proceso de decisiรณn a un pรบblico no tรฉcnico. Elija un bosque cuando la precisiรณn predictiva sea mรกs importante que un diagrama legible.

Ventajas y desventajas del algoritmo Random Forest

Saber dรณnde el algoritmo es fuerte y dรณnde es dรฉbil te indica cuรกndo merece la pena invertir tiempo en su optimizaciรณn.

Ventajas

  • Precisiรณn sin recorte: El promedio de muchos รกrboles no correlacionados ofrece resultados sรณlidos con muy poca correcciรณn manual.
  • Resistente al sobreajuste: Agregar mรกs รกrboles nunca aumenta el error de generalizaciรณn, por lo que se puede aumentar el valor de ntree sin problemas.
  • Maneja datos mixtos: Los predictores numรฉricos y factoriales funcionan conjuntamente, y no es necesario aplicar escalas.
  • Validaciรณn y posicionamiento gratuitos: El error OOB y varImp() no conllevan ningรบn coste computacional adicional.

Desventajas

  • Predicciones opacas: no se puede tracuna รบnica vรญa de decisiรณn, lo cual es importante en entornos regulados.
  • Lento en bosques extensos: El tiempo de entrenamiento y predicciรณn aumenta linealmente con el nรบmero de รกrboles.
  • Puntuaciones de importancia sesgadas: Las variables categรณricas con muchos niveles pueden parecer mรกs importantes de lo que realmente son.
  • Extrapolaciรณn dรฉbil: En el caso de la regresiรณn, el bosque nunca puede predecir valores fuera del rango observado durante el entrenamiento.

Una vez establecidas la teorรญa y analizadas las ventajas y desventajas, los siguientes seis pasos consisten en construir, ajustar y evaluar un bosque aleatorio de extremo a extremo sobre el conjunto de datos de supervivencia del Titanic.

Paso 1) Importar los datos

Para asegurarse de tener el mismo conjunto de datos que en el tutorial para รกrboles de decisiรณnEl conjunto de trenes y el conjunto de pruebas estรกn alojados en lรญnea. Puedes importarlos sin realizar ningรบn cambio.

library(dplyr)
data_train <- read.csv("https://raw.githubusercontent.com/guru99-edu/R-Programming/master/train.csv")
glimpse(data_train)
data_test <- read.csv("https://raw.githubusercontent.com/guru99-edu/R-Programming/master/test.csv") 
glimpse(data_test)

Paso 2) Entrena el modelo

Una forma de evaluar el rendimiento de un modelo es entrenarlo en varios conjuntos de datos mรกs pequeรฑos y evaluarlos sobre otro conjunto de prueba mรกs pequeรฑo. Esto se llama k-fold validaciรณn cruzada. R Esta funciรณn divide aleatoriamente los datos en k subconjuntos de tamaรฑo casi idรฉntico. Por ejemplo, si k = 10, el modelo se entrena con nueve particiones y se evalรบa con la restante. Este proceso se repite hasta que se hayan evaluado todos los subconjuntos. Esta tรฉcnica se utiliza ampliamente para la selecciรณn de modelos, especialmente cuando el modelo requiere el ajuste de parรกmetros.

Ahora que tenemos una forma de evaluar nuestro modelo, necesitamos decidir quรฉ parรกmetros se generalizan mejor a datos no vistos.

El bosque aleatorio elige un subconjunto aleatorio de caracterรญsticas y construye muchos รกrboles de decisiรณn. El modelo promedia todas las predicciones de los รกrboles de Decisiones.

El algoritmo Random Forest tiene algunos parรกmetros que se pueden modificar para mejorar la generalizaciรณn de la predicciรณn. Utilizarรกs la funciรณn randomForest() para entrenar el modelo.

La sintaxis para randomForest() es:

randomForest(formula, ntree=n, mtry=FALSE, maxnodes = NULL)
Arguments:
- Formula: Formula of the fitted model
- ntree: number of trees in the forest
- mtry: Number of candidate variables drawn at each split. By default, it is the square root of the number of predictors for classification.
- maxnodes: Set the maximum number of terminal nodes each tree can have
- importance=TRUE: Whether independent variables importance in the random forest be assessed

Nota: : El bosque aleatorio se puede entrenar con mรกs parรกmetros. Puedes consultar el viรฑeta para ver los diferentes parรกmetros.

Ajustar un modelo es un trabajo tedioso. Existen muchas combinaciones posibles de parรกmetros. No necesariamente tienes tiempo para probarlas todas. Una buena alternativa es dejar que la mรกquina encuentre la mejor combinaciรณn por ti. Hay dos mรฉtodos disponibles:

  • Bรบsqueda aleatoria
  • Bรบsqueda de cuadrรญcula

Ambos mรฉtodos se definen a continuaciรณn, pero este tutorial entrena el modelo utilizando la bรบsqueda en cuadrรญcula.

Definiciรณn de bรบsqueda de cuadrรญcula

El mรฉtodo de bรบsqueda de cuadrรญcula es simple, el modelo se evaluarรก sobre todas las combinaciones que pase en la funciรณn, mediante validaciรณn cruzada.

Por ejemplo, desea probar el modelo con 10, 20, 30 รกrboles y cada รกrbol se probarรก durante un nรบmero de intentos igual a 1, 2, 3, 4, 5. Luego, la mรกquina probarรก 15 modelos diferentes:

    .mtry ntrees
 1      1     10
 2      2     10
 3      3     10
 4      4     10
 5      5     10
 6      1     20
 7      2     20
 8      3     20
 9      4     20
 10     5     20
 11     1     30
 12     2     30
 13     3     30
 14     4     30
 15     5     30	

El algoritmo evaluarรก:

randomForest(formula, ntree=10, mtry=1)
randomForest(formula, ntree=10, mtry=2)
randomForest(formula, ntree=10, mtry=3)
randomForest(formula, ntree=20, mtry=2)
...

Cada combinaciรณn se evalรบa mediante validaciรณn cruzada. La desventaja de la bรบsqueda en cuadrรญcula es la cantidad de experimentos: aumenta exponencialmente cuando el nรบmero de combinaciones es elevado. Para superar este problema, se puede utilizar la bรบsqueda aleatoria.

Definiciรณn de bรบsqueda aleatoria

La principal diferencia entre la bรบsqueda aleatoria y la bรบsqueda en cuadrรญcula radica en que la bรบsqueda aleatoria no evalรบa todas las combinaciones de hiperparรกmetros en el espacio de bรบsqueda. En cambio, selecciona una combinaciรณn al azar en cada iteraciรณn. La ventaja reside en un coste computacional mucho menor.

Establecer el parรกmetro de control

Procederรก de la siguiente manera para construir y evaluar el modelo:

  • Evaluar el modelo con la configuraciรณn predeterminada.
  • Encuentra el mejor nรบmero de mtry
  • Encuentra el mejor nรบmero de maxnodes
  • Encuentra el mejor nรบmero de ntrees
  • Evaluar el modelo en el conjunto de datos de prueba.

Antes de comenzar con la exploraciรณn de parรกmetros, necesita instalar dos bibliotecas.

  • caret: biblioteca de aprendizaje automรกtico de R. Si usted tiene instalar R con r-esencial. ya esta en la biblioteca
  • e1071: biblioteca de aprendizaje automรกtico R.
    • Anaconda: instalaciรณn de conda -c r r-e1071

Puedes importarlos junto con randomForest:

library(randomForest)
library(caret)
library(e1071)

Configuraciรณn predeterminada

La validaciรณn cruzada K-fold estรก controlada por la funciรณn trainControl()

trainControl(method = "cv", number = n, search ="grid")
arguments
- method = "cv": The method used to resample the dataset. 
- number = n: Number of folds to create
- search = "grid": Use the grid search method. For the randomized method, use "random"
Note: You can refer to the vignette to see the other arguments of the function.

Puede intentar ejecutar el modelo con los parรกmetros predeterminados y ver la puntuaciรณn de precisiรณn.

Nota: : Utilizarรกs los mismos controles durante todo el tutorial.

# Define the control
trControl <- trainControl(method = "cv",
    number = 10,
    search = "grid")

Utilizarรก la biblioteca de intercalaciรณn para evaluar su modelo. La biblioteca tiene una funciรณn llamada train() para evaluar casi todos aprendizaje automรกtico Algoritmo. Dicho de otro modo, puedes usar esta funciรณn para entrenar otros algoritmos.

La sintaxis bรกsica es:

train(formula, df, method = "rf", metric= "Accuracy", trControl = trainControl(), tuneGrid = NULL)
argument
- `formula`: Define the formula of the algorithm
- `method`: Define which model to train. Note, at the end of the tutorial, there is a list of all the models that can be trained
- `metric` = "Accuracy": Define how to select the optimal model
- `trControl = trainControl()`: Define the control parameters
- `tuneGrid = NULL`: Return a data frame with all the possible combination

Vamos a construir el modelo con los valores predeterminados.

set.seed(1234)
# Run the model
rf_default <- train(survived~.,
    data = data_train,
    method = "rf",
    metric = "Accuracy",
    trControl = trControl)
# Print the results
print(rf_default)

Code Explicaciรณn

  • trainControl(method=โ€cvโ€, number=10, search=โ€gridโ€): Evalรบa el modelo con una bรบsqueda en cuadrรญcula sobre 10 pliegues.
  • train(โ€ฆ): Entrena un modelo de bosque aleatorio. El modelo mejorado se elige con la medida de precisiรณn.

Salida:

## Random Forest 
## 
## 836 samples
##   7 predictor
##   2 classes: 'No', 'Yes' 
## 
## No pre-processing
## Resampling: Cross-Validated (10 fold) 
## Summary of sample sizes: 753, 752, 753, 752, 752, 752, ... 
## Resampling results across tuning parameters:
## 
##   mtry  Accuracy   Kappa    
##    2    0.7919248  0.5536486
##    6    0.7811245  0.5391611
##   10    0.7572002  0.4939620
## 
## Accuracy was used to select the optimal model using  the largest value.
## The final value used for the model was mtry = 2.

El algoritmo utiliza 500 รกrboles y probรณ tres valores diferentes de mtry: 2, 6, 10.

El valor final utilizado para el modelo fue mtry = 2, con una precisiรณn de validaciรณn cruzada de 0.792. Intentemos obtener una puntuaciรณn mรกs alta.

Busca el mejor mtry

Puedes probar el modelo con valores de mtry del 1 al 10.

set.seed(1234)
tuneGrid <- expand.grid(.mtry = c(1: 10))
rf_mtry <- train(survived~.,
    data = data_train,
    method = "rf",
    metric = "Accuracy",
    tuneGrid = tuneGrid,
    trControl = trControl,
    importance = TRUE,
    nodesize = 14,
    ntree = 300)
print(rf_mtry)

Code Explicaciรณn

  • tuneGrid <- expand.grid(.mtry = c(1:10)): Construye un vector con valores del 1 al 10.

Salida:

## Random Forest 
## 
## 836 samples
##   7 predictor
##   2 classes: 'No', 'Yes' 
## 
## No pre-processing
## Resampling: Cross-Validated (10 fold) 
## Summary of sample sizes: 753, 752, 753, 752, 752, 752, ... 
## Resampling results across tuning parameters:
## 
##   mtry  Accuracy   Kappa    
##    1    0.7572576  0.4647368
##    2    0.7979346  0.5662364
##    3    0.8075158  0.5884815
##    4    0.8110729  0.5970664
##    5    0.8074727  0.5900030
##    6    0.8099111  0.5949342
##    7    0.8050918  0.5866415
##    8    0.8050918  0.5855399
##    9    0.8050631  0.5855035
##   10    0.7978916  0.5707336
## 
## Accuracy was used to select the optimal model using  the largest value.
## The final value used for the model was mtry = 4.

El mejor valor de mtry es 4. Se almacena en:

rf_mtry$bestTune$mtry

Puede almacenarlo y utilizarlo cuando necesite ajustar otros parรกmetros.

max(rf_mtry$results$Accuracy)

Salida:

## [1] 0.8110729
best_mtry <- rf_mtry$bestTune$mtry 
best_mtry

Salida:

## [1] 4

Paso 3) Busca los mejores maxnodes

Debe crear un bucle para evaluar los diferentes valores de maxnodes. En el cรณdigo siguiente, deberรก:

  • Crear una lista
  • Cree una variable con el mejor valor del parรกmetro mtry; Obligatorio
  • Crea el bucle
  • Almacenar el valor actual de maxnode
  • Resumir los resultados
store_maxnode <- list()
tuneGrid <- expand.grid(.mtry = best_mtry)
for (maxnodes in c(5: 15)) {
    set.seed(1234)
    rf_maxnode <- train(survived~.,
        data = data_train,
        method = "rf",
        metric = "Accuracy",
        tuneGrid = tuneGrid,
        trControl = trControl,
        importance = TRUE,
        nodesize = 14,
        maxnodes = maxnodes,
        ntree = 300)
    current_iteration <- toString(maxnodes)
    store_maxnode[[current_iteration]] <- rf_maxnode
}
results_mtry <- resamples(store_maxnode)
summary(results_mtry)

Code explicaciรณn:

  • store_maxnode <- list(): Los resultados del modelo se almacenarรกn en esta lista
  • expand.grid(.mtry=best_mtry): utilice el mejor valor de mtry
  • para (maxnodes en c(5:15)) { โ€ฆ }: Calcula el modelo con valores de maxnodes de 5 a 15.
  • maxnodes = maxnodes: En cada iteraciรณn, maxnodes es igual al valor actual del bucle, es decir, 5, 6, 7, โ€ฆ
  • current_iteration <- toString(maxnodes): Almacena el valor de maxnodes como una cadena.
  • store_maxnode[[current_iteration]] <- rf_maxnode: Guarda el resultado del modelo en la lista.
  • resamples (store_maxnode): organiza los resultados del modelo
  • resumen (resultados_mtry): imprime el resumen de toda la combinaciรณn.

Salida:

## 
## Call:
## summary.resamples(object = results_mtry)
## 
## Models: 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15 
## Number of resamples: 10 
## 
## Accuracy 
##         Min.   1st Qu.    Median      Mean   3rd Qu.      Max. NA's
## 5  0.6785714 0.7529762 0.7903758 0.7799771 0.8168388 0.8433735    0
## 6  0.6904762 0.7648810 0.7784710 0.7811962 0.8125000 0.8313253    0
## 7  0.6904762 0.7619048 0.7738095 0.7788009 0.8102410 0.8333333    0
## 8  0.6904762 0.7627295 0.7844234 0.7847820 0.8184524 0.8433735    0
## 9  0.7261905 0.7747418 0.8083764 0.7955250 0.8258749 0.8333333    0
## 10 0.6904762 0.7837780 0.7904475 0.7895869 0.8214286 0.8433735    0
## 11 0.7023810 0.7791523 0.8024240 0.7943775 0.8184524 0.8433735    0
## 12 0.7380952 0.7910929 0.8144005 0.8051205 0.8288511 0.8452381    0
## 13 0.7142857 0.8005952 0.8192771 0.8075158 0.8403614 0.8452381    0
## 14 0.7380952 0.7941050 0.8203528 0.8098967 0.8403614 0.8452381    0
## 15 0.7142857 0.8000215 0.8203528 0.8075301 0.8378873 0.8554217    0
## 
## Kappa 
##         Min.   1st Qu.    Median      Mean   3rd Qu.      Max. NA's
## 5  0.3297872 0.4640436 0.5459706 0.5270773 0.6068751 0.6717371    0
## 6  0.3576471 0.4981484 0.5248805 0.5366310 0.6031287 0.6480921    0
## 7  0.3576471 0.4927448 0.5192771 0.5297159 0.5996437 0.6508314    0
## 8  0.3576471 0.4848320 0.5408159 0.5427127 0.6200253 0.6717371    0
## 9  0.4236277 0.5074421 0.5859472 0.5601687 0.6228626 0.6480921    0
## 10 0.3576471 0.5255698 0.5527057 0.5497490 0.6204819 0.6717371    0
## 11 0.3794326 0.5235007 0.5783191 0.5600467 0.6126720 0.6717371    0
## 12 0.4460432 0.5480930 0.5999072 0.5808134 0.6296780 0.6717371    0
## 13 0.4014252 0.5725752 0.6087279 0.5875305 0.6576219 0.6678832    0
## 14 0.4460432 0.5585005 0.6117973 0.5911995 0.6590982 0.6717371    0
## 15 0.4014252 0.5689401 0.6117973 0.5867010 0.6507194 0.6955990    0

La mayor precisiรณn media en este rango (0.8099) corresponde a maxnodes = 14, en la parte superior del intervalo analizado. Dado que el mejor valor se encuentra en el borde de la cuadrรญcula, conviene ampliar la bรบsqueda hacia arriba.

store_maxnode <- list()
tuneGrid <- expand.grid(.mtry = best_mtry)
for (maxnodes in c(20: 30)) {
    set.seed(1234)
    rf_maxnode <- train(survived~.,
        data = data_train,
        method = "rf",
        metric = "Accuracy",
        tuneGrid = tuneGrid,
        trControl = trControl,
        importance = TRUE,
        nodesize = 14,
        maxnodes = maxnodes,
        ntree = 300)
    key <- toString(maxnodes)
    store_maxnode[[key]] <- rf_maxnode
}
results_node <- resamples(store_maxnode)
summary(results_node)

Salida:

## 
## Call:
## summary.resamples(object = results_node)
## 
## Models: 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30 
## Number of resamples: 10 
## 
## Accuracy 
##         Min.   1st Qu.    Median      Mean   3rd Qu.      Max. NA's
## 20 0.7142857 0.7821644 0.8144005 0.8075301 0.8447719 0.8571429    0
## 21 0.7142857 0.8000215 0.8144005 0.8075014 0.8403614 0.8571429    0
## 22 0.7023810 0.7941050 0.8263769 0.8099254 0.8328313 0.8690476    0
## 23 0.7023810 0.7941050 0.8263769 0.8111302 0.8447719 0.8571429    0
## 24 0.7142857 0.7946429 0.8313253 0.8135112 0.8417599 0.8690476    0
## 25 0.7142857 0.7916667 0.8313253 0.8099398 0.8408635 0.8690476    0
## 26 0.7142857 0.7941050 0.8203528 0.8123207 0.8528758 0.8571429    0
## 27 0.7023810 0.8060456 0.8313253 0.8135112 0.8333333 0.8690476    0
## 28 0.7261905 0.7941050 0.8203528 0.8111015 0.8328313 0.8690476    0
## 29 0.7142857 0.7910929 0.8313253 0.8087063 0.8333333 0.8571429    0
## 30 0.6785714 0.7910929 0.8263769 0.8063253 0.8403614 0.8690476    0
## 
## Kappa 
##         Min.   1st Qu.    Median      Mean   3rd Qu.      Max. NA's
## 20 0.3956835 0.5316120 0.5961830 0.5854366 0.6661120 0.6955990    0
## 21 0.3956835 0.5699332 0.5960343 0.5853247 0.6590982 0.6919315    0
## 22 0.3735084 0.5560661 0.6221836 0.5914492 0.6422128 0.7189781    0
## 23 0.3735084 0.5594228 0.6228827 0.5939786 0.6657372 0.6955990    0
## 24 0.3956835 0.5600352 0.6337821 0.5992188 0.6604703 0.7189781    0
## 25 0.3956835 0.5530760 0.6354875 0.5912239 0.6554912 0.7189781    0
## 26 0.3956835 0.5589331 0.6136074 0.5969142 0.6822128 0.6955990    0
## 27 0.3735084 0.5852459 0.6368425 0.5998148 0.6426088 0.7189781    0
## 28 0.4290780 0.5589331 0.6154905 0.5946859 0.6356141 0.7189781    0
## 29 0.4070588 0.5534173 0.6337821 0.5901173 0.6423101 0.6919315    0
## 30 0.3297872 0.5534173 0.6202632 0.5843432 0.6590982 0.7189781    0

La mayor precisiรณn media, 0.8135, se obtiene con maxnodes = 24 (maxnodes = 27 empata en la media, pero tiene un tercer cuartil inferior). Por lo tanto, utilizarรก maxnodes = 24 para los pasos restantes.

Paso 4) Busca los mejores ntrees

Ahora que tiene el mejor valor de mtry y maxnode, puede ajustar la cantidad de รกrboles. El mรฉtodo es exactamente el mismo que el de maxnode.

store_maxtrees <- list()
for (ntree in c(250, 300, 350, 400, 450, 500, 550, 600, 800, 1000, 2000)) {
    set.seed(5678)
    rf_maxtrees <- train(survived~.,
        data = data_train,
        method = "rf",
        metric = "Accuracy",
        tuneGrid = tuneGrid,
        trControl = trControl,
        importance = TRUE,
        nodesize = 14,
        maxnodes = 24,
        ntree = ntree)
    key <- toString(ntree)
    store_maxtrees[[key]] <- rf_maxtrees
}
results_tree <- resamples(store_maxtrees)
summary(results_tree)

Salida:

## 
## Call:
## summary.resamples(object = results_tree)
## 
## Models: 250, 300, 350, 400, 450, 500, 550, 600, 800, 1000, 2000 
## Number of resamples: 10 
## 
## Accuracy 
##           Min.   1st Qu.    Median      Mean   3rd Qu.      Max. NA's
## 250  0.7380952 0.7976190 0.8083764 0.8087010 0.8292683 0.8674699    0
## 300  0.7500000 0.7886905 0.8024240 0.8027199 0.8203397 0.8452381    0
## 350  0.7500000 0.7886905 0.8024240 0.8027056 0.8277623 0.8452381    0
## 400  0.7500000 0.7886905 0.8083764 0.8051009 0.8292683 0.8452381    0
## 450  0.7500000 0.7886905 0.8024240 0.8039104 0.8292683 0.8452381    0
## 500  0.7619048 0.7886905 0.8024240 0.8062914 0.8292683 0.8571429    0
## 550  0.7619048 0.7886905 0.8083764 0.8099062 0.8323171 0.8571429    0
## 600  0.7619048 0.7886905 0.8083764 0.8099205 0.8323171 0.8674699    0
## 800  0.7619048 0.7976190 0.8083764 0.8110820 0.8292683 0.8674699    0
## 1000 0.7619048 0.7976190 0.8121510 0.8086723 0.8303571 0.8452381    0
## 2000 0.7619048 0.7886905 0.8121510 0.8086723 0.8333333 0.8452381    0
## 
## Kappa 
##           Min.   1st Qu.    Median      Mean   3rd Qu.      Max. NA's
## 250  0.4061697 0.5667400 0.5836013 0.5856103 0.6335363 0.7196807    0
## 300  0.4302326 0.5449376 0.5780349 0.5723307 0.6130767 0.6710843    0
## 350  0.4302326 0.5449376 0.5780349 0.5723185 0.6291592 0.6710843    0
## 400  0.4302326 0.5482030 0.5836013 0.5774782 0.6335363 0.6710843    0
## 450  0.4302326 0.5449376 0.5780349 0.5750587 0.6335363 0.6710843    0
## 500  0.4601542 0.5449376 0.5780349 0.5804340 0.6335363 0.6949153    0
## 550  0.4601542 0.5482030 0.5857118 0.5884507 0.6396872 0.6949153    0
## 600  0.4601542 0.5482030 0.5857118 0.5884374 0.6396872 0.7196807    0
## 800  0.4601542 0.5667400 0.5836013 0.5910088 0.6335363 0.7196807    0
## 1000 0.4601542 0.5667400 0.5961590 0.5857446 0.6343666 0.6678832    0
## 2000 0.4601542 0.5482030 0.5961590 0.5862151 0.6440678 0.6656337    0

Ya tienes tu modelo final. Puedes entrenar el bosque aleatorio con los siguientes parรกmetros:

  • ntree = 800: Se entrenarรกn 800 รกrboles.
  • mtry = 4: Se dibujan 4 caracterรญsticas candidatas en cada divisiรณn.
  • maxnodes = 24: Cada รกrbol estรก limitado a 24 nodos terminales (hojas).
fit_rf <- train(survived~.,
    data_train,
    method = "rf",
    metric = "Accuracy",
    tuneGrid = tuneGrid,
    trControl = trControl,
    importance = TRUE,
    nodesize = 14,
    ntree = 800,
    maxnodes = 24)

Paso 5) Evaluar el modelo

El cursor de la biblioteca tiene una funciรณn para hacer predicciones.

predict(model, newdata= df)
argument
- `model`: Define the model evaluated before. 
- `newdata`: Define the dataset to make prediction
prediction <-predict(fit_rf, data_test)

Puede utilizar la predicciรณn para calcular la matriz de confusiรณn y ver la puntuaciรณn de precisiรณn.

confusionMatrix(prediction, data_test$survived)

Salida:

## Confusion Matrix and Statistics
## 
##           Reference
## Prediction  No Yes
##        No  110  32
##        Yes  11  56
##                                          
##                Accuracy : 0.7943         
##                  95% CI : (0.733, 0.8469)
##     No Information Rate : 0.5789         
##     P-Value [Acc > NIR] : 3.959e-11      
##                                          
##                   Kappa : 0.5638         
##  Mcnemar's Test P-Value : 0.002289       
##                                          
##             Sensitivity : 0.9091         
##             Specificity : 0.6364         
##          Pos Pred Value : 0.7746         
##          Neg Pred Value : 0.8358         
##              Prevalence : 0.5789         
##          Detection Rate : 0.5263         
##    Detection Prevalence : 0.6794         
##       Balanced Accuracy : 0.7727         
##                                          
##        'Positive' Class : No             
## 

El modelo alcanza una precisiรณn de 0.7943, es decir, un 79.43 % en el conjunto de prueba no visto, superior a la configuraciรณn predeterminada. La sensibilidad es de 0.9091 y la especificidad de 0.6364, por lo que el modelo reconoce a los no supervivientes con mucha mรกs fiabilidad que a los supervivientes.

Paso 6) Visualice el resultado

Por รบltimo, puedes consultar la importancia de las caracterรญsticas con la funciรณn varImp(). Las caracterรญsticas mรกs importantes son el sexo y la edad. Esto no sorprende, ya que las caracterรญsticas importantes suelen aparecer mรกs cerca de la raรญz del รกrbol, mientras que las menos importantes generalmente aparecen mรกs cerca de las hojas.

varImp(fit_rf)

Salida:

## rf variable importance
## 
##              Importance
## sexmale         100.000
## age              28.014
## pclassMiddle     27.016
## fare             21.557
## pclassUpper      16.324
## sibsp            11.246
## parch             5.522
## embarkedC         4.908
## embarkedQ         1.420
## embarkedS         0.000		

Bosque aleatorio en R: Referencia rรกpida de funciones

La tabla que aparece a continuaciรณn enumera todas las funciones utilizadas en los seis pasos, el paquete que las proporciona y los parรกmetros que esperan.

Biblioteca Objetivo Funciรณn Parรกmetro
bosque aleatorio Crea un bosque aleatorio Bosque aleatorio() fรณrmula, ntree=n, mtry=FALSE, maxnodes = NULL
signo de intercalaciรณn Crear validaciรณn cruzada k-fold trenControl() mรฉtodo = โ€œcvโ€, nรบmero = n, bรบsqueda = โ€œcuadrรญculaโ€
signo de intercalaciรณn Entrena un bosque aleatorio entrenar() fรณrmula, df, mรฉtodo = โ€œrfโ€, mรฉtrica = โ€œPrecisiรณnโ€, trControl = trainControl(), tuneGrid = NULL
signo de intercalaciรณn Predecir fuera de la muestra predecir modelo, nuevos datos = df
signo de intercalaciรณn Matriz de confusiรณn y estadรญsticas matriz de confusiรณn() modelo, prueba y
signo de intercalaciรณn Importancia variable varImp() modelo

Apรฉndice: Modelos disponibles en caret

La funciรณn train() puede entrenar muchos mรกs modelos que solo bosques aleatorios. Ejecute el siguiente comando para imprimir todos los identificadores de modelos que admite caret y, a continuaciรณn, pase cualquiera de ellos al argumento del mรฉtodo.

names(getModelInfo())

Salida:

##   [1] "ada"                 "AdaBag"              "AdaBoost.M1"        ##   [4] "adaboost"            "amdai"               "ANFIS"              ##   [7] "avNNet"              "awnb"                "awtan"              ##  [10] "bag"                 "bagEarth"            "bagEarthGCV"        ##  [13] "bagFDA"              "bagFDAGCV"           "bam"                ##  [16] "bartMachine"         "bayesglm"            "binda"              ##  [19] "blackboost"          "blasso"              "blassoAveraged"     ##  [22] "bridge"              "brnn"                "BstLm"              ##  [25] "bstSm"               "bstTree"             "C5.0"               ##  [28] "C5.0Cost"            "C5.0Rules"           "C5.0Tree"           ##  [31] "cforest"             "chaid"               "CSimca"             ##  [34] "ctree"               "ctree2"              "cubist"             ##  [37] "dda"                 "deepboost"           "DENFIS"             ##  [40] "dnn"                 "dwdLinear"           "dwdPoly"            ##  [43] "dwdRadial"           "earth"               "elm"                ##  [46] "enet"                "evtree"              "extraTrees"         ##  [49] "fda"                 "FH.GBML"             "FIR.DM"             ##  [52] "foba"                "FRBCS.CHI"           "FRBCS.W"            ##  [55] "FS.HGD"              "gam"                 "gamboost"           ##  [58] "gamLoess"            "gamSpline"           "gaussprLinear"      ##  [61] "gaussprPoly"         "gaussprRadial"       "gbm_h3o"            ##  [64] "gbm"                 "gcvEarth"            "GFS.FR.MOGUL"       ##  [67] "GFS.GCCL"            "GFS.LT.RS"           "GFS.THRIFT"         ##  [70] "glm.nb"              "glm"                 "glmboost"           ##  [73] "glmnet_h3o"          "glmnet"              "glmStepAIC"         ##  [76] "gpls"                "hda"                 "hdda"               ##  [79] "hdrda"               "HYFIS"               "icr"                ##  [82] "J48"                 "JRip"                "kernelpls"          ##  [85] "kknn"                "knn"                 "krlsPoly"           ##  [88] "krlsRadial"          "lars"                "lars2"              ##  [91] "lasso"               "lda"                 "lda2"               ##  [94] "leapBackward"        "leapForward"         "leapSeq"            ##  [97] "Linda"               "lm"                  "lmStepAIC"          ## [100] "LMT"                 "loclda"              "logicBag"           ## [103] "LogitBoost"          "logreg"              "lssvmLinear"        ## [106] "lssvmPoly"           "lssvmRadial"         "lvq"                ## [109] "M5"                  "M5Rules"             "manb"               ## [112] "mda"                 "Mlda"                "mlp"                ## [115] "mlpKerasDecay"       "mlpKerasDecayCost"   "mlpKerasDropout"    ## [118] "mlpKerasDropoutCost" "mlpML"               "mlpSGD"             ## [121] "mlpWeightDecay"      "mlpWeightDecayML"    "monmlp"             ## [124] "msaenet"             "multinom"            "mxnet"              ## [127] "mxnetAdam"           "naive_bayes"         "nb"                 ## [130] "nbDiscrete"          "nbSearch"            "neuralnet"          ## [133] "nnet"                "nnls"                "nodeHarvest"        ## [136] "null"                "OneR"                "ordinalNet"         ## [139] "ORFlog"              "ORFpls"              "ORFridge"           ## [142] "ORFsvm"              "ownn"                "pam"                ## [145] "parRF"               "PART"                "partDSA"            ## [148] "pcaNNet"             "pcr"                 "pda"                ## [151] "pda2"                "penalized"           "PenalizedLDA"       ## [154] "plr"                 "pls"                 "plsRglm"            ## [157] "polr"                "ppr"                 "PRIM"               ## [160] "protoclass"          "pythonKnnReg"        "qda"                ## [163] "QdaCov"              "qrf"                 "qrnn"               ## [166] "randomGLM"           "ranger"              "rbf"                ## [169] "rbfDDA"              "Rborist"             "rda"                ## [172] "regLogistic"         "relaxo"              "rf"                 ## [175] "rFerns"              "RFlda"               "rfRules"            ## [178] "ridge"               "rlda"                "rlm"                ## [181] "rmda"                "rocc"                "rotationForest"     ## [184] "rotationForestCp"    "rpart"               "rpart1SE"           ## [187] "rpart2"              "rpartCost"           "rpartScore"         ## [190] "rqlasso"             "rqnc"                "RRF"                ## [193] "RRFglobal"           "rrlda"               "RSimca"             ## [196] "rvmLinear"           "rvmPoly"             "rvmRadial"          ## [199] "SBC"                 "sda"                 "sdwd"               ## [202] "simpls"              "SLAVE"               "slda"               ## [205] "smda"                "snn"                 "sparseLDA"          ## [208] "spikeslab"           "spls"                "stepLDA"            ## [211] "stepQDA"             "superpc"             "svmBoundrangeString"## [214] "svmExpoString"       "svmLinear"           "svmLinear2"         ## [217] "svmLinear3"          "svmLinearWeights"    "svmLinearWeights2"  ## [220] "svmPoly"             "svmRadial"           "svmRadialCost"      ## [223] "svmRadialSigma"      "svmRadialWeights"    "svmSpectrumString"  ## [226] "tan"                 "tanSearch"           "treebag"            ## [229] "vbmpRadial"          "vglmAdjCat"          "vglmContRatio"      ## [232] "vglmCumulative"      "widekernelpls"       "WM"                 ## [235] "wsrf"                "xgbLinear"           "xgbTree"            ## [238] "xyf"

Preguntas Frecuentes

Comience con 500, el valor predeterminado de randomForest(). La precisiรณn suele estabilizarse entre 300 y 1000 รกrboles. Aumentar ntree nunca perjudica la precisiรณn, solo el tiempo de ejecuciรณn, asรญ que aumรฉntelo hasta que la curva de error se estabilice.

El error fuera de la muestra evalรบa cada observaciรณn utilizando รบnicamente los รกrboles entrenados sin รฉl. Es una estimaciรณn rรกpida e imparcial y suele reemplazar la validaciรณn cruzada, aunque la validaciรณn cruzada k-fold sigue siendo preferible al comparar cuadrรญculas de ajuste en pliegues idรฉnticos.

Sรญ. Proporcione una respuesta numรฉrica y randomForest() promediarรก las predicciones del รกrbol en lugar de votar. En caret, mantenga method = โ€œrfโ€ y cambie el argumento de la mรฉtrica de Accuracy a RMSE.

Los bosques aleatorios siguen siendo un estรกndar de referencia para problemas de IA tabulares como la deserciรณn de clientes, el fraude y la evaluaciรณn de riesgos. Los equipos suelen evaluar el rendimiento de un bosque aleatorio antes de optar por el aumento de gradiente o las redes neuronales.

Sรญ. Los asistentes de IA pueden diseรฑar cuadrรญculas de ajuste, explicar la salida del remuestreo y sugerir rangos de mtry adecuados. Siempre vuelva a ejecutar el cรณdigo generado con una semilla fija para que la precisiรณn reportada sea reproducible.

Resumir este post con: