Skip to article frontmatterSkip to article content
Site not loading correctly?

This may be due to an incorrect BASE_URL configuration. See the MyST Documentation for reference.

3.16 Whale Shark Interpretability

Open In Colab

Overview

Model interpretability asks not just what a model predicts, but why. GradCAM (Gradient-weighted Class Activation Mapping) highlights the image regions that most influenced a classification decision by computing gradients of the class score with respect to the final convolutional feature map. For whale shark photo-ID, this reveals whether the model attends to the animal’s distinctive spot pattern rather than background water or boat artefacts.

The Wild Me whale shark dataset contains annotated images of individual animals identified by their unique spot patterns, which function like fingerprints. We train a ResNet18 classifier to recognize the top 5 most-photographed individuals, then apply GradCAM to inspect which image regions drive each prediction.

Learning Objectives

By the end of this lesson you will be able to:

  1. Load and parse a Wild Me COCO annotation file to extract individual identity labels.

  2. Fine-tune a pretrained ResNet18 classifier on a small marine photo-ID dataset.

  3. Apply GradCAM to visualize which image regions drive each classification.

  4. Critically evaluate whether model attention aligns with biologically meaningful features.

Dataset: Wild Me Whale Shark Archive

The Wild Me whale shark COCO dataset is hosted on LILA BC (Labeled Information Library of Alexandria: Biology and Conservation). It contains thousands of images of individual whale sharks photographed across multiple field sites, with each animal identified by its unique spot pattern on the flank behind the gills.

The COCO annotation format stores data as a JSON file with three top-level lists: images (file paths and IDs), annotations (bounding boxes and identity labels), and categories (class definitions). Here, the name field in each annotation entry stores the individual whale shark’s string identifier rather than a numeric category, so we extract it directly.

Downloading dataset...
Download complete.
Extracting dataset...
/tmp/ipykernel_1458/1007797648.py:17: DeprecationWarning: Python 3.14 will, by default, filter extracted tar archives and reject files or modify their metadata. Use the filter argument to control this behavior.
  tar.extractall(path=extract_dir)
Extraction complete.

Loading Annotations and Building the Dataset

We parse the COCO JSON to build a mapping from image_id to individual name, then construct a PyTorch Dataset that loads each image, applies a standard ImageNet normalization transform, and returns the image tensor plus the label index. We restrict to the top 5 most-photographed individuals to keep the dataset manageable and the class balance reasonable.

Created dataset with 678 images across 5 individual whale sharks.
Top 5 individuals being classified: ['341569f2-1f34-4884-1dd3-79137be4c77f', 'a785af89-b8c0-5e7b-acec-c4874ec5483f', '24e404dd-094c-e5bd-defa-e52125422443', '21027863-1d99-20d5-e7f9-dd3f7ad3b166', '26560de1-6930-ddaf-5069-f7b85acd40fb']

Training the Classifier

We fine-tune ResNet18 pre-trained on ImageNet. The final fully connected layer is replaced with a fresh 5-class output head sized for our individual count. The model trains for 30 epochs with the Adam optimizer and cross-entropy loss. On a Colab T4 GPU this takes approximately 5 to 10 minutes. On CPU, plan for significantly longer.

The key transfer learning insight: ImageNet pre-training gives the model strong low-level feature detectors (edges, textures, color gradients). Fine-tuning on whale shark images redirects those detectors toward the spot pattern, fin shape, and body markings that distinguish individuals.

Starting training...
Epoch 1/30 - Loss: 1.7029
Epoch 2/30 - Loss: 1.3642
Epoch 3/30 - Loss: 1.3048
Epoch 4/30 - Loss: 1.1987
Epoch 5/30 - Loss: 0.9746
Epoch 6/30 - Loss: 0.9277
Epoch 7/30 - Loss: 0.8428
Epoch 8/30 - Loss: 0.6293
Epoch 9/30 - Loss: 0.6125
Epoch 10/30 - Loss: 0.4083
Epoch 11/30 - Loss: 0.5477
Epoch 12/30 - Loss: 0.3892
Epoch 13/30 - Loss: 0.3364
Epoch 14/30 - Loss: 0.2693
Epoch 15/30 - Loss: 0.1908
Epoch 16/30 - Loss: 0.3070
Epoch 17/30 - Loss: 0.2449
Epoch 18/30 - Loss: 0.2288
Epoch 19/30 - Loss: 0.1971
Epoch 20/30 - Loss: 0.1876
Epoch 21/30 - Loss: 0.1744
Epoch 22/30 - Loss: 0.0899
Epoch 23/30 - Loss: 0.0791
Epoch 24/30 - Loss: 0.0795
Epoch 25/30 - Loss: 0.0628
Epoch 26/30 - Loss: 0.0873
Epoch 27/30 - Loss: 0.0754
Epoch 28/30 - Loss: 0.0613
Epoch 29/30 - Loss: 0.0717
Epoch 30/30 - Loss: 0.0524
Training complete.

