Explicabilidad y Grad-CAM (Abriendo la Caja Negra)#

Open In Colab

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.GradientTape para 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#

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#

Hide code cell source

# @title *Esta celda clona el repositorio, descarga los datos e importa utilidades*
import sys
import os

IN_COLAB = "google.colab" in sys.modules

if IN_COLAB:
    import subprocess
    REPO_NAME = "applied-ai-engineering"
    if not os.path.exists(REPO_NAME):
        subprocess.run(["git", "clone", f"https://github.com/AxelSkrauba/{REPO_NAME}.git"], check=True)
    os.chdir(f"/content/{REPO_NAME}")
    sys.path.append(f"/content/{REPO_NAME}")

    # Descarga del dataset de Enfermedades del Té
    if not os.path.exists("tea_sickness_dataset"):
        import gdown
        gdown.download(f"https://drive.google.com/uc?id=1ac6EkoBBxCEnJNcJO438DGHhGQKaTELx", "tea_sickness_dataset.zip", quiet=False)
        subprocess.run(["unzip", "-qq", "tea_sickness_dataset.zip"], check=True)
        os.remove("tea_sickness_dataset.zip")
else:
    os.chdir(f"../../")

from utils.plots import setup_plot_style
setup_plot_style()

os.environ["KERAS_BACKEND"] = "tensorflow"

import keras
import tensorflow as tf
import numpy as np
import matplotlib.pyplot as plt
import matplotlib.cm as cm
from PIL import Image as PILImage
from pathlib import Path
import time

# Herramientas para LIME "casero"
from skimage.segmentation import slic, mark_boundaries
from sklearn.linear_model import Ridge

keras.utils.set_random_seed(42)

gpus = tf.config.list_physical_devices('GPU')
if gpus:
    tf.config.experimental.set_memory_growth(gpus[0], True)
    print(f"GPU detectada: {gpus[0].name}")
else:
    print("ADVERTENCIA: Sin GPU. Los métodos de explicabilidad serán lentos.")
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…

Hide code cell source

import keras
from keras import layers, ops
import tensorflow as tf

# Hacemos un split rápido en memoria para simplificar
DATA_DIR = Path('./tea_sickness_dataset')
IMG_SIZE = 224
BATCH_SIZE = 32

# Cargamos todo el dataset
full_ds = keras.utils.image_dataset_from_directory(
    DATA_DIR, label_mode='categorical', image_size=(IMG_SIZE, IMG_SIZE), batch_size=BATCH_SIZE, seed=42
)
class_names = full_ds.class_names
n_classes = len(class_names)

# Split manual rápido (80% Train, 20% Val para simplificar)
dataset_size = len(full_ds)
train_size = int(0.8 * dataset_size)
ds_train = full_ds.take(train_size)
ds_val = full_ds.skip(train_size)

# Optimización
AUTOTUNE = tf.data.AUTOTUNE
ds_train = ds_train.cache().prefetch(buffer_size=AUTOTUNE)
ds_val = ds_val.cache().prefetch(buffer_size=AUTOTUNE)

print(f"Clases ({n_classes}): {class_names}")

from keras.applications.efficientnet import preprocess_input as eff_preprocess

# Instanciamos el backbone SIN pooling
backbone_base = keras.applications.EfficientNetB0(
    include_top=False, weights="imagenet", input_shape=(IMG_SIZE, IMG_SIZE, 3), pooling=None
)
backbone_base.trainable = False


@keras.saving.register_keras_serializable()
class SpatialAttention(layers.Layer):
    """
    Aprende qué REGIONES del mapa espacial (7x7) son más relevantes.
    """
    def __init__(self, kernel_size=7, **kwargs):
        super().__init__(**kwargs)
        self.kernel_size = kernel_size

    def build(self, input_shape):
        self.conv = layers.Conv2D(1, self.kernel_size, padding='same', activation='sigmoid')
        super().build(input_shape)

    def call(self, x):
        # Calculamos el promedio y el máximo a través de los canales
        avg_out = tf.reduce_mean(x, axis=-1, keepdims=True)
        max_out = tf.reduce_max(x, axis=-1, keepdims=True)
        concat = tf.concat([avg_out, max_out], axis=-1)

        # Generamos un mapa de atención (pesos entre 0 y 1)
        attn = self.conv(concat)

        # Multiplicamos los features originales por el mapa de atención
        return x * attn

    def get_config(self):
        cfg = super().get_config()
        cfg.update({"kernel_size": self.kernel_size})
        return cfg

