当前位置:首页 > 技术 > 正文内容

基于PyTorch的卷积神经网络图像分类与特征图可视化实践

访客 技术 2026年8月9日 1

在深度学习图像分类任务中,构建卷积神经网络(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
)

相关文章

Linux crontab 详解

1) crontab 是什么cron 是 Linux 的定时任务守护进程;crontab 是用来编辑/查看“按时间周期执行命令”的表(cron table)。常见两类:用户 crontab:每个用户一份(crontab -e 编辑)系统级 crontab / cron.d:可指定执行用户(/etc/crontab、/etc/cron.d/*)2) crontab 时间...

富文本里可以允许的 HTML 属性

一、所有标签默认允许的安全属性(极少)class        (可选)id           (通常建议禁用)title️ 注意:id 容易被滥用做锚点注入,很多系统直接禁用class 允许的话最好只允许固定前缀(如 editor-*)二、a 标签允许属性<a href="" t...

Mac 安装 Node.js 指南

方法一:通过官网安装包(最简单,适合初学者)如果你只是想快速安装并开始使用,这是最直接的方法。访问 Node.js 官网。页面会显示两个版本:LTS (Recommended For Most Users):长期支持版,最稳定。建议选这个。Current:最新特性版,包含最新功能但可能不够稳定。下载 .pkg 安装包并运行。按照安装向导点击“下一步”即可完成。方法二:使用 Homebrew 安装(...

Dom\HTML_NO_DEFAULT_NS 的副作用:自动加闭合标签

在使用Dom\HTMLDocument时,Dom\HTML_NO_DEFAULT_NS 将禁止在解析过程中设置元素的命名空间, 此设置是为了与DOMDocument向后兼容而存在的。当使用它时,已知的一个副作用就是:自动加闭合标签例如 </img> 为什么会这样?当你使用:Dom\HTML_NO_DEFAULT_NS文档会变成 无命名空间模式,此时内部更接近 XML...

Laravel 事件和监听器创建

在 Laravel 中,使用 Artisan 命令创建 Events(事件) 和 Listeners(监听器) 是非常高效的。你可以通过以下几种方式来实现:1. 手动创建单个 Event如果你只想创建一个事件类,可以使用 make:event 命令:Bashphp artisan make:event UserRegistered执行后,文件将生成在 app/Even...

自定义域名解析神器 dnsmasq

什么是 dnsmasq?dnsmasq 是一个轻量级、功能强大的网络服务工具,专为小型和中等规模网络设计。它是一个综合的网络基础设施解决方案[1]。dnsmasq 能做什么?功能说明应用场景DNS 转发与缓存将 DNS 查询转发到上游服务器(ISP、Google DNS 等),并在本地缓存结果加快 DNS 查询速度,减少外部 DNS 流量本地 DNS解析本地网络设备的主机名,无需编辑&n...

发表评论

访客

◎欢迎参与讨论,请在这里发表您的看法和观点。