Skip to content

第 10 章 卷积神经网络

学习目标

  • 掌握图像张量的形状约定 (N, C, H, W)
  • 掌握 nn.Conv2d 的常用参数与输出尺寸计算
  • 掌握 nn.MaxPool2d 池化的作用
  • 能搭一个小型 CNN 并逐层推算形状
  • 能用合成图片数据训练一个图像分类器

10.1 图像在 PyTorch 里是什么形状

一张灰度图(高 28、宽 28)在 PyTorch 里是形状 (1, 28, 28) 的张量:通道 × 高 × 宽(第 9 章 ToTensor 的约定)。一批 8 张就是 (8, 1, 28, 28),即 (N, C, H, W):批大小、通道数、高、宽。彩色图则是 (N, 3, H, W)

10.2 卷积层 nn.Conv2d

卷积层(convolutional layer)用一个可学习的「小窗口」(卷积核,如 3×3)扫过整张图,在每个位置做加权求和,提取局部特征。nn.Conv2d 的常用参数:

  • in_channels:输入通道数(灰度图 1,彩色图 3);
  • out_channels:输出通道数,即用多少个卷积核、提取多少种特征;
  • kernel_size:卷积核大小,如 3 表示 3×3;
  • padding:边缘补 0 的圈数,padding=1 保持高宽不变;
  • stride:滑动的步长,默认 1。
python
import torch
import torch.nn as nn

torch.manual_seed(0)
conv = nn.Conv2d(in_channels=1, out_channels=4, kernel_size=3, padding=1)
x = torch.randn(8, 1, 28, 28)     # 8 张 28×28 灰度图
out = conv(x)
print("带 padding 输出:", out.shape)

conv_nopad = nn.Conv2d(1, 4, kernel_size=3)
print("无 padding 输出:", conv_nopad(x).shape)

输出:

带 padding 输出: torch.Size([8, 4, 28, 28])
无 padding 输出: torch.Size([8, 4, 26, 26])

输出高宽的计算公式:

输出高 = (输入高 - kernel_size + 2 × padding) / stride + 1

28×28 的图:3×3 核、padding=1(28 - 3 + 2) / 1 + 1 = 28,高宽不变;不 padding → (28 - 3) / 1 + 1 = 26。输出通道数 4 与图大小无关,只是特征数。

10.3 池化:缩小尺寸、保留特征

池化(pooling)在局部取最大值(或平均值),把图缩小,同时保留主要特征、减少计算量。nn.MaxPool2d(2) 表示 2×2 窗口、步长 2,高宽各减半:

python
import torch
import torch.nn as nn

pool = nn.MaxPool2d(2)
feature = torch.randn(8, 4, 28, 28)
print(pool(feature).shape)

输出:

torch.Size([8, 4, 14, 14])

28×28 → 14×14。典型的卷积网络模式是「卷积提取特征 → 激活 → 池化缩小」,重复几轮,特征图越来越小、通道越来越多。

10.4 一个小型 CNN:逐层推算形状

搭一个两层卷积网络,输入 28×28 灰度图,输出 10 类分数:

python
import torch
import torch.nn as nn

class TinyCNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.features = nn.Sequential(
            nn.Conv2d(1, 4, kernel_size=3, padding=1),   # (B, 1, 28, 28) → (B, 4, 28, 28)
            nn.ReLU(),
            nn.MaxPool2d(2),                              # → (B, 4, 14, 14)
            nn.Conv2d(4, 8, kernel_size=3, padding=1),   # → (B, 8, 14, 14)
            nn.ReLU(),
            nn.MaxPool2d(2),                              # → (B, 8, 7, 7)
        )
        self.classifier = nn.Linear(8 * 7 * 7, 10)

    def forward(self, x):
        h = self.features(x)
        h = h.flatten(1)          # (B, 8, 7, 7) → (B, 392)
        return self.classifier(h)

net = TinyCNN()
x = torch.randn(8, 1, 28, 28)
print(net(x).shape)

输出:

torch.Size([8, 10])

形状流:每张 28×28 图 → 池化两次减半到 7×7,8 个通道 → 展平成 8×7×7 = 392 个特征 → 全连接层输出 10 个分数。全连接层的输入维度必须与展平后的特征数一致,这是新手最常算错的地方。

10.5 完整示例:训练一个图像分类器

合成数据训练 TinyCNN:类别 0 的图是「暗图」(像素均值低),类别 1 是「亮图」(像素均值高),共 200 张。模型要学的是「看像素亮暗判断类别」:

python
import torch
import torch.nn as nn

