Back to stuffs
Deep Learning

MNIST with PyTorch + GradCAM

•4 min read
#deep learning
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
import torchvision.transforms as transforms
import torchvision.datasets as datasets
# Define a transform to normalize the data
transform = transforms.Compose([
    transforms.ToTensor(), 
    transforms.Normalize((0.5,), (0.5,))  # Normalize with mean=0.5, std=0.5
])

train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform)
test_dataset = datasets.MNIST(root='./data', train=False, download=True, transform=transform)

train_loader = DataLoader(dataset=train_dataset, batch_size=64, shuffle=True)
test_loader = DataLoader(dataset=test_dataset, batch_size=64, shuffle=False)

Build the model

class CNN(nn.Module):
    def __init__(self):
        super(CNN, self).__init__()
        self.conv1 = nn.Conv2d(1, 32, kernel_size=3, stride=1, padding=1)
        self.relu = nn.ReLU()
        self.pool = nn.MaxPool2d(kernel_size=2, stride=2)
        self.fc1 = nn.Linear(32 * 14 * 14, 10)  # 14x14 after pooling

    def forward(self, x):
        x = self.conv1(x)
        x = self.relu(x)
        x = self.pool(x)
        x = x.view(x.size(0), -1)  # Flatten the tensor
        x = self.fc1(x)
        return x

# Instantiate the model
model = CNN()

Assign Optimizer & Loss Function

criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)

Train the model

for epoch in range(1):  # One epoch
    for batch_idx, (data, target) in enumerate(train_loader):
        optimizer.zero_grad()  # Zero the gradients
        output = model(data)  # Forward pass
        loss = criterion(output, target)  # Compute loss
        loss.backward()  # Backward pass
        optimizer.step()  # Update weights

        if batch_idx % 100 == 0:
            print(f'Epoch: {epoch}, Batch: {batch_idx}, Loss: {loss.item()}')
Epoch: 0, Batch: 0, Loss: 2.2880163192749023

Epoch: 0, Batch: 100, Loss: 0.40903452038764954

Epoch: 0, Batch: 200, Loss: 0.2600383460521698

Epoch: 0, Batch: 300, Loss: 0.25161853432655334

Epoch: 0, Batch: 400, Loss: 0.18361416459083557

Epoch: 0, Batch: 500, Loss: 0.07740028202533722

Epoch: 0, Batch: 600, Loss: 0.15956029295921326

Epoch: 0, Batch: 700, Loss: 0.08738137781620026

Epoch: 0, Batch: 800, Loss: 0.10718843340873718

Epoch: 0, Batch: 900, Loss: 0.11650585383176804
model.eval()  # Set the model to evaluation mode
correct = 0
total = 0

with torch.no_grad():  # Disable gradient calculation for evaluation
    for data, target in test_loader:
        output = model(data)
        _, predicted = torch.max(output.data, 1)  # Get the predicted class
        total += target.size(0)
        correct += (predicted == target).sum().item()

print(f'Test Accuracy: {100 * correct / total:.2f}%')
Test Accuracy: 97.59%
# Save the model
torch.save(model.state_dict(), 'mnist_cnn.pth')

# Load the model
model = CNN()  # Recreate the model architecture
model.load_state_dict(torch.load('mnist_cnn.pth'))
model.eval()  # Set to evaluation mode
CNN(

  (conv1): Conv2d(1, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))

  (relu): ReLU()

  (pool): MaxPool2d(kernel_size=2, stride=2, padding=0, dilation=1, ceil_mode=False)

  (fc1): Linear(in_features=6272, out_features=10, bias=True)

)

Visualize

import matplotlib.pyplot as plt
import numpy as np

# Set the model to evaluation mode
model.eval()

# Get a batch of test images
dataiter = iter(test_loader)
images, labels = next(dataiter)

# Make predictions
with torch.no_grad():
    outputs = model(images)
    _, predicted = torch.max(outputs, 1)

# Convert images to numpy arrays for visualization
images = images.numpy()

# Plot the images with their predicted and actual labels
fig, axes = plt.subplots(4, 4, figsize=(10, 10))
for idx, ax in enumerate(axes.ravel()):
    ax.imshow(np.squeeze(images[idx]), cmap='gray')
    ax.set_title(f'Pred: {predicted[idx].item()}\nTrue: {labels[idx].item()}')
    ax.axis('off')

plt.tight_layout()
plt.show()

png

Adding Grad-CAM (Gradient-weighted Class Activation Mapping) is a great way to visualize which parts of an image the model is focusing on when making predictions. Grad-CAM highlights the regions of the image that were most influential in the model's decision.

To implement Grad-CAM in PyTorch, we need to:

  1. Hook into the activations of the last convolutional layer.

  2. Compute gradients of the target class with respect to those activations.

  3. Combine the activations and gradients to produce the heatmap.

import torch.nn.functional as F

# Hook into the activations of the last convolutional layer
activations = None
def hook_fn(module, input, output):
    global activations
    activations = output

# Register the hook
model.conv1.register_forward_hook(hook_fn)

# Get a batch of test images
dataiter = iter(test_loader)
images, labels = next(dataiter)

# Forward pass
images.requires_grad_()  # Enable gradient computation for the input images
outputs = model(images)
_, predicted = torch.max(outputs, 1)

# Backward pass to get gradients
model.zero_grad()
loss = F.cross_entropy(outputs, labels)
loss.backward(retain_graph=True)  # Retain the graph for Grad-CAM computation

# Grad-CAM computation
grads = torch.autograd.grad(loss, activations, retain_graph=True)[0]  # Retain graph here as well
pooled_grads = grads.mean(dim=(2, 3), keepdim=True)
heatmap = (activations * pooled_grads).sum(dim=1, keepdim=True)
heatmap = F.relu(heatmap)  # Apply ReLU to the heatmap
heatmap /= heatmap.max()  # Normalize the heatmap

# Convert heatmap to numpy for visualization
heatmap = heatmap.squeeze().detach().numpy()

# Plot the original image and heatmap
fig, axes = plt.subplots(4, 4, figsize=(10, 10))
for idx, ax in enumerate(axes.ravel()):
    ax.imshow(np.squeeze(images[idx].detach().numpy()), cmap='gray')
    ax.imshow(heatmap[idx], cmap='jet', alpha=0.5)  # Overlay heatmap
    ax.set_title(f'Pred: {predicted[idx].item()}\nTrue: {labels[idx].item()}')
    ax.axis('off')

plt.tight_layout()
plt.show()

png