基于PyTorch的卷积神经网络图像分类与特征图可视化实践
在深度学习图像分类任务中,构建卷积神经网络(CNN)并理解其内部特征是提升模型性能与可解释性的关键。本文将详细探讨如何使用 PyTorch 实现一个完整的图像分类 Pipeline,并通过 Grad-CAM 和中间层特征图可视化技术来解析模型的决策依据。
一、 图像分类模型构建与训练
首先,我们需要准备数据集、定义数据增强策略、构建 CNN 架构,并编写训练与验证循环。以下代码展示了一个针对乐器图像分类的完整实现,包含了数据加载、模型定义、训练调度以及 Grad-CAM 热力图生成。
import os
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import Dataset, DataLoader, random_split
from torchvision import transforms
from PIL import Image
from pytorch_grad_cam import GradCAM
from pytorch_grad_cam.utils.image import show_cam_on_image
# 1. 自定义数据集类
class InstrumentDataset(Dataset):
def __init__(self, data_dir, metadata_csv, img_transform=None):
self.data_dir = data_dir
self.metadata = pd.read_csv(metadata_csv)
self.img_transform = img_transform
self.class_names = sorted(self.metadata['class'].unique())
self.class_to_idx = {cls: idx for idx, cls in enumerate(self.class_names)}
self.file_paths = []
self.targets = []
for _, row in self.metadata.iterrows():
cls_name = row['class']
cls_path = os.path.join(data_dir, cls_name)
img_files = sorted([f for f in os.listdir(cls_path) if f.lower().endswith('.jpg')])
for img_file in img_files:
self.file_paths.append(os.path.join(cls_path, img_file))
self.targets.append(self.class_to_idx[cls_name])
def __len__(self):
return len(self.file_paths)
def __getitem__(self, index):
img_path = self.file_paths[index]
image = Image.open(img_path).convert('RGB')
label = self.targets[index]
if self.img_transform:
image = self.img_transform(image)
return image, label
# 2. 数据增强与预处理
data_transforms = transforms.Compose([
transforms.RandomHorizontalFlip(p=0.5),
transforms.RandomAffine(degrees=10, translate=(0.1, 0.1)),
transforms.ColorJitter(brightness=0.1, contrast=0.1, saturation=0.1),
transforms.Resize((256, 256)),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])
# 3. 实例化数据集与加载器
base_dir = "./data/music_instruments"
csv_path = "./data/music_instruments/stats.csv"
full_dataset = InstrumentDataset(base_dir, csv_path, img_transform=data_transforms)
train_ratio = 0.8
train_size = int(train_ratio * len(full_dataset))
val_size = len(full_dataset) - train_size
train_subset, val_subset = random_split(
full_dataset, [train_size, val_size],
generator=torch.Generator().manual_seed(99)
)
train_loader = DataLoader(train_subset, batch_size=64, shuffle=True, num_workers=2)
val_loader = DataLoader(val_subset, batch_size=64, shuffle=False, num_workers=2)
# 4. 定义卷积神经网络
class InstrumentClassifier(nn.Module):
def __init__(self, num_classes):
super(InstrumentClassifier, self).__init__()
self.feature_extractor = nn.Sequential(
# Block 1
nn.Conv2d(3, 16, kernel_size=3, padding=1),
nn.BatchNorm2d(16),
nn.ReLU(inplace=True),
nn.MaxPool2d(2, 2),
# Block 2
nn.Conv2d(16, 32, kernel_size=3, padding=1),
nn.BatchNorm2d(32),
nn.ReLU(inplace=True),
nn.MaxPool2d(2, 2),
# Block 3
nn.Conv2d(32, 64, kernel_size=3, padding=1),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.MaxPool2d(2, 2)
)
# 224 -> 112 -> 56 -> 28. 28*28*64 = 50176
self.classifier_head = nn.Sequential(
nn.Flatten(),
nn.Linear(64 * 28 * 28, 256),
nn.ReLU(inplace=True),
nn.Dropout(0.4),
nn.Linear(256, num_classes)
)
def forward(self, x):
features = self.feature_extractor(x)
logits = self.classifier_head(features)
return logits
# 5. 训练与评估逻辑
def execute_training(model, train_dl, val_dl, loss_fn, optimizer, scheduler, device, epochs):
history = {'train_loss': [], 'val_loss': [], 'train_acc': [], 'val_acc': []}
for epoch in range(epochs):
model.train()
running_loss, correct, total = 0.0, 0, 0
for inputs, labels in train_dl:
inputs, labels = inputs.to(device), labels.to(device)
optimizer.zero_grad()
outputs = model(inputs)
loss = loss_fn(outputs, labels)
loss.backward()
optimizer.step()
running_loss += loss.item() * inputs.size(0)
_, preds = torch.max(outputs, 1)
correct += (preds == labels).sum().item()
total += labels.size(0)
epoch_train_loss = running_loss / total
epoch_train_acc = correct / total
model.eval()
val_loss, val_correct, val_total = 0.0, 0, 0
with torch.no_grad():
for inputs, labels in val_dl:
inputs, labels = inputs.to(device), labels.to(device)
outputs = model(inputs)
loss = loss_fn(outputs, labels)
val_loss += loss.item() * inputs.size(0)
_, preds = torch.max(outputs, 1)
val_correct += (preds == labels).sum().item()
val_total += labels.size(0)
epoch_val_loss = val_loss / val_total
epoch_val_acc = val_correct / val_total
scheduler.step(epoch_val_loss)
history['train_loss'].append(epoch_train_loss)
history['train_acc'].append(epoch_train_acc)
history['val_loss'].append(epoch_val_loss)
history['val_acc'].append(epoch_val_acc)
print(f"Epoch {epoch+1}/{epochs} | "
f"Train Loss: {epoch_train_loss:.4f} Acc: {epoch_train_acc:.4f} | "
f"Val Loss: {epoch_val_loss:.4f} Acc: {epoch_val_acc:.4f}")
return history
# 初始化与执行
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
num_classes = len(full_dataset.class_names)
net = InstrumentClassifier(num_classes).to(device)
criterion = nn.CrossEntropyLoss()
opt = optim.AdamW(net.parameters(), lr=1e-3, weight_decay=1e-3)
lr_scheduler = optim.lr_scheduler.ReduceLROnPlateau(opt, mode='min', factor=0.5, patience=3)
train_history = execute_training(net, train_loader, val_loader, criterion, opt, lr_scheduler, device, epochs=25)
# 6. Grad-CAM 可视化
target_layer = net.feature_extractor[8] # 选择 Block 3 的 Conv2d 层
cam_extractor = GradCAM(model=net, target_layers=[target_layer])
sample_imgs, sample_labels = next(iter(val_loader))
input_tensor = sample_imgs[0:1].to(device)
cam_mask = cam_extractor(input_tensor=input_tensor, targets=None)
inv_normalize = transforms.Normalize(
mean=[-0.485/0.229, -0.456/0.224, -0.406/0.225],
std=[1/0.229, 1/0.224, 1/0.225]
)
orig_img = inv_normalize(sample_imgs[0]).permute(1, 2, 0).cpu().numpy()
orig_img = np.clip(orig_img, 0, 1)
cam_visualization = show_cam_on_image(orig_img, cam_mask[0], use_rgb=True)
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(10, 5))
ax1.imshow(orig_img)
ax1.set_title("Original Input")
ax1.axis('off')
ax2.imshow(cam_visualization)
ax2.set_title("Grad-CAM Heatmap")
ax2.axis('off')
plt.show()
二、 模型可解释性:中间层特征图提取
除了 Grad-CAM,直接观察网络中间层的特征图也是理解 CNN 工作机制的有效手段。通过注册前向钩子(Forward Hook),我们可以无损地捕获网络各层的输出张量。以下代码采用面向对象的方式管理钩子,并利用 torchvision.utils.make_grid 实现多通道特征的高效网格化展示。
import torchvision.utils as vutils
class ActivationHook:
"""用于捕获模块输出的钩子类"""
def __init__(self):
self.features = None
def __call__(self, module, module_in, module_out):
self.features = module_out.detach().cpu()
def extract_and_visualize_features(model, data_loader, target_layer_names, device, num_samples=2, channels_to_show=8):
"""
提取并可视化指定层的特征图
"""
model.eval()
hooks = {}
# 注册钩子到目标层
for name, module in model.named_modules():
if name in target_layer_names:
hook = ActivationHook()
module.register_forward_hook(hook)
hooks[name] = hook
# 获取样本数据
images, labels = next(iter(data_loader))
images = images[:num_samples].to(device)
with torch.no_grad():
_ = model(images)
# 反归一化函数
inv_normalize = transforms.Normalize(
mean=[-0.485/0.229, -0.456/0.224, -0.406/0.225],
std=[1/0.229, 1/0.224, 1/0.225]
)
for i in range(num_samples):
orig_img = inv_normalize(images[i]).cpu()
fig = plt.figure(figsize=(15, 5 * len(target_layer_names)))
# 绘制原图
ax = fig.add_subplot(len(target_layer_names), 1, 1)
ax.imshow(orig_img.permute(1, 2, 0).numpy())
ax.set_title(f"Sample {i+1} - Original Image")
ax.axis('off')
# 绘制各层特征图
for row_idx, layer_name in enumerate(target_layer_names):
if layer_name in hooks:
feature_maps = hooks[layer_name].features[i]
# 限制显示的通道数并增加单通道维度以适配 make_grid
maps_to_show = feature_maps[:channels_to_show].unsqueeze(1)
grid = vutils.make_grid(
maps_to_show,
nrow=4,
padding=2,
normalize=True,
scale_each=True,
pad_value=0.5
)
ax = fig.add_subplot(len(target_layer_names), 1, row_idx + 2)
ax.imshow(grid.permute(1, 2, 0).numpy(), cmap='viridis')
ax.set_title(f"Layer: {layer_name} (First {channels_to_show} channels)")
ax.axis('off')
plt.tight_layout()
plt.show()
# 调用示例:检查三个卷积块的第一个卷积层
layers_to_inspect = ['feature_extractor.0', 'feature_extractor.4', 'feature_extractor.8']
extract_and_visualize_features(
model=net,
data_loader=val_loader,
target_layer_names=layers_to_inspect,
device=device,
num_samples=2,
channels_to_show=12
)