torch.manual_seed(1)
n = 200
x0 = torch.randn(n // 2, 1, 28, 28) * 0.12 + 0.30   # 暗图
x1 = torch.randn(n // 2, 1, 28, 28) * 0.12 + 0.70   # 亮图
x = torch.cat([x0, x1])
y = torch.cat([torch.zeros(n // 2), torch.ones(n // 2)]).long()

model = TinyCNN()
model.classifier = nn.Linear(8 * 7 * 7, 2)   # 两类
loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.005)

for epoch in range(6):
    logits = model(x)
    loss = loss_fn(logits, y)
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()
    acc = (logits.argmax(1) == y).float().mean().item()
    print(f"epoch {epoch}: loss = {loss.item():.4f}, acc = {acc:.2f}")

输出:

epoch 0: loss = 0.6892, acc = 0.50
epoch 1: loss = 0.6586, acc = 0.56
epoch 2: loss = 0.6202, acc = 1.00
epoch 3: loss = 0.5738, acc = 1.00
epoch 4: loss = 0.5222, acc = 1.00
epoch 5: loss = 0.4639, acc = 1.00

第 2 轮后准确率到 100%:两类图差异明显,卷积网络很快学会「高均值 → 类别 1」。注意 logits.argmax(1) == y 是算准确率的标准写法:取每个样本分数最高的类别,与真实标签比较。

epoch 0 准确率恰好 50% 不是巧合:两类样本各占一半,瞎猜就是这个数。

动手实践

  1. 分别用 kernel_size=5padding=2kernel_size=3padding=0 对 28×28 输入做卷积,用公式算出输出高宽并验证。
  2. 把 TinyCNN 的第二个池化删掉,运行前向传播,记录报错,解释全连接层输入维度错在哪里。
  3. 把 10.5 的暗图均值改成 0.45、亮图均值改成 0.55(差异变小),重训,观察准确率变化。
  4. 把 10.5 的 Adam 换成 SGD(lr=0.01),观察前 6 轮的准确率,感受优化器差异。
  5. 打印 model.features 的结构,并手算从输入到展平前的每层形状,与代码注释核对。

常见错误

错误写法现象原因
输入 (B, H, W) 忘了通道维RuntimeError: Expected 3D (unbatched) or 4D (batched) input to conv2d卷积要求 4 维 (N, C, H, W),用 unsqueeze(1) 补通道维
全连接层输入维度算错RuntimeError: mat1 and mat2 shapes cannot be multiplied展平后特征数 = 通道 × 高 × 宽,用 flatten(1) 后打印形状核对
分类输出层维度与类别数不符损失形状错乱或报错CrossEntropyLoss 要求输出维度 = 类别数
CrossEntropyLoss 的标签不是整数RuntimeError: Expected target size [...] 或 dtype 错误标签用 torch.long(整数),如 .long()
池化后忘了更新全连接维度前向报形状错误池化改变高宽,展平特征数随之变化
H×W×C 的图片直接喂 Conv2d通道维在最后,形状不匹配图像张量必须 C×H×W,用 ToTensor 转换

章末练习

基础

  1. 写出 nn.Conv2d(1, 8, kernel_size=3, padding=1)(16, 1, 32, 32) 输入的输出形状,并说明每个数字的含义。
  2. nn.MaxPool2d(2) 处理 (16, 8, 32, 32),写出输出形状。
  3. 搭一个 CNN:输入 (4, 1, 28, 28),卷积 3×3 padding=1 → ReLU → 池化 → 展平 → Linear(392, 5),打印最终输出形状。

提高

  1. 解释 padding=1 的作用,并用输出尺寸公式证明「3×3 核 + padding=1 保持高宽不变」。
  2. 把 10.5 的数据噪声从 0.12 改成 0.35,重训,记录准确率曲线,解释为什么更难学了。

挑战

  1. 设计一个 3 层卷积的 CNN(每层卷积 + ReLU + 池化),输入 (8, 1, 64, 64),手算每层形状,再写代码验证全连接层维度。
  2. 不运行代码,预测:输入 (2, 3, 16, 16)nn.Conv2d(3, 6, kernel_size=3, padding=1)、再 nn.MaxPool2d(2) 后的形状;运行验证并解释。

章末自测

每题选择一个最佳答案。本书不附答案:完成后交由老师或 AI 老师批改讲解。

  1. 一批 8 张 28×28 灰度图在 PyTorch 里的形状是?
    • A. (8, 28, 28)
    • B. (8, 1, 28, 28)
    • C. (28, 28, 8)
    • D. (1, 8, 28, 28)
  2. nn.Conv2d(1, 4, kernel_size=3, padding=1)(8, 1, 28, 28) 输入,输出形状是?
    • A. (8, 4, 28, 28)
    • B. (8, 4, 26, 26)
    • C. (8, 1, 28, 28)
    • D. (4, 8, 28, 28)
  3. 28×28 的输入、3×3 核、无 padding,输出高宽是?
    • A. 28
    • B. 27
    • C. 26
    • D. 25
  4. nn.MaxPool2d(2)(8, 4, 28, 28) 的输出形状是?
    • A. (8, 4, 28, 28)
    • B. (8, 4, 14, 14)
    • C. (8, 8, 14, 14)
    • D. (4, 8, 14, 14)
  5. flatten(1)(8, 8, 7, 7) 的结果形状是?
    • A. (8, 392)
    • B. (392, 8)
    • C. (8, 8, 49)
    • D. (3136,)
  6. 卷积层 in_channels 表示?
    • A. 输出特征数
    • B. 输入通道数
    • C. 卷积核大小
    • D. 批大小
  7. 池化的作用是?
    • A. 增加通道数
    • B. 缩小高宽、保留主要特征
    • C. 把图片变亮
    • D. 增加参数量
  8. 多分类标签送给 CrossEntropyLoss 前,应该转成?
    • A. float
    • B. long(整数)
    • C. bool
    • D. 不用转
  9. 计算分类准确率的标准写法是?
    • A. (logits == y).mean()
    • B. (logits.argmax(1) == y).float().mean()
    • C. (logits.max(1) == y).mean()
    • D. loss.item()
  10. 卷积网络接全连接层前,通常先?
    • A. flatten 展平特征图
    • B. 转置
    • C. 归一化到整数
    • D. 什么都不做