Después de preparar el conjunto de datos, el siguiente paso es dividir los datos en subconjuntos de formación y prueba para garantizar una evaluación efectiva del modelo de aprendizaje automático. Ahora, analicemos el proceso de partición en detalle.
Para crear y evaluar un modelo predictivo de forma eficaz, el conjunto de datos de entrada se divide aleatoriamente en dos subconjuntos separados: formación y pruebas. Si la columna "ID" no se especifica explícitamente, se supone que la primera columna del dataframe contiene el "ID". Para obtener más detalles, consulte [1].
No hay un porcentaje de partición universalmente óptimo para los subconjuntos mencionados anteriormente. La elección de la partición debe alinearse con los objetivos específicos del proyecto predictivo.
Los porcentajes de partición comunes incluyen:
- Formación: 80 % / Pruebas: 20 %
- Formación: 70% / Pruebas: 30%
- Formación: 60 % / Pruebas: 40 %
En este caso, el objetivo es maximizar los datos disponibles para la formación y, al mismo tiempo, garantizar que queden suficientes puntos de datos para una evaluación de modelo sólida. Por lo tanto, se ha seleccionado la siguiente partición de datos:
Formación: 85% / Pruebas: 15%Además, el subconjunto de formación se resume categorizando a los empleados en dos grupos:
- Empleados que permanecieron en la empresa en los últimos 12 meses.
- Empleados que se fueron dentro del mismo período.
La columna 'FLIGHT_RISK' sirve como indicador, marcando si un empleado ha dejado la empresa en los últimos 12 meses. La columna 'N' representa el recuento de empleados en cada categoría.
123456789101112
# Split the station classification dataframe into a training and test subset
df_train, df_test, df_val = train_test_val_split(data=hdf_employeechurn, id_column='EMPLOYEE_ID',
random_seed=1234,
partition_method='stratified', stratified_column='FLIGHT_RISK',
training_percentage=0.85,
testing_percentage=0.15,
validation_percentage=0.00)
#df_train.describe().collect()
df_train.agg([('count', 'EMPLOYEE_ID', 'N')], group_by='FLIGHT_RISK').collect()
| ITEM_NUMBER | FLIGHT_RISK | N |
|---|
| 0 | No | 14463 |
| 1 | Sí | 1785 |