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,
(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.
- Compute gradients of the class score with respect to each feature map.
- Global-average-pool the gradients to get an importance weight per channel:
- Take a weighted combination of feature maps and keep positive evidence:
- 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.
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:
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.