GradCAM Interpretability Analysis

GradCAM works through five steps:

  1. Select a target class (e.g., individual whale shark number 3).

  2. Compute the gradient of that class score with respect to the activations of the final convolutional layer (layer4 in ResNet18).

  3. Global average pool those gradients to produce per-channel importance weights.

  4. Compute a weighted sum of the feature maps using those weights.

  5. Apply ReLU to keep only positive activations (regions that support the predicted class).

The resulting heatmap is resized to the input image dimensions and overlaid using show_cam_on_image(). Red regions indicate strong positive influence on the prediction; blue regions indicate low or negative influence.

We target model.layer4[-1] because it is the deepest convolutional block, producing the most semantically meaningful (but spatially coarsest) feature maps before the global average pool.

Requirement already satisfied: grad-cam in /usr/local/lib/python3.12/dist-packages (1.5.5)
Requirement already satisfied: numpy in /usr/local/lib/python3.12/dist-packages (from grad-cam) (2.0.2)
Requirement already satisfied: Pillow in /usr/local/lib/python3.12/dist-packages (from grad-cam) (11.3.0)
Requirement already satisfied: torch>=1.7.1 in /usr/local/lib/python3.12/dist-packages (from grad-cam) (2.11.0+cu128)
Requirement already satisfied: torchvision>=0.8.2 in /usr/local/lib/python3.12/dist-packages (from grad-cam) (0.26.0+cu128)
Requirement already satisfied: ttach in /usr/local/lib/python3.12/dist-packages (from grad-cam) (0.0.3)
Requirement already satisfied: tqdm in /usr/local/lib/python3.12/dist-packages (from grad-cam) (4.67.3)
Requirement already satisfied: opencv-python in /usr/local/lib/python3.12/dist-packages (from grad-cam) (4.13.0.92)
Requirement already satisfied: matplotlib in /usr/local/lib/python3.12/dist-packages (from grad-cam) (3.10.0)
Requirement already satisfied: scikit-learn in /usr/local/lib/python3.12/dist-packages (from grad-cam) (1.6.1)
Requirement already satisfied: filelock in /usr/local/lib/python3.12/dist-packages (from torch>=1.7.1->grad-cam) (3.29.3)
Requirement already satisfied: typing-extensions>=4.10.0 in /usr/local/lib/python3.12/dist-packages (from torch>=1.7.1->grad-cam) (4.15.0)
Requirement already satisfied: setuptools<82 in /usr/local/lib/python3.12/dist-packages (from torch>=1.7.1->grad-cam) (75.2.0)
Requirement already satisfied: sympy>=1.13.3 in /usr/local/lib/python3.12/dist-packages (from torch>=1.7.1->grad-cam) (1.14.0)
Requirement already satisfied: networkx>=2.5.1 in /usr/local/lib/python3.12/dist-packages (from torch>=1.7.1->grad-cam) (3.6.1)
Requirement already satisfied: jinja2 in /usr/local/lib/python3.12/dist-packages (from torch>=1.7.1->grad-cam) (3.1.6)
Requirement already satisfied: fsspec>=0.8.5 in /usr/local/lib/python3.12/dist-packages (from torch>=1.7.1->grad-cam) (2025.3.0)
Requirement already satisfied: cuda-toolkit==12.8.1 in /usr/local/lib/python3.12/dist-packages (from cuda-toolkit[cublas,cudart,cufft,cufile,cupti,curand,cusolver,cusparse,nvjitlink,nvrtc,nvtx]==12.8.1; platform_system == "Linux"->torch>=1.7.1->grad-cam) (12.8.1)
Requirement already satisfied: cuda-bindings<13,>=12.9.4 in /usr/local/lib/python3.12/dist-packages (from torch>=1.7.1->grad-cam) (12.9.7)
Requirement already satisfied: nvidia-cudnn-cu12==9.19.0.56 in /usr/local/lib/python3.12/dist-packages (from torch>=1.7.1->grad-cam) (9.19.0.56)
Requirement already satisfied: nvidia-cusparselt-cu12==0.7.1 in /usr/local/lib/python3.12/dist-packages (from torch>=1.7.1->grad-cam) (0.7.1)
Requirement already satisfied: nvidia-nccl-cu12==2.28.9 in /usr/local/lib/python3.12/dist-packages (from torch>=1.7.1->grad-cam) (2.28.9)
Requirement already satisfied: nvidia-nvshmem-cu12==3.4.5 in /usr/local/lib/python3.12/dist-packages (from torch>=1.7.1->grad-cam) (3.4.5)
Requirement already satisfied: triton==3.6.0 in /usr/local/lib/python3.12/dist-packages (from torch>=1.7.1->grad-cam) (3.6.0)
Requirement already satisfied: nvidia-cublas-cu12==12.8.4.1.* in /usr/local/lib/python3.12/dist-packages (from cuda-toolkit[cublas,cudart,cufft,cufile,cupti,curand,cusolver,cusparse,nvjitlink,nvrtc,nvtx]==12.8.1; platform_system == "Linux"->torch>=1.7.1->grad-cam) (12.8.4.1)
Requirement already satisfied: nvidia-cuda-runtime-cu12==12.8.90.* in /usr/local/lib/python3.12/dist-packages (from cuda-toolkit[cublas,cudart,cufft,cufile,cupti,curand,cusolver,cusparse,nvjitlink,nvrtc,nvtx]==12.8.1; platform_system == "Linux"->torch>=1.7.1->grad-cam) (12.8.90)
Requirement already satisfied: nvidia-cufft-cu12==11.3.3.83.* in /usr/local/lib/python3.12/dist-packages (from cuda-toolkit[cublas,cudart,cufft,cufile,cupti,curand,cusolver,cusparse,nvjitlink,nvrtc,nvtx]==12.8.1; platform_system == "Linux"->torch>=1.7.1->grad-cam) (11.3.3.83)
Requirement already satisfied: nvidia-cufile-cu12==1.13.1.3.* in /usr/local/lib/python3.12/dist-packages (from cuda-toolkit[cublas,cudart,cufft,cufile,cupti,curand,cusolver,cusparse,nvjitlink,nvrtc,nvtx]==12.8.1; platform_system == "Linux"->torch>=1.7.1->grad-cam) (1.13.1.3)
Requirement already satisfied: nvidia-cuda-cupti-cu12==12.8.90.* in /usr/local/lib/python3.12/dist-packages (from cuda-toolkit[cublas,cudart,cufft,cufile,cupti,curand,cusolver,cusparse,nvjitlink,nvrtc,nvtx]==12.8.1; platform_system == "Linux"->torch>=1.7.1->grad-cam) (12.8.90)
Requirement already satisfied: nvidia-curand-cu12==10.3.9.90.* in /usr/local/lib/python3.12/dist-packages (from cuda-toolkit[cublas,cudart,cufft,cufile,cupti,curand,cusolver,cusparse,nvjitlink,nvrtc,nvtx]==12.8.1; platform_system == "Linux"->torch>=1.7.1->grad-cam) (10.3.9.90)
Requirement already satisfied: nvidia-cusolver-cu12==11.7.3.90.* in /usr/local/lib/python3.12/dist-packages (from cuda-toolkit[cublas,cudart,cufft,cufile,cupti,curand,cusolver,cusparse,nvjitlink,nvrtc,nvtx]==12.8.1; platform_system == "Linux"->torch>=1.7.1->grad-cam) (11.7.3.90)
Requirement already satisfied: nvidia-cusparse-cu12==12.5.8.93.* in /usr/local/lib/python3.12/dist-packages (from cuda-toolkit[cublas,cudart,cufft,cufile,cupti,curand,cusolver,cusparse,nvjitlink,nvrtc,nvtx]==12.8.1; platform_system == "Linux"->torch>=1.7.1->grad-cam) (12.5.8.93)
Requirement already satisfied: nvidia-nvjitlink-cu12==12.8.93.* in /usr/local/lib/python3.12/dist-packages (from cuda-toolkit[cublas,cudart,cufft,cufile,cupti,curand,cusolver,cusparse,nvjitlink,nvrtc,nvtx]==12.8.1; platform_system == "Linux"->torch>=1.7.1->grad-cam) (12.8.93)
Requirement already satisfied: nvidia-cuda-nvrtc-cu12==12.8.93.* in /usr/local/lib/python3.12/dist-packages (from cuda-toolkit[cublas,cudart,cufft,cufile,cupti,curand,cusolver,cusparse,nvjitlink,nvrtc,nvtx]==12.8.1; platform_system == "Linux"->torch>=1.7.1->grad-cam) (12.8.93)
Requirement already satisfied: nvidia-nvtx-cu12==12.8.90.* in /usr/local/lib/python3.12/dist-packages (from cuda-toolkit[cublas,cudart,cufft,cufile,cupti,curand,cusolver,cusparse,nvjitlink,nvrtc,nvtx]==12.8.1; platform_system == "Linux"->torch>=1.7.1->grad-cam) (12.8.90)
Requirement already satisfied: contourpy>=1.0.1 in /usr/local/lib/python3.12/dist-packages (from matplotlib->grad-cam) (1.3.3)
Requirement already satisfied: cycler>=0.10 in /usr/local/lib/python3.12/dist-packages (from matplotlib->grad-cam) (0.12.1)
Requirement already satisfied: fonttools>=4.22.0 in /usr/local/lib/python3.12/dist-packages (from matplotlib->grad-cam) (4.63.0)
Requirement already satisfied: kiwisolver>=1.3.1 in /usr/local/lib/python3.12/dist-packages (from matplotlib->grad-cam) (1.5.0)
Requirement already satisfied: packaging>=20.0 in /usr/local/lib/python3.12/dist-packages (from matplotlib->grad-cam) (26.2)
Requirement already satisfied: pyparsing>=2.3.1 in /usr/local/lib/python3.12/dist-packages (from matplotlib->grad-cam) (3.3.2)
Requirement already satisfied: python-dateutil>=2.7 in /usr/local/lib/python3.12/dist-packages (from matplotlib->grad-cam) (2.9.0.post0)
Requirement already satisfied: scipy>=1.6.0 in /usr/local/lib/python3.12/dist-packages (from scikit-learn->grad-cam) (1.16.3)
Requirement already satisfied: joblib>=1.2.0 in /usr/local/lib/python3.12/dist-packages (from scikit-learn->grad-cam) (1.5.3)
Requirement already satisfied: threadpoolctl>=3.1.0 in /usr/local/lib/python3.12/dist-packages (from scikit-learn->grad-cam) (3.6.0)
Requirement already satisfied: cuda-pathfinder~=1.1 in /usr/local/lib/python3.12/dist-packages (from cuda-bindings<13,>=12.9.4->torch>=1.7.1->grad-cam) (1.5.5)
Requirement already satisfied: six>=1.5 in /usr/local/lib/python3.12/dist-packages (from python-dateutil>=2.7->matplotlib->grad-cam) (1.17.0)
Requirement already satisfied: mpmath<1.4,>=1.1.0 in /usr/local/lib/python3.12/dist-packages (from sympy>=1.13.3->torch>=1.7.1->grad-cam) (1.3.0)
Requirement already satisfied: MarkupSafe>=2.0 in /usr/local/lib/python3.12/dist-packages (from jinja2->torch>=1.7.1->grad-cam) (3.0.3)

Visualizing Grad-CAM for individual: 341569f2-1f34-4884-1dd3-79137be4c77f
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>

Visualizing Grad-CAM for individual: a785af89-b8c0-5e7b-acec-c4874ec5483f
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>

Visualizing Grad-CAM for individual: 24e404dd-094c-e5bd-defa-e52125422443
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>

Visualizing Grad-CAM for individual: 21027863-1d99-20d5-e7f9-dd3f7ad3b166
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>

Visualizing Grad-CAM for individual: 26560de1-6930-ddaf-5069-f7b85acd40fb
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>
<Figure size 1000x500 with 2 Axes>

Reflecting on Results

Examine the GradCAM overlays carefully before answering these questions:

  1. Does GradCAM highlight the whale shark’s spot pattern, or does it focus on other image regions such as background water, the boat, or the diver? What does this tell you about what the model has learned?

  2. How would you design an experiment to distinguish between a model that learned the spot pattern versus one that learned to recognize correlated background cues such as water color or location?

  3. What modifications to the training pipeline (data augmentation, cropping strategy, loss function) might improve the biological meaningfulness of the learned features?