def build_model():
    inputs = keras.Input(shape=(IMG_SIZE, IMG_SIZE, 3))

    # Data Augmentation Integrado
    aug = keras.Sequential([
        layers.RandomFlip("horizontal_and_vertical"),
        layers.RandomRotation(0.2),
        layers.RandomZoom(0.1)
    ])
    x = aug(inputs)

    # Preprocesamiento y Backbone
    x = eff_preprocess(x)
    feats = backbone_base(x, training=False) # Backbone congelado

    # Atención Espacial + Label Smoothing
    x = SpatialAttention()(feats)
    x = layers.GlobalAveragePooling2D()(x)
    x = layers.BatchNormalization()(x)
    x = layers.Dropout(0.4)(x)
    x = layers.Dense(256, activation="relu", kernel_regularizer=keras.regularizers.l2(1e-4))(x)
    outputs = layers.Dense(n_classes, activation="softmax")(x)

    model = keras.Model(inputs, outputs, name="Head_SpatialAttention")

    model.compile(
        optimizer=keras.optimizers.Adam(learning_rate=1e-3),
        loss=keras.losses.CategoricalCrossentropy(label_smoothing=0.1),
        metrics=["accuracy"]
    )
    return model

modelo = build_model()

callbacks = [
    keras.callbacks.EarlyStopping(monitor='val_loss', patience=3, restore_best_weights=True)
]

hist = modelo.fit(ds_train, validation_data=ds_val, epochs=20, callbacks=callbacks, verbose=1)

def unfreeze_backbone_safely(modelo, num_layers_to_unfreeze=20):
    """
    Descongela las últimas N capas del backbone, manteniendo BatchNormalization congelado.
    """
    # Buscamos el backbone dentro del modelo
    backbone = [layer for layer in modelo.layers if isinstance(layer, keras.Model)][0]

    backbone.trainable = True
    total_layers = len(backbone.layers)
    freeze_until = total_layers - num_layers_to_unfreeze

    for i, layer in enumerate(backbone.layers):
        if i < freeze_until:
            layer.trainable = False
        else:
            # Regla de Oro: BN siempre congelado en Fine-Tuning
            if isinstance(layer, layers.BatchNormalization):
                layer.trainable = False
            else:
                layer.trainable = True

    print(f"Descongeladas {num_layers_to_unfreeze} capas. Capas BN protegidas.")

# Aplicamos el descongelamiento seguro
unfreeze_backbone_safely(modelo, num_layers_to_unfreeze=20)

# RE-COMPILACIÓN CRÍTICA: Learning Rate microscópico
modelo.compile(
    optimizer=keras.optimizers.Adam(learning_rate=1e-5), # 100x más pequeño que antes
    loss=keras.losses.CategoricalCrossentropy(label_smoothing=0.1),
    metrics=["accuracy"]
)

print("\nIniciando Fine-Tuning Progresivo...")
callbacks_ft = [
    keras.callbacks.EarlyStopping(monitor='val_loss', patience=4, restore_best_weights=True),
    keras.callbacks.ReduceLROnPlateau(monitor='val_loss', factor=0.5, patience=2)
]

hist_ft = modelo.fit(ds_train, validation_data=ds_val, epochs=25, callbacks=callbacks_ft, verbose=1)

Hide code cell output

