""" app.py — Demo Gradio: Detección temprana de riesgo de diabetes Space: Jesusrodriguezf90/diabetes-risk-demo Carga el pipeline LightGBM entrenado sobre BRFSS 2015 y ofrece: - Formulario de 22 variables clínicas con threshold dinámico - Predicción binaria con probabilidad - Gráfico SHAP waterfall individual (explicación de la predicción) Autor: Jesús Rodríguez Fernández """ # --- Librerías estándar --- import sys import types import warnings warnings.filterwarnings("ignore") # --- Librerías terceros --- import gradio as gr import joblib import matplotlib import matplotlib.pyplot as plt import numpy as np import pandas as pd import shap from huggingface_hub import hf_hub_download # --- Usar backend no interactivo para matplotlib (necesario en Spaces) --- matplotlib.use("Agg") # --------------------------------------------------------------------------- # MÓDULO SRC — necesario para deserializar el .pkl # El pipeline fue serializado con src.preprocessing.preprocessing_pipeline, # por lo que hay que registrar ese módulo antes de llamar a joblib.load() # --------------------------------------------------------------------------- # Constantes del preprocesador (copiadas de preprocessing_pipeline.py) _CATEGORICAL_NOMINAL = ["BPHIGH4", "_RACE"] _BINARY_VARS = [ "BPMEDS", "BLOODCHO", "HAVARTH3", "QLACTLM2", "USEEQUIP", "BLIND", "DECIDE", "DIFFWALK", "DIFFALON", "DIFFDRES", "SMOKE100", "ADDEPEV2", "SEX", ] _CATEGORICAL_ORDINAL = ["GENHLTH", "_PACAT1", "_AGEG5YR", "_BMI5CAT"] _NUMERIC_VARS = ["EXEROFT1", "_FRUTSUM", "_VEGESUM"] def _boosting_deterministic_preproc(X_input: pd.DataFrame) -> pd.DataFrame: """Preprocesamiento determinista: reemplaza -1 por NaN, normaliza binarias, fuerza tipo categórico en nominales.""" X = X_input.copy() X = X.replace(-1, np.nan) for col in _BINARY_VARS: if col in X.columns: X[col] = (X[col] == 1).astype(int) for col in _CATEGORICAL_NOMINAL: if col in X.columns: X[col] = X[col].astype("category") return X def _cap_outliers_numeric(X_input, numeric_vars): """Recorta valores extremos al percentil 1-99.""" X = X_input.copy() for col in numeric_vars: if col in X.columns: low = np.nanpercentile(X[col], 1) high = np.nanpercentile(X[col], 99) X[col] = X[col].clip(low, high) return X def _register_src_module(): """Registra src.preprocessing.preprocessing_pipeline en sys.modules para que joblib.load() pueda deserializar el pipeline.""" src_mod = types.ModuleType("src") preprocessing_mod = types.ModuleType("src.preprocessing") pipeline_mod = types.ModuleType("src.preprocessing.preprocessing_pipeline") pipeline_mod.boosting_deterministic_preproc = _boosting_deterministic_preproc pipeline_mod.cap_outliers_numeric = _cap_outliers_numeric pipeline_mod.CATEGORICAL_NOMINAL = _CATEGORICAL_NOMINAL pipeline_mod.BINARY_VARS = _BINARY_VARS pipeline_mod.CATEGORICAL_ORDINAL = _CATEGORICAL_ORDINAL pipeline_mod.NUMERIC_VARS = _NUMERIC_VARS src_mod.preprocessing = preprocessing_mod preprocessing_mod.preprocessing_pipeline = pipeline_mod sys.modules.setdefault("src", src_mod) sys.modules.setdefault("src.preprocessing", preprocessing_mod) sys.modules.setdefault( "src.preprocessing.preprocessing_pipeline", pipeline_mod ) # --------------------------------------------------------------------------- # CARGA DEL MODELO # --------------------------------------------------------------------------- _register_src_module() _MODEL_REPO = "Jesusrodriguezf90/lgbm-diabetes-early-detection" _MODEL_FILE = "lgbm_diabetes_pipeline.pkl" _model_path = hf_hub_download(repo_id=_MODEL_REPO, filename=_MODEL_FILE) PIPELINE = joblib.load(_model_path) # --------------------------------------------------------------------------- # CONSTANTES DE LA DEMO # --------------------------------------------------------------------------- # Orden exacto de columnas de salida del ColumnTransformer (verificado del pkl): # nom → bin → ord → num FEATURE_NAMES = [ "BPHIGH4", "_RACE", "BPMEDS", "BLOODCHO", "HAVARTH3", "QLACTLM2", "USEEQUIP", "BLIND", "DECIDE", "DIFFWALK", "DIFFALON", "DIFFDRES", "SMOKE100", "ADDEPEV2", "SEX", "GENHLTH", "_PACAT1", "_AGEG5YR", "_BMI5CAT", "EXEROFT1", "_FRUTSUM", "_VEGESUM", ] # Etiquetas legibles para el gráfico SHAP FEATURE_LABELS = { "BPHIGH4": "Presión arterial alta", "_RACE": "Raza / etnia", "BPMEDS": "Medicación tensión arterial", "BLOODCHO": "Revisión colesterol", "HAVARTH3": "Artritis / reumatismo", "QLACTLM2": "Limitación actividades", "USEEQUIP": "Equipo de asistencia", "BLIND": "Ceguera / dif. visión", "DECIDE": "Dif. concentración / memoria", "DIFFWALK": "Dif. caminar / escaleras", "DIFFALON": "Dif. recados solo", "DIFFDRES": "Dif. vestirse / bañarse", "SMOKE100": "Fumador (>100 cigarrillos)", "ADDEPEV2": "Depresión diagnosticada", "SEX": "Sexo", "GENHLTH": "Salud general", "_PACAT1": "Actividad física (categoría)", "_AGEG5YR": "Grupo de edad", "_BMI5CAT": "Categoría IMC", "EXEROFT1": "Frecuencia ejercicio (sem.)", "_FRUTSUM": "Consumo de frutas (porc./día)", "_VEGESUM": "Consumo de verduras (porc./día)", } THRESHOLD_DEFAULT = 0.4733 # --------------------------------------------------------------------------- # LÓGICA DE INFERENCIA Y SHAP # --------------------------------------------------------------------------- def _build_dataframe( bphigh4, race, bpmeds, bloodcho, havarth3, qlactlm2, useequip, blind, decide, diffwalk, diffalon, diffdres, smoke100, addepev2, sex, genhlth, pacat1, ageg5yr, bmi5cat, exeroft1, frutsum, vegesum, ) -> pd.DataFrame: """Construye el DataFrame de entrada con los nombres de columna correctos.""" return pd.DataFrame([{ "GENHLTH": int(genhlth), "BPHIGH4": int(bphigh4), "BPMEDS": int(bpmeds), "BLOODCHO": int(bloodcho), "HAVARTH3": int(havarth3), "ADDEPEV2": int(addepev2), "SEX": int(sex), "QLACTLM2": int(qlactlm2), "USEEQUIP": int(useequip), "BLIND": int(blind), "DECIDE": int(decide), "DIFFWALK": int(diffwalk), "DIFFDRES": int(diffdres), "DIFFALON": int(diffalon), "SMOKE100": int(smoke100), "EXEROFT1": float(exeroft1), "_RACE": int(race), "_AGEG5YR": int(ageg5yr), "_BMI5CAT": int(bmi5cat), "_FRUTSUM": float(frutsum), "_VEGESUM": float(vegesum), "_PACAT1": int(pacat1), }]) def predict_and_explain( bphigh4, race, bpmeds, bloodcho, havarth3, qlactlm2, useequip, blind, decide, diffwalk, diffalon, diffdres, smoke100, addepev2, sex, genhlth, pacat1, ageg5yr, bmi5cat, exeroft1, frutsum, vegesum, threshold, ): """ Ejecuta inferencia y genera el gráfico SHAP waterfall individual. Returns: resultado (str): Decisión clínica + probabilidad. fig (matplotlib.Figure): Gráfico SHAP waterfall. """ df = _build_dataframe( bphigh4, race, bpmeds, bloodcho, havarth3, qlactlm2, useequip, blind, decide, diffwalk, diffalon, diffdres, smoke100, addepev2, sex, genhlth, pacat1, ageg5yr, bmi5cat, exeroft1, frutsum, vegesum, ) # --- Predicción --- proba = float(PIPELINE.predict_proba(df)[0][1]) decision = ( "⚠️ Realizar prueba HbA1c" if proba >= threshold else "✅ No se recomienda prueba HbA1c" ) resultado = ( f"**{decision}**\n\n" f"Probabilidad estimada de riesgo: **{proba * 100:.1f}%**\n\n" f"Umbral utilizado: {threshold:.4f}" ) # --- SHAP waterfall individual --- # Aplicar solo deterministic + preprocessor (todos los steps menos el modelo) preproc = PIPELINE[:-1] X_transformed = preproc.transform(df) lgb_model = PIPELINE.named_steps["model"] explainer = shap.TreeExplainer(lgb_model) shap_values = explainer.shap_values(X_transformed) # shap_values puede ser lista [clase0, clase1] o array 2D según versión # Tomamos clase positiva (índice 1) if isinstance(shap_values, list): sv_positive = shap_values[1][0] base_val = explainer.expected_value[1] else: sv_positive = shap_values[0] base_val = explainer.expected_value # Etiquetas legibles para el eje Y del waterfall feature_labels = [FEATURE_LABELS.get(n, n) for n in FEATURE_NAMES] fig, ax = plt.subplots(figsize=(9, 7)) plt.sca(ax) shap.waterfall_plot( shap.Explanation( values=sv_positive, base_values=base_val, data=X_transformed[0], feature_names=feature_labels, ), max_display=15, show=False, ) fig = plt.gcf() fig.tight_layout() return resultado, fig # --------------------------------------------------------------------------- # INTERFAZ GRADIO # --------------------------------------------------------------------------- # NOTA: en gr.Dropdown, choices es una lista de tuplas (nombre_visible, valor) # El usuario ve el nombre_visible; la función recibe el valor numérico. # --------------------------------------------------------------------------- with gr.Blocks(title="Detección de Riesgo de Diabetes") as demo: gr.Markdown( """ # 🩺 Detección temprana de riesgo de diabetes Modelo LightGBM entrenado sobre ~258.000 observaciones clínicas del BRFSS 2015 (CDC). Rellena el formulario con los datos del paciente y pulsa **Predecir**. > ⚠️ **Aviso**: Esta herramienta es un apoyo al cribado, no un diagnóstico médico. > La decisión clínica final corresponde siempre al profesional sanitario. """ ) with gr.Row(): # --- Columna izquierda: formulario --- with gr.Column(scale=1): gr.Markdown("### Datos del paciente") genhlth = gr.Dropdown( label="Estado general de salud", choices=[ ("Excelente", 1), ("Muy bueno", 2), ("Bueno", 3), ("Regular", 4), ("Malo", 5), ], value=3, ) ageg5yr = gr.Dropdown( label="Grupo de edad", choices=[ ("18–24 años", 1), ("25–29 años", 2), ("30–34 años", 3), ("35–39 años", 4), ("40–44 años", 5), ("45–49 años", 6), ("50–54 años", 7), ("55–59 años", 8), ("60–64 años", 9), ("65–69 años", 10), ("70–74 años", 11), ("75–79 años", 12), ("80 años o más",13), ], value=7, ) sex = gr.Dropdown( label="Sexo", choices=[ ("Masculino", 1), ("Femenino", 2), ], value=1, ) race = gr.Dropdown( label="Raza / etnia", choices=[ ("Blanco/a", 1), ("Negro/a", 2), ("Indígena americano/a", 3), ("Asiático/a", 4), ("Nativo/a de Hawái / Pacífico", 5), ("Otra raza", 6), ("Multirracial", 7), ("Hispano/a", 8), ], value=1, ) bmi5cat = gr.Dropdown( label="Categoría IMC", choices=[ ("Bajo peso", 1), ("Peso normal", 2), ("Sobrepeso", 3), ("Obesidad", 4), ], value=2, ) gr.Markdown("### Condiciones médicas") bphigh4 = gr.Dropdown( label="¿Le han diagnosticado presión arterial alta?", choices=[ ("Sí", 1), ("Sí, solo durante el embarazo",2), ("No", 3), ("Borderline / prehipertensión",4), ], value=3, ) bpmeds = gr.Dropdown( label="¿Toma medicación para la tensión arterial?", choices=[ ("Sí", 1), ("No", 2), ], value=2, ) bloodcho = gr.Dropdown( label="¿Se ha revisado el colesterol alguna vez?", choices=[ ("Sí", 1), ("No", 2), ], value=1, ) havarth3 = gr.Dropdown( label="¿Le han diagnosticado artritis / reumatismo / gota?", choices=[ ("Sí", 1), ("No", 2), ], value=2, ) addepev2 = gr.Dropdown( label="¿Le han diagnosticado depresión?", choices=[ ("Sí", 1), ("No", 2), ], value=2, ) gr.Markdown("### Limitaciones funcionales") qlactlm2 = gr.Dropdown( label="¿Tiene limitaciones en actividades por problemas de salud?", choices=[ ("Sí", 1), ("No", 2), ], value=2, ) useequip = gr.Dropdown( label="¿Usa equipo especial (bastón, silla de ruedas...)?", choices=[ ("Sí", 1), ("No", 2), ], value=2, ) blind = gr.Dropdown( label="¿Tiene ceguera o dificultad grave de visión?", choices=[ ("Sí", 1), ("No", 2), ], value=2, ) decide = gr.Dropdown( label="¿Tiene dificultad para concentrarse o tomar decisiones?", choices=[ ("Sí", 1), ("No", 2), ], value=2, ) diffwalk = gr.Dropdown( label="¿Tiene dificultad para caminar o subir escaleras?", choices=[ ("Sí", 1), ("No", 2), ], value=2, ) diffalon = gr.Dropdown( label="¿Tiene dificultad para hacer recados solo?", choices=[ ("Sí", 1), ("No", 2), ], value=2, ) diffdres = gr.Dropdown( label="¿Tiene dificultad para vestirse o bañarse?", choices=[ ("Sí", 1), ("No", 2), ], value=2, ) gr.Markdown("### Hábitos de vida") smoke100 = gr.Dropdown( label="¿Ha fumado más de 100 cigarrillos en su vida?", choices=[ ("Sí", 1), ("No", 2), ], value=2, ) pacat1 = gr.Dropdown( label="Nivel de actividad física", choices=[ ("Muy activo/a", 1), ("Activo/a", 2), ("Insuficientemente activo/a", 3), ("Inactivo/a", 4), ], value=2, ) exeroft1 = gr.Dropdown( label="Frecuencia de ejercicio semanal", choices=[ ("Menos de 1 vez por semana", 0.47), ("1 vez por semana", 1.0), ("2 veces por semana", 2.0), ("3 veces por semana", 3.0), ("4 veces por semana", 4.0), ("5 veces por semana", 5.0), ("6 veces por semana", 6.0), ("Todos los días", 7.0), ], value=3.0, ) frutsum = gr.Dropdown( label="Consumo diario de frutas", choices=[ ("Nunca o casi nunca", 0.0), ("Menos de 1 porción al día", 0.43), ("1 porción al día", 1.0), ("1–2 porciones al día", 1.43), ("2 porciones al día", 2.0), ("3 porciones al día", 3.0), ("4 o más porciones al día", 4.0), ], value=1.0, ) vegesum = gr.Dropdown( label="Consumo diario de verduras", choices=[ ("Nunca o casi nunca", 0.0), ("Menos de 1 porción al día", 0.57), ("1 porción al día", 1.0), ("1–2 porciones al día", 1.57), ("2 porciones al día", 2.0), ("3 porciones al día", 3.0), ("4 o más porciones al día", 4.0), ], value=2.0, ) gr.Markdown("### Umbral de decisión") threshold = gr.Slider( minimum=0.10, maximum=0.90, step=0.05, value=THRESHOLD_DEFAULT, label=f"Threshold (por defecto: {THRESHOLD_DEFAULT})", info=( "Valores bajos → más sensible (detecta más casos, más falsos positivos). " "Valores altos → más específico (menos falsos positivos, puede perder casos)." ), ) btn = gr.Button("Predecir", variant="primary") # --- Columna derecha: resultado + SHAP --- with gr.Column(scale=1): gr.Markdown("### Resultado") resultado_md = gr.Markdown( value="Rellena el formulario y pulsa **Predecir**." ) gr.Markdown("### Explicación SHAP — contribución de cada variable") gr.Markdown( "_Las barras rojas aumentan el riesgo predicho; las azules lo reducen. " "Los valores están en escala log-odds (no son probabilidades directas). " "E[f(X)] es la predicción media del modelo; f(x) es la predicción para " "este paciente concreto_" ) shap_plot = gr.Plot(label="SHAP waterfall") # --- Evento --- btn.click( fn=predict_and_explain, inputs=[ bphigh4, race, bpmeds, bloodcho, havarth3, qlactlm2, useequip, blind, decide, diffwalk, diffalon, diffdres, smoke100, addepev2, sex, genhlth, pacat1, ageg5yr, bmi5cat, exeroft1, frutsum, vegesum, threshold, ], outputs=[resultado_md, shap_plot], ) gr.Markdown( """ --- **Modelo**: LightGBM · **Dataset**: BRFSS 2015 (CDC) · ~258.000 observaciones · 22 variables **Métricas** (conjunto test, threshold=0.4733): ROC-AUC=0.839 · Recall=0.813 · Precision=0.256 **Repositorio**: [GitHub — TFM](https://github.com/Jesusrodriguezf90/TFM) · **Modelo en HF Hub**: [lgbm-diabetes-early-detection](https://huggingface.co/Jesusrodriguezf90/lgbm-diabetes-early-detection) *Esta herramienta es un apoyo al cribado. No reemplaza el diagnóstico médico.* """ ) if __name__ == "__main__": demo.launch()