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()

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:
-
Hook into the activations of the last convolutional layer.
-
Compute gradients of the target class with respect to those activations.
-
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()
