Skip to content

PyTorch - 卷积网络可视化

理解和可视化您的输入数据是任何深度学习项目(尤其是在处理图像任务的卷积神经网络,即 ConvNets 或 CNNs 时)至关重要的第一步。本章演示了如何使用 torchvision 加载一个标准数据集(如 MNIST),并使用 matplotlib 可视化其中的一些图像。

我们将需要 torch,用于数据集和图像转换的 torchvision,以及用于绘图的 matplotlib.pyplot。

import torch
import torchvision
import torchvision.transforms as transforms
import matplotlib.pyplot as plt
import numpy as np # For converting tensors to numpy arrays for plotting
# 将张量(Tensor)转换为 NumPy 数组(array)以进行绘图

torchvision.datasets 提供了轻松访问流行数据集(如 MNIST、CIFAR10 等)的功能。我们还将使用 torch.utils.data.DataLoader 以批次(batch)为单位迭代数据。transforms.ToTensor() 将 PIL 图像或 NumPy ndarray 转换为 PyTorch 张量(Tensor),并将像素值缩放到 [0.0, 1.0] 的范围。

# 为数据集定义转换(transform)
# ToTensor() 将 PIL Image 或 numpy.ndarray 转换为 FloatTensor。
# 它还会将图像的像素强度值缩放到 [0., 1.] 的范围。
transform = transforms.Compose([
transforms.ToTensor(),
# 您可以根据模型需要添加其他转换,例如 Normalize
# transforms.Normalize((0.5,), (0.5,)) # 单通道图像的示例, (均值,), (标准差,)
])
# 加载 MNIST 训练数据集
# download=True 将在根目录(root directory)中找不到数据集时下载它
train_dataset = torchvision.datasets.MNIST(root='./data',
train=True,
transform=transform,
download=True)
# 创建一个 DataLoader 以迭代数据集
# batch_size 定义了每个批次加载多少样本
train_loader = torch.utils.data.DataLoader(dataset=train_dataset,
batch_size=4,
shuffle=True) # 打乱顺序以获得更好的训练效果

此代码下载 MNIST 数据集(如果 ./data 目录中尚未存在),应用定义的转换(transform),并准备好数据加载器(data loader)。

现在,让我们从 train_loader 中获取一个图像批次(batch)并显示它们。

# 显示图像的函数
def imshow(img_tensor, title=None):
"""Imshow for Tensor."""
# PyTorch 张量通常是 (通道数 C, 高 H, 宽 W) 的格式
# Matplotlib 需要 (高 H, 宽 W, 通道数 C) 或 (高 H, 宽 W)(用于灰度图)的格式
img_numpy = img_tensor.numpy() # 将张量转换为 numpy 数组
# 如果是彩色图像,调整通道顺序。对于灰度图像 (1, H, W),压缩通道维度。
if img_numpy.shape[0] == 1: # 灰度图
plt.imshow(np.squeeze(img_numpy), cmap='gray')
else: # 彩色图
plt.imshow(np.transpose(img_numpy, (1, 2, 0)))
if title is not None:
plt.title(title)
plt.axis('off') # 隐藏坐标轴
# 获取一个训练数据批次
data_iter = iter(train_loader)
images, labels = next(data_iter) # In Python 3, next(data_iter) is preferred over data_iter.next()
# 创建一个图像网格并显示
# torchvision.utils.make_grid 接受一个图像批次 (批量大小 B, 通道数 C, 高 H, 宽 W)
# 并将其排列成一个网格图像。
img_grid = torchvision.utils.make_grid(images)
plt.figure(figsize=(8, 2)) # 根据需要调整图形大小
imshow(img_grid, title=[str(x.item()) for x in labels]) # 标题为对应标签
plt.show()
# 显示批次中的单张图像:
# plt.figure()
# imshow(images[0], title=f"Label: {labels[0].item()}") # 标题为对应标签
# plt.show()

此代码片段将显示一个 MNIST 数据集图像网格,以及它们对应的标签。可视化数据(Visualizing your data)有助于您确认数据是否正确加载,并理解其特征,然后再将其输入到 ConvNet 中。对于更高级的可视化,例如模型激活(activation)或特征图(feature map),通常会结合 PyTorch 使用 TensorBoard 等工具。