Cómo hacer pruning de árboles en Machine Learning: paso a paso
Se explica cómo reducir el overfitting en modelos de árbol mediante pruning (poda) y por qué esto mejora la generalización. Vamos a ver técnicas de pre-poda y post-poda con un ejemplo práctico en Python usando DecisionTreeClassifier. La idea central es simple: árboles muy profundos se ajustan en exceso al ruido del entrenamiento; al podar ramas irrelevantes se reduce la varianza y se obtiene un modelo más sencillo e interpretable.
Prerequisitos
- Python 3.8+ instalado
- Bibliotecas: scikit-learn, pandas, numpy, matplotlib (ej.: pip install scikit-learn pandas numpy matplotlib)
- Conocimientos básicos de clasificación y train/test (train_test_split, métricas como accuracy)
- Máquina con recursos básicos: Iris ocupa 150 filas, CPUs modernas son suficientes
Paso 1: Por qué pruning en Machine Learning
Los árboles de decisión suelen ajustarse en exceso a los datos de entrenamiento (overfitting). El pruning elimina ramas poco relevantes para simplificar el árbol y mejorar el rendimiento en datos nuevos. Existen dos enfoques comunes: pre-poda (limitando max_depth, min_samples_leaf, etc.) y post-poda (cost-complexity pruning). La pre-poda actúa durante el entrenamiento (impide el crecimiento), mientras que la post-poda hace crecer el árbol completo y luego corta ramas con bajo beneficio. En la práctica, cuando la diferencia entre la accuracy en entrenamiento y en test es grande (por ejemplo >10 puntos porcentuales), es señal para aplicar pruning.
Paso 2: Preparar un conjunto de datos de ejemplo
Usa un dataset estándar para demostrar: el conjunto Iris o un dataset simulado. Aquí usamos Iris para simplificar la reproducción. Iris tiene 150 muestras (50 por clase), por lo que con test_size=0.3 obtienes 105 muestras de entrenamiento y 45 de test — suficiente para ver tendencias de overfitting sin complicar.
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
import pandas as pd
data = load_iris()
X = pd.DataFrame(data.data, columns=data.feature_names)
y = data.target
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3, random_state=42, stratify=y)
# X_train.shape => (105, 4), X_test.shape => (45, 4)
Paso 3: Entrenar un árbol sin pruning (baseline)
Entrena un DecisionTreeClassifier sin restricciones para observar el overfitting. Evalúa entrenamiento y test. Típicamente un árbol sin restricciones alcanza 100% en entrenamiento (ajuste perfecto) y puede tener 95–99% en test en Iris; la diferencia revela el grado de overfitting.
from sklearn.tree import DecisionTreeClassifier
from sklearn.metrics import accuracy_score
clf = DecisionTreeClassifier(random_state=42)
clf.fit(X_train, y_train)
train_acc = accuracy_score(y_train, clf.predict(X_train))
test_acc = accuracy_score(y_test, clf.predict(X_test))
print(f"Baseline train acc: {train_acc:.3f}, test acc: {test_acc:.3f}")
Paso 4: Pre-poda — usar parámetros para controlar la complejidad
Prueba a limitar max_depth, min_samples_leaf y min_samples_split. Estos parámetros impiden que el árbol crezca arbitrariamente. Por ejemplo, con max_depth=3 y min_samples_leaf=5 normalmente reduces la profundidad y el número de nodos drásticamente; en Iris esto tiende a mantener la test accuracy cercana al baseline (p.ej. 0.93–0.97) pero con una brecha entrenamiento/test mucho menor.
clf_pre = DecisionTreeClassifier(random_state=42, max_depth=3, min_samples_leaf=5)
clf_pre.fit(X_train, y_train)
print("Pre-pruning train:", accuracy_score(y_train, clf_pre.predict(X_train)))
print("Pre-pruning test:", accuracy_score(y_test, clf_pre.predict(X_test)))
Paso 5: Post-poda — cost-complexity pruning (ccp_alpha)
scikit-learn implementa cost-complexity pruning. Primero se obtiene una secuencia de ccp_alpha posibles y después se elige el mejor por rendimiento en validación. El algoritmo devuelve típicamente 5–30 valores de alpha, desde muy pequeño (casi sin poda) hasta grande (árbol reducido al nodo raíz). Evalúa cada alpha y elige el que maximice la métrica en el conjunto de validación o test. Atención: evita elegir el alpha basándote solo en el test final; usa validación cruzada cuando sea posible.
# obtener camino de poda
path = DecisionTreeClassifier(random_state=42).cost_complexity_pruning_path(X_train, y_train)
ccp_alphas = path.ccp_alphas
clfs = []
for ccp in ccp_alphas:
clf_tmp = DecisionTreeClassifier(random_state=42, ccp_alpha=ccp)
clf_tmp.fit(X_train, y_train)
clfs.append((ccp, clf_tmp))
# evaluar cada alpha
import numpy as np
train_scores = [accuracy_score(y_train, m.predict(X_train)) for _, m in clfs]
test_scores = [accuracy_score(y_test, m.predict(X_test)) for _, m in clfs]
best_idx = int(np.argmax(test_scores))
best_alpha, best_model = clfs[best_idx]
print(f"Mejor ccp_alpha: {best_alpha}, test acc: {test_scores[best_idx]:.3f}")
Paso 6: Elegir el modelo final y validar con cross-validation
Confirma la robustez del alpha elegido usando cross-validation para evitar seleccionar por casualidad en el conjunto de test. Por ejemplo, haz 5-fold CV y compara la media y la desviación estándar de las accuracies. Valores consistentes (p.ej. mean=0.95, std<0.03) indican estabilidad.
from sklearn.model_selection import cross_val_score
final_clf = DecisionTreeClassifier(random_state=42, ccp_alpha=best_alpha)
scores = cross_val_score(final_clf, X, y, cv=5)
print("CV accuracy:", scores.mean(), scores)
Verificar el resultado
Compara las métricas: si el modelo podado tiene rendimiento similar en test pero con menor diferencia entrenamiento-test, el pruning funcionó. Verifica también la simplicidad del árbol (profundidad y número de nodos) para confirmar la reducción de la complejidad. Ej.: Baseline depth: 5, nodes: 35 vs Pruned depth: 3, nodes: 9 — esto se traduce en un modelo más interpretable y menos sensible al ruido.
print("Baseline depth:", clf.get_depth(), "nodes:", clf.tree_.node_count)
print("Pruned depth:", best_model.get_depth(), "nodes:", best_model.tree_.node_count)
Conclusión
El pruning (pre-poda y post-poda) ayuda a reducir el overfitting y a obtener modelos de árbol más interpretable y robustos. Próximos pasos: probar con conjuntos más complejos (p.ej. cientos de miles de registros), usar validación en k-fold para seleccionar parámetros y comparar con ensembles (RandomForest, GradientBoosting) que reducen la varianza sin tanto tuning manual. Consejo práctico: cuando veas una gran diferencia entre entrenamiento y test, prueba primero max_depth y min_samples_leaf; si necesitas más precisión usa cost-complexity pruning (ccp_alpha) con validación cruzada.