๐Ÿ‘๏ธ Computer Vision ยท Lecture 26 of 27

Explaining Vision Models: Saliency Maps, Grad-CAM and Their Limits

Which pixels made the model decide? We study gradient saliency, Grad-CAM, integrated gradients and occlusion, show how they reveal shortcuts, and discuss sanity checks that expose unreliable explanations.

A model classifies a chest X-ray as "pneumonia" with 94% confidence. A doctor asks: why? Did it look at the lungs โ€” or at a hospital's text marker in the corner? Explanation methods for vision produce heatmaps highlighting the image regions that most influenced a prediction. They are invaluable for debugging and building appropriate trust โ€” but they can also mislead if used naively.

Gradient saliency#

The simplest idea (Simonyan et al., 2013): compute the gradient of the class score $y_c$ with respect to the input pixels,

$$ S(\mathbf{x}) = \left|\frac{\partial y_c}{\partial\mathbf{x}}\right| $$

(maximum over colour channels). Large values mark pixels whose small changes would most affect the score. Vanilla gradients are noisy; SmoothGrad averages gradients over several noisy copies of the input for cleaner maps.

Grad-CAM#

Grad-CAM (Selvaraju et al., 2017) is the most widely used method for CNNs. It works at the last convolutional layer, whose feature maps $A^k$ retain spatial layout while encoding high-level semantics.

  1. Compute gradients of the class score with respect to each feature map.
  2. Global-average-pool the gradients to get an importance weight per channel:
$$ \alpha_k^c = \frac{1}{Z}\sum_{i,j}\frac{\partial y_c}{\partial A_{ij}^k} $$
  1. Take a weighted combination of feature maps and keep positive evidence:
$$ L^c_{\text{Grad-CAM}} = \text{ReLU}\left(\sum_k\alpha_k^c A^k\right) $$
  1. Upsample to image size and overlay as a heatmap.

Grad-CAM is class-discriminative (different classes highlight different regions), needs no retraining and works with any CNN.

python
import torch
import torch.nn.functional as F
from torchvision import models

model = models.resnet50(weights=models.ResNet50_Weights.DEFAULT).eval()
activations, gradients = {}, {}
layer = model.layer4
layer.register_forward_hook(lambda m, i, o: activations.__setitem__("a", o))
layer.register_full_backward_hook(lambda m, gi, go: gradients.__setitem__("g", go[0]))

def grad_cam(x, class_idx=None):
    logits = model(x)
    c = logits.argmax(1).item() if class_idx is None else class_idx
    model.zero_grad(); logits[0, c].backward()
    A, G = activations["a"], gradients["g"]              # (1, K, h, w)
    weights = G.mean(dim=(2, 3), keepdim=True)           # alpha_k
    cam = F.relu((weights * A).sum(1, keepdim=True))
    cam = F.interpolate(cam, size=x.shape[2:], mode="bilinear", align_corners=False)
    return ((cam - cam.min()) / (cam.max() - cam.min() + 1e-8))[0, 0].detach(), c

x = torch.randn(1, 3, 224, 224)                          # replace with a normalised real image
heatmap, predicted = grad_cam(x)
print(heatmap.shape, predicted)

The pytorch-grad-cam library provides Grad-CAM and many variants (Grad-CAM++, Score-CAM, Eigen-CAM) for CNNs and ViTs.

Integrated gradients#

Plain gradients can be near zero when a feature is saturated even though it matters. Integrated Gradients (Sundararajan et al., 2017) integrates gradients along a path from a baseline $\mathbf{x}'$ (e.g. a black image) to the input:

$$ \text{IG}_i(\mathbf{x}) = (x_i - x_i')\int_0^1\frac{\partial F(\mathbf{x}' + \alpha(\mathbf{x} - \mathbf{x}'))}{\partial x_i}\,d\alpha $$

It satisfies completeness: attributions sum to $F(\mathbf{x}) - F(\mathbf{x}')$. The choice of baseline matters and should be justified.

Perturbation methods#

Occlusion: slide a grey patch over the image and record how much the class score drops โ€” model-agnostic and intuitive, but slow. RISE averages random masks weighted by the resulting scores. LIME fits a local interpretable model over superpixels (see the XAI lecture in the Ethics track).

Explanations reveal shortcuts#

Saliency maps have exposed many shortcut behaviours: models attending to watermarks or copyright text, to snow in "wolf" photos, to hospital markers in X-rays, or to the background rather than the object. Routinely inspect explanations for a sample of correct and incorrect predictions during model development.

Explanations can mislead#

Beyond heatmaps#

  • Concept-based explanations (TCAV) test whether a human-defined concept (e.g. "striped") influences a class.
  • Prototype networks explain by comparing parts of the input with learned prototypical parts ("this looks like that").
  • Counterfactual explanations show a minimally changed image that would change the decision.
  • Mechanistic interpretability studies individual neurons and circuits inside networks.
JA
Written by

Janin A Apurba

B.Sc. in CSE, AUST ยท Advanced ICT Officer, CNRS-UNHCR. Teaching AI, ML and Deep Learning to the next generation of engineers and researchers.

Keep learning

Related lectures

๐Ÿ‘๏ธ Computer Vision

Vision Foundation Models: Segment Anything and Promptable Vision

Vision is following language towards general-purpose foundation models. We study the Segment Anything Model's promptable design and data engine, open-vocabulary detection, and how foundation models change vision workflows.

Advancedโฑ 5 min#159
๐Ÿ‘๏ธ Computer Vision

Adversarial Examples: Fooling Neural Networks and Defending Them

Imperceptible perturbations can make a network confidently wrong. We derive FGSM and PGD attacks, explain why adversarial examples exist, cover physical and black-box attacks, and evaluate defences including adversarial training.

Advancedโฑ 5 min#161
๐Ÿ‘๏ธ Computer Vision

AI in Medical Imaging: Opportunities, Pitfalls and Validation

Deep learning can detect disease in X-rays, retinal scans and pathology slides. We survey modalities and tasks, discuss data and labelling challenges, shortcut learning, rigorous clinical validation, and deployment responsibilities.

Intermediateโฑ 5 min#158