Demo_churn / scripts /generar_graficas.py
Larxmind's picture
Despliegue inicial limpio sin binarios
1ec1a50
Raw
History Blame Contribute Delete
4.17 kB
import importlib.util
import os
import sys
from pathlib import Path
PROJECT_ROOT = Path(__file__).resolve().parent.parent
VENV_PYTHON = PROJECT_ROOT / 'env_academia' / 'bin' / 'python'
if (
(importlib.util.find_spec('seaborn') is None or importlib.util.find_spec('shap') is None)
and VENV_PYTHON.exists()
and Path(sys.executable).resolve() != VENV_PYTHON.resolve()
):
os.execv(str(VENV_PYTHON), [str(VENV_PYTHON), str(Path(__file__).resolve()), *sys.argv[1:]])
import pandas as pd
import joblib
import matplotlib.pyplot as plt
import seaborn as sns
from sklearn.metrics import confusion_matrix, roc_curve, auc
from sklearn.model_selection import train_test_split
import shap
# 1. Configuración de rutas (Ajustadas a tu árbol actual)
MODEL_PATH = 'modelos/modelo_rf_optimizado.pkl'
DATA_PATH = 'datos_procesados/dataset_preparado_ml.csv'
OUTPUT_DIR = 'outputs/charts'
os.makedirs(OUTPUT_DIR, exist_ok=True)
print("📦 Cargando modelo y datos preparados...")
model = joblib.load(MODEL_PATH)
df = pd.read_csv(DATA_PATH)
# 2. Separar características (X) y variable objetivo (y)
target_column = 'churn' if 'churn' in df.columns else 'abandono'
if target_column not in df.columns:
raise KeyError("No se encontró la columna objetivo 'churn' ni 'abandono' en el dataset preparado")
X = df.drop(columns=[target_column, 'estudiante_id'], errors='ignore')
y = df[target_column]
# 3. REPLICAR LA DIVISIÓN DEL NOTEBOOK
# ¡Crucial! Usa el mismo test_size y random_state que usaste en app.ipynb
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)
print(f"📊 Evaluando sobre {len(X_test)} registros de validación...")
# Predicciones sobre el conjunto de test
y_pred = model.predict(X_test)
y_prob = model.predict_proba(X_test)[:, 1]
print("🎨 Generando representaciones gráficas...")
# -----------------------------------------------------------------
# GRÁFICO 1: MATRIZ DE CONFUSIÓN
# -----------------------------------------------------------------
plt.figure(figsize=(6, 5))
cm = confusion_matrix(y_test, y_pred)
sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',
xticklabels=['No Abandona', 'Abandona'],
yticklabels=['No Abandona', 'Abandona'])
plt.title('Matriz de Confusión - Predicción de Abandono')
plt.ylabel('Clase Real')
plt.xlabel('Clase Predicha')
plt.tight_layout()
plt.savefig(os.path.join(OUTPUT_DIR, 'matriz_confusion.png'), dpi=300)
plt.close()
print("✅ Matriz de confusión guardada.")
# -----------------------------------------------------------------
# GRÁFICO 2: CURVA ROC
# -----------------------------------------------------------------
fpr, tpr, _ = roc_curve(y_test, y_prob)
roc_auc = auc(fpr, tpr)
plt.figure(figsize=(6, 5))
plt.plot(fpr, tpr, color='darkorange', lw=2, label=f'Curva ROC (AUC = {roc_auc:.4f})')
plt.plot([0, 1], [0, 1], color='navy', lw=2, linestyle='--')
plt.xlim([0.0, 1.0])
plt.ylim([0.0, 1.05])
plt.xlabel('Tasa de Falsos Positivos (FPR)')
plt.ylabel('Tasa de Verdaderos Positivos (TPR)')
plt.title('Curva ROC (Receiver Operating Characteristic)')
plt.legend(loc="lower right")
plt.grid(True, linestyle='--', alpha=0.6)
plt.tight_layout()
plt.savefig(os.path.join(OUTPUT_DIR, 'curva_roc.png'), dpi=300)
plt.close()
print("✅ Curva ROC guardada.")
# -----------------------------------------------------------------
# GRÁFICO 3: IMPORTANCIA DE VARIABLES (SHAP VALUES)
# -----------------------------------------------------------------
# SHAP puede ser computacionalmente pesado, usamos el X_test
explainer = shap.TreeExplainer(model)
shap_values = explainer.shap_values(X_test)
plt.figure(figsize=(8, 5))
if isinstance(shap_values, list):
shap_vals_to_plot = shap_values[1]
else:
shap_vals_to_plot = shap_values
shap.summary_plot(shap_vals_to_plot, X_test, plot_type="bar", show=False)
plt.title('Importancia Global de las Variables (SHAP)', fontsize=14, pad=15)
plt.tight_layout()
plt.savefig(os.path.join(OUTPUT_DIR, 'importancia_shap.png'), dpi=300)
plt.close()
print("✅ Gráfico de importancia SHAP guardado.")
print(f"\n🎉 ¡Proceso completado! Archivos listos en: '{OUTPUT_DIR}/'")