Explicabilidad y Grad-CAM (Abriendo la Caja Negra)#
Objetivos y Restricciones#
Objetivo principal: Implementar métodos de explicabilidad para verificar que el modelo basa sus predicciones en características morfológicas reales de la planta y no en artefactos del fondo.
Restricciones de Ingeniería:
Cero Dependencias Frágiles: Implementamos la matemática de Grad-CAM de forma nativa usando
tf.GradientTapepara garantizar compatibilidad total con Keras 3.Latencia de Explicación: El método elegido para producción debe generar el mapa de calor en
< 1000 ms.
Prerrequisitos#
Haber completado el Capítulo 4 (Deep Learning), especialmente Transfer Learning y Fine-Tuning
Contexto del Problema#
Entrenamos con éxito nuestro modelo EfficientNetB0 para detectar enfermedades en hojas de té en el notebook Diseño de la Cabeza de Clasificación. El modelo reporta aprox. un 90% de Accuracy en el Test Set.
Sin embargo, la empresa agrícola que nos contrata se niega a desplegar el modelo en sus drones. El ingeniero agrónomo en jefe plantea una duda razonable: “En las fotos de entrenamiento, las hojas sanas estaban sobre un fondo azul y las enfermas sobre un fondo blanco. ¿Cómo sé que la IA está mirando las manchas de la hoja y no simplemente el color de la mesa?”.
A este fenómeno se le conoce como el Efecto Clever Hans (aprender atajos espurios). En este caso de estudio, construimos herramientas de explicabilidad (Explainable AI - XAI) para auditar visualmente las decisiones de nuestra red neuronal.
Configuración del Entorno y Datos#
Downloading...
From (original): https://drive.google.com/uc?id=1ac6EkoBBxCEnJNcJO438DGHhGQKaTELx
From (redirected): https://drive.google.com/uc?id=1ac6EkoBBxCEnJNcJO438DGHhGQKaTELx&confirm=t&uuid=b188edff-5308-41e6-b6d5-85992f89252c
To: /content/applied-ai-engineering/tea_sickness_dataset.zip
100%|██████████| 775M/775M [00:09<00:00, 83.0MB/s]
GPU detectada: /physical_device:GPU:0
Modelo y Preparación#
Para este notebook, construimos y entrenamos el modelo avanzado del notebook Diseño de la Cabeza de Clasificación.
La siguiente celda, hace todo en un solo bloque…
# Función de utilidad para cargar imágenes
def load_image_for_model(image_path):
img = tf.io.read_file(str(image_path))
img = tf.image.decode_image(img, channels=3, expand_animations=False)
img = tf.image.resize(img, [IMG_SIZE, IMG_SIZE])
img_display = tf.cast(img, tf.float32) / 255.0 # Para matplotlib [0, 1]
img_array = tf.expand_dims(img, 0) # Para el modelo (1, 224, 224, 3)
return img_array.numpy(), img_display.numpy()
# Seleccionamos una imagen de prueba (ej. Anthracnose)
sample_cls = "Anthracnose"
sample_path = list((DATA_DIR / sample_cls).glob("*.jpg"))[0]
img_array, img_display = load_image_for_model(sample_path)
1. Identificando la Capa Objetivo#
Para aplicar métodos basados en gradientes, necesitamos encontrar la última capa convolucional del modelo. Es aquí donde la red tiene la información semántica más rica (formas complejas) antes de colapsarla en un vector 1D.
def find_last_conv_layer(model):
"""Busca la última capa convolucional, incluso si está dentro de un backbone anidado."""
for layer in reversed(model.layers):
# Si es un submodelo (como nuestro EfficientNetB0 base)
if isinstance(layer, keras.Model):
for inner_layer in reversed(layer.layers):
if isinstance(inner_layer, keras.layers.Conv2D):
return inner_layer.name, layer
# Si es una capa plana
if isinstance(layer, keras.layers.Conv2D):
return layer.name, model
raise ValueError("No se encontró ninguna capa Conv2D.")
LAST_CONV_NAME, BACKBONE_MODEL = find_last_conv_layer(modelo)
print(f"Capa objetivo para Grad-CAM: '{LAST_CONV_NAME}'")
Capa objetivo para Grad-CAM: 'top_conv'
# La capa encontrada está dentro del submodelo "efficientnetb0"
modelo.summary()
Model: "Head_SpatialAttention"
┏━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━┓ ┃ Layer (type) ┃ Output Shape ┃ Param # ┃ ┡━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━╇━━━━━━━━━━━━━━━━━━━━━━━━╇━━━━━━━━━━━━━━━┩ │ input_layer_1 (InputLayer) │ (None, 224, 224, 3) │ 0 │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ sequential (Sequential) │ (None, 224, 224, 3) │ 0 │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ efficientnetb0 (Functional) │ (None, 7, 7, 1280) │ 4,049,571 │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ spatial_attention │ (None, 7, 7, 1280) │ 99 │ │ (SpatialAttention) │ │ │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ global_average_pooling2d │ (None, 1280) │ 0 │ │ (GlobalAveragePooling2D) │ │ │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ batch_normalization │ (None, 1280) │ 5,120 │ │ (BatchNormalization) │ │ │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ dropout (Dropout) │ (None, 1280) │ 0 │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ dense (Dense) │ (None, 256) │ 327,936 │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ dense_1 (Dense) │ (None, 8) │ 2,056 │ └─────────────────────────────────┴────────────────────────┴───────────────┘
Total params: 5,050,086 (19.26 MB)
Trainable params: 332,651 (1.27 MB)
Non-trainable params: 4,052,131 (15.46 MB)
Optimizer params: 665,304 (2.54 MB)
2. Implementación Nativa de Grad-CAM#
Grad-CAM (Gradient-weighted Class Activation Mapping) responde a la pregunta: “¿Qué partes de la imagen activaron más fuertemente a la clase ganadora?”.
Matemáticamente, calcula el gradiente (la derivada) del score de la clase ganadora respecto a los mapas de características de la última capa convolucional. Luego, promedia estos gradientes para obtener la “importancia” de cada canal, y multiplica los mapas por esta importancia.
Criterio de Ingeniería: Implementar esto con
tf.GradientTapenos libera de librerías de terceros que se rompen constantemente. La “cinta” graba las operaciones matemáticas en tiempo real para poder calcular las derivadas hacia atrás.
def make_gradcam_heatmap(img_array, model, backbone_model_instance, last_conv_layer_name, pred_index=None):
# 1. Creamos un sub-modelo que escupe las activaciones de la capa conv Y las predicciones finales
# Usamos la instancia del backbone ya identificada y la asignamos directamente.
backbone = backbone_model_instance
grad_model = keras.Model(
inputs=backbone.inputs,
outputs=[backbone.get_layer(last_conv_layer_name).output, backbone.output]
)
# 2. Grabamos las operaciones con GradientTape
img_tensor = tf.cast(img_array, tf.float32)
# Aplicamos el preprocesamiento de la capa de entrada de nuestro modelo principal
x_preprocessed = eff_preprocess(img_tensor)
with tf.GradientTape() as tape:
# Le decimos a la cinta que vigile el tensor de entrada
tape.watch(x_preprocessed)
# Obtenemos las activaciones (ej. 7x7x1280) y la salida del backbone
conv_outputs, backbone_preds = grad_model(x_preprocessed, training=False)
# Pasamos la salida del backbone por la cabeza de clasificación de nuestro modelo principal
x = backbone_preds
for layer in model.layers[model.layers.index(backbone)+1:]:
x = layer(x, training=False)
preds = x
if pred_index is None:
pred_index = tf.argmax(preds[0])
# El "Score" de la clase ganadora (lo que queremos maximizar)
class_channel = preds[:, pred_index]
# 3. Calculamos los gradientes del score respecto a las activaciones convolucionales
grads = tape.gradient(class_channel, conv_outputs)
# 4. Promediamos los gradientes espacialmente (Global Average Pooling manual)
# Esto nos da un vector de pesos de importancia para cada canal
pooled_grads = tf.reduce_mean(grads, axis=(0, 1, 2))
# 5. Multiplicamos cada canal por su "importancia" y los sumamos
conv_outputs = conv_outputs[0]
heatmap = conv_outputs @ pooled_grads[..., tf.newaxis]
heatmap = tf.squeeze(heatmap)
# 6. Aplicamos ReLU (solo nos importan las características que tienen influencia POSITIVA)
heatmap = tf.maximum(heatmap, 0) / tf.math.reduce_max(heatmap)
return heatmap.numpy(), pred_index.numpy(), preds[0, pred_index].numpy()
# Generamos el Heatmap
heatmap, clase_idx, prob = make_gradcam_heatmap(img_array, modelo, BACKBONE_MODEL, LAST_CONV_NAME)
print(f"Heatmap generado con éxito. Forma: {heatmap.shape}")
Heatmap generado con éxito. Forma: (7, 7)
Visualización: El Diagnóstico del Experto#
Vamos a superponer este mapa de calor (que actualmente es de 7x7 píxeles) sobre la imagen original de 224x224.
def display_gradcam(img_display, heatmap, alpha=0.5):
# Redimensionamos el heatmap al tamaño de la imagen
hm_resized = np.array(PILImage.fromarray(np.uint8(255 * heatmap)).resize(
(img_display.shape[1], img_display.shape[0]), PILImage.BILINEAR)) / 255.0
# Aplicamos un mapa de colores (Jet: Azul=Bajo, Rojo=Alto)
cmap = plt.colormaps["jet"]
hm_colored = cmap(hm_resized)[..., :3]
# Superponemos
overlay = np.clip((1 - alpha) * img_display + alpha * hm_colored, 0, 1)
fig, axes = plt.subplots(1, 3, figsize=(15, 5))
axes[0].imshow(img_display); axes[0].set_title("Original"); axes[0].axis('off')
axes[1].imshow(heatmap, cmap='jet'); axes[1].set_title("Grad-CAM (Crudo)"); axes[1].axis('off')
axes[2].imshow(overlay); axes[2].set_title(f"Superposición\nPredicción: {class_names[clase_idx]} ({prob*100:.1f}%)"); axes[2].axis('off')
plt.show()
display_gradcam(img_display, heatmap)
Auditoría de Ingeniería:
Observar la imagen superpuesta. ¿Dónde están las zonas rojas (alta activación)?
Si el rojo está sobre las manchas marrones de la hoja, ¡felicidades! El modelo “aprendió botánica”.
Si el rojo está sobre el fondo blanco/rosado, el modelo es un Clever Hans. Aprendió a clasificar el color de la mesa. Se debería volver al notebook Diseño de la Cabeza de Clasificación y aplicar recortes (cropping) o segmentación de fondo, experimentar con variantes en las capas de “data augmentation” para intentar mitigar los sesgos.
¿Ocurre lo mismo en todas las imágenes?
Va un pequeño bench sobre varias imágenes de una clase. Notar qué pasa cuando la predicción es incorrecta…
print(f"\n--- Benchmarking Grad-CAM para la clase: {sample_cls} ---")
# Obtenemos todas las rutas de imágenes para la clase de ejemplo
all_sample_paths = list((DATA_DIR / sample_cls).glob("*.jpg"))
# Seleccionamos un número de imágenes para el benchmark (ej. 5 imágenes)
num_images_for_bench = min(5, len(all_sample_paths))
selected_paths = all_sample_paths[:num_images_for_bench]
for i, path in enumerate(selected_paths):
print(f"\nProcesando imagen {i+1}/{num_images_for_bench}: {path.name}")
img_array, img_display = load_image_for_model(path)
# Generar el Heatmap
heatmap, clase_idx, prob = make_gradcam_heatmap(img_array, modelo, BACKBONE_MODEL, LAST_CONV_NAME)
# Mostrar el resultado de Grad-CAM
display_gradcam(img_display, heatmap)
--- Benchmarking Grad-CAM para la clase: Anthracnose ---
Procesando imagen 1/5: IMG_20220503_145116.jpg
Procesando imagen 2/5: IMG_20220503_143401.jpg
Procesando imagen 3/5: IMG_20220503_143647.jpg
Procesando imagen 4/5: IMG_20220503_143639.jpg
Procesando imagen 5/5: IMG_20220503_145507.jpg
3. LIME: Explicaciones Agnósticas (Black-Box)#
Grad-CAM es excelente, pero requiere acceso a los gradientes internos de la red (White-box). ¿Qué pasa si estamos auditando un modelo de un proveedor externo al que solo le podemos enviar imágenes y recibir predicciones (API)?
Usamos LIME (Local Interpretable Model-agnostic Explanations).
LIME divide la imagen en “superpíxeles”. Luego, apaga y enciende estos superpíxeles aleatoriamente miles de veces, enviando las imágenes rotas al modelo. Al ver cómo cambia la predicción, LIME deduce qué superpíxeles son los responsables de la decisión.
Va una implementación simple y rudimentaria para entender el concepto:
def lime_explain_image(img_array, img_display, model, n_segments=40, n_samples=150):
print("Generando perturbaciones LIME (esto tomará unos segundos)...")
# 1. Segmentación en superpíxeles (SLIC)
img_uint8 = (img_display * 255).astype(np.uint8)
segments = slic(img_uint8, n_segments=n_segments, compactness=10, start_label=0)
n_segs = segments.max() + 1
# 2. Generamos perturbaciones binarias (1=encendido, 0=apagado)
perturbations = np.random.randint(0, 2, (n_samples, n_segs))
preds = np.zeros(n_samples)
# Color de fondo para los píxeles "apagados" (gris medio)
bg_color = img_array[0].mean(axis=(0, 1))
# 3. Evaluamos cada perturbación en el modelo
for i, perm in enumerate(perturbations):
perturbed_img = img_array[0].copy()
for seg_id in range(n_segs):
if perm[seg_id] == 0: # Si está apagado, lo pintamos de gris
perturbed_img[segments == seg_id] = bg_color
# Predicción del modelo black-box
pred = model.predict(tf.expand_dims(perturbed_img, 0), verbose=0)[0]
preds[i] = pred[clase_idx] # Guardamos la prob de la clase original
# 4. Entrenamos una regresión lineal simple para encontrar la importancia de cada segmento
ridge = Ridge(alpha=1.0)
ridge.fit(perturbations, preds)
coefs = ridge.coef_
# 5. Visualización
pos_mask = np.zeros_like(segments, dtype=bool)
neg_mask = np.zeros_like(segments, dtype=bool)
for seg_id in range(n_segs):
if coefs[seg_id] > 0: pos_mask |= (segments == seg_id)
else: neg_mask |= (segments == seg_id)
overlay_lime = img_display.copy()
# Pintamos de verde lo que apoya la decisión, de rojo lo que va en contra
overlay_lime[pos_mask] = overlay_lime[pos_mask] * 0.5 + np.array([0, 0.8, 0]) * 0.5
overlay_lime[neg_mask] = overlay_lime[neg_mask] * 0.5 + np.array([0.8, 0, 0]) * 0.5
fig, axes = plt.subplots(1, 2, figsize=(10, 5))
axes[0].imshow(mark_boundaries(img_display, segments))
axes[0].set_title(f"Superpíxeles ({n_segs} regiones)")
axes[0].axis('off')
axes[1].imshow(overlay_lime)
axes[1].set_title("LIME\nVerde = A favor | Rojo = En contra")
axes[1].axis('off')
plt.show()
lime_explain_image(img_array, img_display, modelo)
Generando perturbaciones LIME (esto tomará unos segundos)...
print(f"\n--- Benchmarking LIME para la clase: {sample_cls} ---")
all_sample_paths_lime = list((DATA_DIR / sample_cls).glob("*.jpg"))
# Seleccionamos un número de imágenes para el benchmark (ej. 5 imágenes)
num_images_for_bench_lime = min(5, len(all_sample_paths_lime))
selected_paths_lime = all_sample_paths_lime[:num_images_for_bench_lime]
for i, path in enumerate(selected_paths_lime):
print(f"\nProcesando imagen {i+1}/{num_images_for_bench_lime}: {path.name}")
img_array_lime, img_display_lime = load_image_for_model(path)
# Medir el tiempo de inferencia de LIME
start_time = time.perf_counter()
lime_explain_image(img_array_lime, img_display_lime, modelo)
end_time = time.perf_counter()
latencia_ms_lime = (end_time - start_time) * 1000
print(f"Latencia de LIME para esta imagen: {latencia_ms_lime:.1f} ms")
print("\n--- Benchmark de LIME completado ---")
--- Benchmarking LIME para la clase: Anthracnose ---
Procesando imagen 1/5: IMG_20220503_145116.jpg
Generando perturbaciones LIME (esto tomará unos segundos)...
Latencia de LIME para esta imagen: 14298.7 ms
Procesando imagen 2/5: IMG_20220503_143401.jpg
Generando perturbaciones LIME (esto tomará unos segundos)...
Latencia de LIME para esta imagen: 14291.2 ms
Procesando imagen 3/5: IMG_20220503_143647.jpg
Generando perturbaciones LIME (esto tomará unos segundos)...
Latencia de LIME para esta imagen: 14279.4 ms
Procesando imagen 4/5: IMG_20220503_143639.jpg
Generando perturbaciones LIME (esto tomará unos segundos)...
Latencia de LIME para esta imagen: 14424.7 ms
Procesando imagen 5/5: IMG_20220503_145507.jpg
Generando perturbaciones LIME (esto tomará unos segundos)...
Latencia de LIME para esta imagen: 15439.2 ms
--- Benchmark de LIME completado ---
Notar qué pasa si se aplica LIME sobre la misma imagen, la explicabilidad de las regiones no siempre es igual… (la última del bench es la misma analizada individualmente al inicio) La latencia, es considerablemente mayor a Grad-CAM.
Comparativa Grad-CAM vs LIME#
Método |
Tipo |
Ventajas |
Desventajas |
Uso Recomendado |
|---|---|---|---|---|
Grad-CAM |
White-box |
Rápido, alta resolución, gratuito |
Requiere acceso interno |
Auditoría interna, debugging |
LIME |
Model-agnostic |
Funciona con cualquier modelo/API |
Lento, inestable, baja resolución |
Auditoría de modelos externos |
Recomendación de Ingeniería: Usar Grad-CAM durante desarrollo y LIME cuando se auditan modelos de terceros o APIs black-box.
4. Explicabilidad en Producción: El Microservicio#
En un entorno de producción real, no se ejecutarían celdas de Jupyter. Se construye una clase que encapsula el modelo y la lógica de Grad-CAM, lista para ser inyectada en una API (ej. FastAPI).
Va un pequeño ejemplo también:
class ProductionExplainer:
"""
Microservicio de explicabilidad. Cachea el sub-modelo de gradientes
para garantizar latencias < 1000 ms en producción.
"""
def __init__(self, model, backbone_model_instance, last_conv_layer_name, class_names):
self.model = model
self.class_names = class_names
self.last_conv_layer_name = last_conv_layer_name
# Cacheamos el sub-modelo de Grad-CAM en memoria
# Usamos la instancia del backbone ya identificada y la asignamos directamente.
self.backbone = backbone_model_instance
self.grad_model = keras.Model(
inputs=self.backbone.inputs,
outputs=[self.backbone.get_layer(self.last_conv_layer_name).output, self.backbone.output]
)
def explain(self, img_array):
t0 = time.perf_counter()
# Lógica optimizada de Grad-CAM
x_preprocessed = eff_preprocess(tf.cast(img_array, tf.float32))
with tf.GradientTape() as tape:
tape.watch(x_preprocessed)
conv_outputs, backbone_preds = self.grad_model(x_preprocessed, training=False)
x = backbone_preds
for layer in self.model.layers[self.model.layers.index(self.backbone)+1:]:
x = layer(x, training=False)
preds = x
pred_idx = tf.argmax(preds[0])
class_channel = preds[:, pred_idx]
grads = tape.gradient(class_channel, conv_outputs)
pooled_grads = tf.reduce_mean(grads, axis=(0, 1, 2))
heatmap = tf.maximum(tf.squeeze(conv_outputs[0] @ pooled_grads[..., tf.newaxis]), 0)
heatmap = heatmap / (tf.math.reduce_max(heatmap) + 1e-8)
latencia_ms = (time.perf_counter() - t0) * 1000
return {
"prediccion": self.class_names[pred_idx.numpy()],
"confianza": float(preds[0, pred_idx]),
"heatmap_crudo": heatmap.numpy(),
"latencia_ms": latencia_ms
}
# Simulamos una llamada a la API
explainer_api = ProductionExplainer(modelo, BACKBONE_MODEL, LAST_CONV_NAME, class_names)
respuesta = explainer_api.explain(img_array)
print("--- Respuesta de la API ---")
print(f"Predicción: {respuesta['prediccion']} ({respuesta['confianza']*100:.1f}%)")
print(f"Latencia de Explicación: {respuesta['latencia_ms']:.1f} ms")
--- Respuesta de la API ---
Predicción: gray light (30.4%)
Latencia de Explicación: 618.9 ms
Lecciones de Ingeniería y Checklist#
La explicabilidad no es magia; tiene limitaciones matemáticas severas:
Correlación no es Causalidad: Que el mapa de calor se ilumine en una mancha de la hoja no significa que la red entienda qué es una enfermedad. Solo significa que los píxeles de esa zona activaron matemáticamente los filtros que conducen a esa clase.
Resolución Espacial: Grad-CAM produce mapas de 7x7 píxeles (para imágenes de 224x224). Es excelente para ubicar regiones generales, pero inútil para segmentación a nivel de píxel.
Checklist de Auditoría Pre-Deploy:
Ejecutar Grad-CAM en al menos 10 imágenes de cada clase.
Verificar que el mapa de calor se concentra en el objeto de interés y no en el fondo.
Ejecutar LIME en los casos donde el modelo falla con alta confianza (Falsos Positivos graves) para entender qué “engañó” a la red.
Referencias y Lecturas Recomendadas#
Artículos Fundacionales#
Selvaraju, R. R., Cogswell, M., Das, A., Vedantam, R., Parikh, D., & Batra, D. (2016). Grad-CAM: Visual Explanations from Deep Networks via Gradient-based Localization. ICCV.
[arXiv]
(Paper original de Grad-CAM — lectura obligatoria).Ribeiro, M. T., Singh, S., & Guestrin, C. (2016). “Why Should I Trust You?”: Explaining the Predictions of Any Classifier. KDD.
[arXiv]
(Paper original de LIME).
Recursos Prácticos y Actuales#
Keras Team. Grad-CAM implementation example. Keras Documentation.
Molnar, C. (2020). Interpretable Machine Learning. (Libro gratuito excelente sobre explicabilidad).