Demo_churn / scripts /explicabilidad_shap.py
Larxmind's picture
Despliegue inicial limpio sin binarios
1ec1a50
Raw
History Blame Contribute Delete
4.33 kB
import pandas as pd
import shap
import joblib
import numpy as np
def obtener_explicacion_alumno(id_alumno, ruta_modelo, ruta_datos_ml, ruta_datos_maestro):
"""
Analiza a un estudiante usando SHAP y extrae su perfil sociodemográfico.
Devuelve un JSON estructurado para alimentar al Agente LLM del tutor.
"""
try:
modelo = joblib.load(ruta_modelo)
df_ml = pd.read_csv(ruta_datos_ml)
df_maestro = pd.read_csv(ruta_datos_maestro)
except FileNotFoundError as e:
return {"error": f"Archivo no encontrado: {e}"}
# Buscar la fila del alumno
indice_alumno = df_maestro.index[df_maestro['id_alumno'] == id_alumno].tolist()
if not indice_alumno:
return {"error": "Alumno no encontrado en la base de datos."}
idx = indice_alumno[0]
# 1. EXTRACCIÓN Y MAPEO DEL CONTEXTO SOCIODEMOGRÁFICO (Lenguaje Natural para el LLM)
gen_raw = str(df_maestro.loc[idx, 'genero'])
disc_raw = str(df_maestro.loc[idx, 'discapacidad'])
# Mapeos simples para claridad del LLM
texto_genero = "Mujer" if gen_raw == "F" else "Hombre" if gen_raw == "M" else "No especificado"
texto_discapacidad = "Sí, requiere adaptaciones o atención especial" if disc_raw == "Y" else "No consta"
contexto_alumno = {
"perfil_demografico": texto_genero,
"grupo_de_edad": str(df_maestro.loc[idx, 'rango_edad']) + " años",
"nivel_estudios_previo": str(df_maestro.loc[idx, 'nivel_educativo']),
"indice_nivel_economico": f"Banda de ingresos {df_maestro.loc[idx, 'nivel_economico']}",
"veces_matriculado_anteriormente": int(df_maestro.loc[idx, 'intentos_previos']),
"tiene_discapacidad_reconocida": texto_discapacidad
}
# 2. PREDICCIÓN Y EXPLICABILIDAD MATEMÁTICA (XAI)
X = df_ml.drop(columns=['churn', 'code_presentation'], errors='ignore')
perfil_alumno = X.iloc[[idx]]
prob_abandono = modelo.predict_proba(perfil_alumno)[0][1]
if prob_abandono < 0.50:
return {
"alerta": False,
"id_alumno": id_alumno,
"mensaje": "Sin riesgo crítico. No requiere intervención del tutor."
}
explainer = shap.TreeExplainer(modelo)
shap_values = explainer.shap_values(perfil_alumno)
# --- CORRECCIÓN DEL ERROR DE DIMENSIONES NUMPY ---
# Interceptamos la matriz de SHAP y forzamos la extracción de escalares 1D
if isinstance(shap_values, list):
valores_shap_clase1 = shap_values[1][0]
else:
shap_array = np.array(shap_values)
if len(shap_array.shape) == 3: # (muestras, features, clases)
valores_shap_clase1 = shap_array[0, :, 1]
elif len(shap_array.shape) == 2: # (muestras, features)
valores_shap_clase1 = shap_array[0, :]
else:
valores_shap_clase1 = shap_array.flatten()
# Aplanamos y aseguramos que es un array 1D
valores_shap_clase1 = np.array(valores_shap_clase1).flatten()
impacto_variables = {}
for nombre_columna, valor_shap in zip(X.columns, valores_shap_clase1):
if valor_shap > 0:
impacto_variables[nombre_columna] = valor_shap
impacto_ordenado = sorted(impacto_variables.items(), key=lambda item: item[1], reverse=True)
top_causas = []
for var, impacto in impacto_ordenado[:3]:
top_causas.append({
"variable_tecnica": var,
"valor_actual": float(perfil_alumno[var].iloc[0]),
"peso_en_riesgo": round(impacto, 3)
})
# 3. EMPAQUETADO FINAL PARA EL LLM
return {
"alerta": True,
"id_alumno": id_alumno,
"diagnostico_matematico": {
"riesgo_abandono_porcentaje": round(prob_abandono * 100, 2),
"causas_shap": top_causas
},
"contexto_alumno": contexto_alumno
}
if __name__ == "__main__":
print("Iniciando generación de payload para el tutor...")
resultado = obtener_explicacion_alumno(
id_alumno=30268,
ruta_modelo='datos_procesados/modelo_rf_optimizado.pkl',
ruta_datos_ml='datos_procesados/dataset_preparado_ml.csv',
ruta_datos_maestro='datos_procesados/dataset_maestro_churn.csv'
)
import json
print(json.dumps(resultado, indent=4, ensure_ascii=False))