Skip to content
rutwik.devALL ACCESS
ALL POSTS

Explainable AI for Cataract Detection: What GradCAM Taught Us

AUGUST 20, 20248 MIN READAIComputer VisionResearchHealthcare

When we published our cataract detection research at IEEE PuneCon 2023, the headline number was 97% accuracy. But the part that actually mattered, and the part that took the most work, was making the model explain itself.

Here's what we built, what we learned, and why explainability isn't just a nice-to-have in medical AI.

The problem with black-box medical AI

Imagine a doctor using an AI tool that classifies a retinal image as "cataract detected." The doctor has two choices:

  1. Trust the model and recommend surgery
  2. Ignore the model and rely on their own assessment

If the model is a black box with no explanation, just a confidence score, most doctors will choose option 2. And they should. A confidence score of 97% doesn't tell you why the model thinks there's a cataract, which means the doctor can't verify whether the model is looking at the right part of the image.

This is the explainability problem in medical AI. Our research was specifically about solving it.

The architecture: VGG-19 + GradCAM

We used VGG-19, a 19-layer convolutional neural network pretrained on ImageNet, fine-tuned on retinal fundus images from the ORIGA dataset (650 images, class-balanced).

Fine-tuning strategy:

  • Froze all layers except the last 4 convolutional blocks
  • Added a custom classification head: GlobalAveragePooling → Dense(256, ReLU) → Dropout(0.5) → Dense(1, Sigmoid)
  • Trained for 50 epochs with early stopping (patience=10)
  • Data augmentation: rotation ±20°, horizontal flip, brightness ±20%

The 97% accuracy came after about 3 weeks of hyperparameter tuning. But accuracy was the easy part.

How GradCAM works (and why it works for this task)

GradCAM (Gradient-weighted Class Activation Mapping) computes a heatmap showing which parts of the image most influenced the model's decision. The math is surprisingly simple:

def compute_gradcam(model, image, layer_name, class_idx):
    # Get the output of the target convolutional layer
    grad_model = tf.keras.models.Model(
        inputs=model.inputs,
        outputs=[model.get_layer(layer_name).output, model.output]
    )
    
    with tf.GradientTape() as tape:
        conv_outputs, predictions = grad_model(image)
        loss = predictions[:, class_idx]
    
    # Gradients of the class score w.r.t. the conv layer output
    grads = tape.gradient(loss, conv_outputs)
    
    # Pool the gradients over all axes (global average pooling)
    pooled_grads = tf.reduce_mean(grads, axis=(0, 1, 2))
    
    # Weight the conv outputs by the pooled gradients
    conv_outputs = conv_outputs[0]
    heatmap = conv_outputs @ pooled_grads[..., tf.newaxis]
    heatmap = tf.squeeze(heatmap)
    
    # Normalize and resize to original image dimensions
    heatmap = tf.maximum(heatmap, 0) / tf.math.reduce_max(heatmap)
    return cv2.resize(heatmap.numpy(), (image.shape[2], image.shape[1]))

For cataract detection, a well-trained model should produce heatmaps that highlight the lens region, the area where cataracts physically form. If the model highlights the optic disc or the sclera instead, something is wrong with the training data or the model's learned features.

What the heatmaps revealed

When we first generated GradCAM heatmaps, about 15% of them highlighted regions that had nothing to do with the lens. The model was picking up on artifacts: image compression boundaries, metadata pixels, and even the retinal camera's reflection pattern.

This was invisible in the accuracy metric. The model could still classify correctly by chance or by learning spurious correlations. But the heatmaps made the problem obvious.

We fixed it by:

  1. Preprocessing all images to remove known artifacts (center crop to remove border artifacts)
  2. Adding more diverse augmentation to break spurious correlations
  3. Retraining with the last 6 convolutional blocks unfrozen instead of 4

After retraining, >90% of heatmaps correctly highlighted the lens region. Accuracy improved slightly to 97.3%, but more importantly, the model was now looking at the right things.

The clinical implication

We showed these heatmaps to two ophthalmologists who reviewed the model's outputs. Their feedback was unanimous: they would be willing to use a tool that shows them what it's looking at. A heatmap superimposed on a retinal image takes 2 seconds to review and immediately tells a clinician whether the AI is trustworthy on this specific case.

This is the core argument for explainability in medical AI: it shifts the decision from "do I trust this model in general?" to "does this model's reasoning on this specific case make sense?" The latter is a much easier and more reliable question for a human expert to answer.

What I'd do differently

Use a larger dataset. 650 images is small by modern standards. We used cross-validation to compensate, but a model trained on 10K+ images with varied camera hardware would generalize much better across clinics.

Try Segment Anything for preprocessing. Manual cropping worked, but SAM (Meta's Segment Anything Model) would make lens region isolation automatic and more precise, which would remove the need for the center crop heuristic entirely.

Publish the heatmap generation code. We open-sourced the model weights but not the GradCAM visualization pipeline. In retrospect, the visualization code is often more useful to practitioners than the weights.

The published paper

The full paper, "XAI meets Ophthalmology: An Explainable Approach to Cataract Detection using VGG-19 and Grad-CAM", is available on IEEE Xplore. The live demo of the cataract detection system is at cataractdetectionwithxai.streamlit.app.