Found 885 files belonging to 8 classes.
Clases (8): ['Anthracnose', 'algal leaf', 'bird eye spot', 'brown blight', 'gray light', 'healthy', 'red leaf spot', 'white spot']
Downloading data from https://storage.googleapis.com/keras-applications/efficientnetb0_notop.h5
16705208/16705208 ━━━━━━━━━━━━━━━━━━━━ 1s 0us/step
Epoch 1/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 77s 3s/step - accuracy: 0.5014 - loss: 1.7525 - val_accuracy: 0.5470 - val_loss: 1.7708
Epoch 2/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 101ms/step - accuracy: 0.6804 - loss: 1.2495 - val_accuracy: 0.5801 - val_loss: 1.7837
Epoch 3/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 101ms/step - accuracy: 0.7727 - loss: 1.1311 - val_accuracy: 0.5801 - val_loss: 1.6596
Epoch 4/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 91ms/step - accuracy: 0.8295 - loss: 0.9974 - val_accuracy: 0.6685 - val_loss: 1.6481
Epoch 5/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 92ms/step - accuracy: 0.8097 - loss: 1.0242 - val_accuracy: 0.6519 - val_loss: 1.6181
Epoch 6/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 92ms/step - accuracy: 0.8125 - loss: 0.9918 - val_accuracy: 0.7459 - val_loss: 1.5451
Epoch 7/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 92ms/step - accuracy: 0.8523 - loss: 0.9244 - val_accuracy: 0.7403 - val_loss: 1.4643
Epoch 8/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 94ms/step - accuracy: 0.8679 - loss: 0.8904 - val_accuracy: 0.7569 - val_loss: 1.3611
Epoch 9/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 104ms/step - accuracy: 0.8608 - loss: 0.8918 - val_accuracy: 0.7182 - val_loss: 1.3575
Epoch 10/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 93ms/step - accuracy: 0.8835 - loss: 0.8596 - val_accuracy: 0.7790 - val_loss: 1.2559
Epoch 11/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 87ms/step - accuracy: 0.8580 - loss: 0.9047 - val_accuracy: 0.7459 - val_loss: 1.3079
Epoch 12/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 92ms/step - accuracy: 0.8793 - loss: 0.8778 - val_accuracy: 0.8398 - val_loss: 1.1819
Epoch 13/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 93ms/step - accuracy: 0.8807 - loss: 0.8447 - val_accuracy: 0.8122 - val_loss: 1.1694
Epoch 14/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 93ms/step - accuracy: 0.8849 - loss: 0.8423 - val_accuracy: 0.8011 - val_loss: 1.1170
Epoch 15/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 104ms/step - accuracy: 0.8963 - loss: 0.8154 - val_accuracy: 0.8122 - val_loss: 1.1031
Epoch 16/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 87ms/step - accuracy: 0.8878 - loss: 0.8232 - val_accuracy: 0.8177 - val_loss: 1.1112
Epoch 17/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 92ms/step - accuracy: 0.9148 - loss: 0.7910 - val_accuracy: 0.8177 - val_loss: 1.0236
Epoch 18/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 92ms/step - accuracy: 0.9148 - loss: 0.7874 - val_accuracy: 0.8232 - val_loss: 1.0110
Epoch 19/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 92ms/step - accuracy: 0.9091 - loss: 0.7799 - val_accuracy: 0.8343 - val_loss: 0.9557
Epoch 20/20
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 88ms/step - accuracy: 0.9219 - loss: 0.7739 - val_accuracy: 0.8011 - val_loss: 1.0202
Descongeladas 20 capas. Capas BN protegidas.

