PyTorch - 卷积网络可视化
PyTorch - ConvNets 图像数据可视化
Section titled “PyTorch - ConvNets 图像数据可视化”理解和可视化您的输入数据是任何深度学习项目(尤其是在处理图像任务的卷积神经网络,即 ConvNets 或 CNNs 时)至关重要的第一步。本章演示了如何使用 torchvision 加载一个标准数据集(如 MNIST),并使用 matplotlib 可视化其中的一些图像。
步骤 1:导入所需模块
Section titled “步骤 1:导入所需模块”我们将需要 torch,用于数据集和图像转换的 torchvision,以及用于绘图的 matplotlib.pyplot。
import torchimport torchvisionimport torchvision.transforms as transformsimport matplotlib.pyplot as pltimport numpy as np # For converting tensors to numpy arrays for plotting
# 将张量(Tensor)转换为 NumPy 数组(array)以进行绘图步骤 2:加载数据集
Section titled “步骤 2:加载数据集”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)。
步骤 3:可视化部分图像
Section titled “步骤 3:可视化部分图像”现在,让我们从 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 等工具。