从零构建CNN:用PyTorch拆解卷积、激活与池化的核心机制
当你第一次看到卷积神经网络(CNN)的架构图时,那些堆叠的Conv、ReLU和Pooling层是否让你感到困惑?这些看似简单的组件如何协同工作,从原始像素中提取出高级特征?本文将用PyTorch代码逐层拆解这个"黑箱",让你真正理解每个操作背后的数学原理和视觉意义。
1. 环境准备与数据加载
在开始构建CNN之前,我们需要准备好Python环境和MNIST数据集。这个经典的手写数字数据集包含60,000张28x28的灰度图像,非常适合用来理解基础视觉概念。
import torch
import torch.nn as nn
import torchvision
from torchvision import transforms
from torch.utils.data import DataLoader
import matplotlib.pyplot as plt
# 数据预处理管道
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,))
])
# 加载MNIST数据集
train_set = torchvision.datasets.MNIST(
root='./data',
train=True,
download=True,
transform=transform
)
test_set = torchvision.datasets.MNIST(
root='./data',
train=False,
download=True,
transform=transform
)
# 创建数据加载器
train_loader = DataLoader(train_set, batch_size=64, shuffle=True)
test_loader = DataLoader(test_set, batch_size=1000, shuffle=False)
提示:Normalize变换将像素值从[0,1]范围缩放到[-1,1],这对神经网络的训练稳定性很重要。
让我们可视化一些样本,了解我们要处理的数据:
def show_images(images, labels, nrows=3, ncols=3):
fig, axes = plt.subplots(nrows, ncols, figsize=(8,8))
for i, ax in enumerate(axes.flat):
ax.imshow(images[i].squeeze(), cmap='gray')
ax.set_title(f'Label: {labels[i]}')
ax.axis('off')
plt.tight_layout()
plt.show()
# 获取一个batch的数据
dataiter = iter(train_loader)
images, labels = next(dataiter)
show_images(images, labels)
这些手写数字图像虽然简单,但已经包含了计算机视觉的核心挑战:同一数字的不同书写风格、位置变化和轻微形变等。
2. 卷积层:特征提取的艺术
卷积层是CNN的核心组件,它通过一组可学习的滤波器在图像上滑动,提取局部特征。让我们深入理解这个过程的每个细节。
2.1 卷积操作原理
一个卷积滤波器本质上是一个小矩阵(通常3x3或5x5),它在输入图像上滑动,计算局部区域的点积。这个过程可以理解为在图像中寻找特定模式(如边缘、纹理等)。
在PyTorch中,我们使用nn.Conv2d创建卷积层:
# 创建一个卷积层示例
conv_layer = nn.Conv2d(
in_channels=1, # 输入通道数(灰度图为1)
out_channels=16, # 输出通道数/滤波器数量
kernel_size=3, # 滤波器大小
stride=1, # 滑动步长
padding=1 # 边缘填充
)
# 查看滤波器权重
print(conv_layer.weight.shape) # 输出:torch.Size([16, 1, 3, 3])
理解输出尺寸的计算至关重要。对于一个输入尺寸为(W,H)的图像,输出尺寸由以下公式决定:
输出宽度 = floor((W - F + 2P)/S + 1)
输出高度 = floor((H - F + 2P)/S + 1)
其中:
- F:滤波器大小
- P:填充像素数
- S:步长
2.2 可视化卷积效果
让我们实际应用一个卷积层并观察它对图像的影响:
def apply_conv_and_visualize(image, conv_layer):
# 应用卷积
with torch.no_grad():
output = conv_layer(image.unsqueeze(0)) # 添加batch维度
# 可视化
fig, axes = plt.subplots(4, 4, figsize=(10,10))
for i, ax in enumerate(axes.flat):
if i < output.shape[1]: # 只显示前16个特征图
ax.imshow(output[0,i].detach().numpy(), cmap='gray')
ax.set_title(f'Filter {i+1}')
ax.axis('off')
plt.tight_layout()
plt.show()
# 选择一张图像并应用卷积
sample_image = images[0]
apply_conv_and_visualize(sample_image


464

被折叠的 条评论
为什么被折叠?