Iniciando Fine-Tuning Progresivo...
Epoch 1/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 16s 260ms/step - accuracy: 0.9091 - loss: 0.7799 - val_accuracy: 0.8343 - val_loss: 0.9383 - learning_rate: 1.0000e-05
Epoch 2/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 94ms/step - accuracy: 0.9048 - loss: 0.7880 - val_accuracy: 0.8453 - val_loss: 0.9246 - learning_rate: 1.0000e-05
Epoch 3/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 93ms/step - accuracy: 0.9162 - loss: 0.7769 - val_accuracy: 0.8508 - val_loss: 0.9102 - learning_rate: 1.0000e-05
Epoch 4/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 93ms/step - accuracy: 0.9006 - loss: 0.7835 - val_accuracy: 0.8564 - val_loss: 0.8987 - learning_rate: 1.0000e-05
Epoch 5/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 94ms/step - accuracy: 0.9375 - loss: 0.7430 - val_accuracy: 0.8619 - val_loss: 0.8899 - learning_rate: 1.0000e-05
Epoch 6/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 112ms/step - accuracy: 0.9290 - loss: 0.7590 - val_accuracy: 0.8619 - val_loss: 0.8843 - learning_rate: 1.0000e-05
Epoch 7/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 96ms/step - accuracy: 0.9247 - loss: 0.7687 - val_accuracy: 0.8674 - val_loss: 0.8781 - learning_rate: 1.0000e-05
Epoch 8/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 93ms/step - accuracy: 0.9134 - loss: 0.7819 - val_accuracy: 0.8729 - val_loss: 0.8717 - learning_rate: 1.0000e-05
Epoch 9/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 93ms/step - accuracy: 0.9048 - loss: 0.7791 - val_accuracy: 0.8785 - val_loss: 0.8664 - learning_rate: 1.0000e-05
Epoch 10/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 94ms/step - accuracy: 0.9247 - loss: 0.7676 - val_accuracy: 0.8785 - val_loss: 0.8632 - learning_rate: 1.0000e-05
Epoch 11/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 94ms/step - accuracy: 0.9162 - loss: 0.7757 - val_accuracy: 0.8785 - val_loss: 0.8581 - learning_rate: 1.0000e-05
Epoch 12/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 3s 103ms/step - accuracy: 0.9276 - loss: 0.7669 - val_accuracy: 0.8785 - val_loss: 0.8541 - learning_rate: 1.0000e-05
Epoch 13/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 94ms/step - accuracy: 0.9261 - loss: 0.7688 - val_accuracy: 0.8840 - val_loss: 0.8521 - learning_rate: 1.0000e-05
Epoch 14/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 94ms/step - accuracy: 0.9375 - loss: 0.7452 - val_accuracy: 0.8785 - val_loss: 0.8486 - learning_rate: 1.0000e-05
Epoch 15/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 95ms/step - accuracy: 0.9190 - loss: 0.7718 - val_accuracy: 0.8785 - val_loss: 0.8486 - learning_rate: 1.0000e-05
Epoch 16/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 95ms/step - accuracy: 0.9304 - loss: 0.7597 - val_accuracy: 0.8785 - val_loss: 0.8481 - learning_rate: 1.0000e-05
Epoch 17/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 3s 130ms/step - accuracy: 0.9304 - loss: 0.7630 - val_accuracy: 0.8785 - val_loss: 0.8451 - learning_rate: 1.0000e-05
Epoch 18/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 3s 122ms/step - accuracy: 0.9290 - loss: 0.7564 - val_accuracy: 0.8785 - val_loss: 0.8456 - learning_rate: 1.0000e-05
Epoch 19/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 114ms/step - accuracy: 0.9418 - loss: 0.7439 - val_accuracy: 0.8785 - val_loss: 0.8466 - learning_rate: 1.0000e-05
Epoch 20/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 110ms/step - accuracy: 0.9347 - loss: 0.7479 - val_accuracy: 0.8785 - val_loss: 0.8447 - learning_rate: 5.0000e-06
Epoch 21/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 109ms/step - accuracy: 0.9332 - loss: 0.7671 - val_accuracy: 0.8785 - val_loss: 0.8443 - learning_rate: 5.0000e-06
Epoch 22/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 105ms/step - accuracy: 0.9134 - loss: 0.7678 - val_accuracy: 0.8785 - val_loss: 0.8440 - learning_rate: 5.0000e-06
Epoch 23/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 102ms/step - accuracy: 0.9247 - loss: 0.7829 - val_accuracy: 0.8785 - val_loss: 0.8424 - learning_rate: 5.0000e-06
Epoch 24/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 95ms/step - accuracy: 0.9403 - loss: 0.7390 - val_accuracy: 0.8785 - val_loss: 0.8415 - learning_rate: 5.0000e-06
Epoch 25/25
22/22 ━━━━━━━━━━━━━━━━━━━━ 2s 96ms/step - accuracy: 0.9190 - loss: 0.7685 - val_accuracy: 0.8785 - val_loss: 0.8409 - learning_rate: 5.0000e-06
# 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.GradientTape nos 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)
../../../_images/cd9e8edaede06c80ad935c448d64f4a6b359fc5ab138fc947136a98381a34141.png

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
../../../_images/cd9e8edaede06c80ad935c448d64f4a6b359fc5ab138fc947136a98381a34141.png
Procesando imagen 2/5: IMG_20220503_143401.jpg
../../../_images/b3087823d44abd35ddb844747bdf128e2b1be9ca6c9badbe940fd9d1ac8199b6.png
Procesando imagen 3/5: IMG_20220503_143647.jpg
../../../_images/730353a36c0cdd72aea2e9ea49a6c43d0a231daa5a61df0094f04878909e4a2d.png
Procesando imagen 4/5: IMG_20220503_143639.jpg
../../../_images/16e462beef7d638eff1073713db544aeb6f840016d886e42fc26f239118bc0f2.png
Procesando imagen 5/5: IMG_20220503_145507.jpg
../../../_images/e9b4beb4b501591f237266c5ea718fc9e6f3a016c291f806ac6cabbf994dcfa8.png

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)...
../../../_images/14c578ca81dda9a0551b555aaeb5bcf8bf64073d152e6e996e7a8a8136bc83a5.png
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)...
../../../_images/81cf2148bc51a03f8ac625e675676644ee489c72ea321e5393f6022f2f95b041.png
Latencia de LIME para esta imagen: 14298.7 ms

Procesando imagen 2/5: IMG_20220503_143401.jpg
Generando perturbaciones LIME (esto tomará unos segundos)...
../../../_images/9d0ed7867988e09c3ef808a5abf104c64bdf5c681761294b9a676c7bb7ce3e86.png
Latencia de LIME para esta imagen: 14291.2 ms

Procesando imagen 3/5: IMG_20220503_143647.jpg
Generando perturbaciones LIME (esto tomará unos segundos)...
../../../_images/b6787a5582a0233264b2c5596458ed7fa6a373cd74ba7413b60c1967196eac6b.png
Latencia de LIME para esta imagen: 14279.4 ms

Procesando imagen 4/5: IMG_20220503_143639.jpg
Generando perturbaciones LIME (esto tomará unos segundos)...
../../../_images/25c459edadeda98e0e133a221225e789e11dafc8a1bfec22294a6056964d75e1.png
Latencia de LIME para esta imagen: 14424.7 ms

Procesando imagen 5/5: IMG_20220503_145507.jpg
Generando perturbaciones LIME (esto tomará unos segundos)...
../../../_images/4af71c4cbc0ddc65edb7c96aa6030279e6576dd0156b08862d8e0e4a3a46bdde.png
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:

  1. 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.

  2. 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#

  1. 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).

  2. 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#


Entorno de Ejecución#

Hide code cell source

from utils.environment import environment_table
environment_table(include_all=True)

Hide code cell output

Reproducibility Environment Information
Package Version
Python 3.12.13
Platform Linux-6.6.122+-x86_64-with-glibc2.35
Cython 3.0.12
IPython 7.34.0
OpenSSL 26.2.0
PIL 11.3.0
anywidget 0.9.21
argparse 1.1
astunparse 1.6.3
attr 26.1.0
backcall 0.2.0
bottleneck 1.4.2
brotli 1.2.0
bs4 4.13.5
certifi 2026.5.20
chardet 5.2.0
charset_normalizer 3.4.7
cloudpickle 3.1.2
cryptography 48.0.0
csv 1.0
ctypes 1.1.0
cuda 12.9.6
cv2 4.13.0
cycler 0.12.1
cython 3.0.12
dateutil 2.9.0.post0
debugpy 1.8.15
decimal 1.70
decorator 4.4.2
defusedxml 0.7.1
dill 0.3.8
entrypoints 0.4
etils 1.14.0
filelock 3.29.0
flatbuffers 25.12.19
gast 0.7.0
gdown 5.2.2
google 3.0.0
google_auth_httplib2 0.4.0
googleapiclient 0.1.3
h5py 3.16.0
html5lib 1.1
http 0.6
httplib2 0.31.2
huggingface_hub 1.16.1
idna 17.0.0
ipaddress 1.0
ipykernel 6.17.1
ipython_genutils 0.2.0
ipywidgets 7.7.1
jax 0.7.2
jax_cuda12_plugin 0.7.2
jaxlib 0.7.2
joblib 1.5.3
json 2.0.9
jupyter_client 7.4.9
jupyter_core 5.9.1
keras 3.13.2
kiwisolver 1.5.0
lazy_loader 0.5
logging 0.5.1.2
lxml 6.1.1
matplotlib 3.10.0
matplotlib_inline 0.2.2
ml_dtypes 0.5.4
namex 0.1.0
numexpr 2.14.1
numpy 2.0.2
oauth2client 4.1.3
opt_einsum 3.4.0
optree 0.19.1
packaging 26.2
pandas 2.2.2
patsy 1.0.2
pexpect 4.9.0
pickleshare 0.7.5
platformdirs 4.9.6
prompt_toolkit 3.0.52
psutil 5.9.5
psygnal 0.15.1
ptyprocess 0.7.0
pyarrow 18.1.0
pyasn1 0.6.3
pyasn1_modules 0.4.2
pydevd 3.2.3
pydot 4.0.1
pygments 2.20.0
pyparsing 3.3.2
pytz 2025.2
rapids_dask_dependency 26.2.0
re 2.2.1
requests 2.32.4
rich 13.9.4
rsa 4.9.1
scipy 1.16.3
seaborn 0.13.2
simplejson 4.1.1
six 1.17.0
skimage 0.25.2
sklearn 1.5.3
socketserver 0.4
socks 1.7.1
soupsieve 2.8.3
statsmodels 0.14.6
tblib 3.2.2
tensorflow 2.20.0
termcolor 3.3.0
threadpoolctl 3.6.0
tornado 6.5.1
tqdm 4.67.3
traitlets 5.7.1
typing_extensions 4.15.0
uritemplate 4.2.0
urllib 3.12
urllib3 2.5.0
wcwidth 0.7.0
webencodings 0.5.1
wrapt 2.2.0
xmlrpc 3.12
zlib 1.0
zmq 26.2.1
zstandard 0.25.0