这篇接着 AlexNet 的花卉分类练习,换成 VGG 网络。VGG 的代码重点不在某个复杂算子,而在于用配置列表批量生成卷积和池化层。列表中的数字表示卷积核数量,M 表示最大池化层,所以 VGG11、VGG13、VGG16、VGG19 可以共用同一套构建函数。

代码仍然使用 flower_data/trainflower_data/val,类别数为 5。数据划分和预测预处理可以沿用上一篇,下面只保留 VGG 自己的部分。

定义网络:model.py

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
import torch
import torch.nn as nn


CONFIGS = {
"vgg11": [64, "M", 128, "M", 256, 256, "M", 512, 512, "M", 512, 512, "M"],
"vgg13": [64, 64, "M", 128, 128, "M", 256, 256, "M", 512, 512, "M", 512, 512, "M"],
"vgg16": [64, 64, "M", 128, 128, "M", 256, 256, 256, "M", 512, 512, 512, "M", 512, 512, 512, "M"],
"vgg19": [64, 64, "M", 128, 128, "M", 256, 256, 256, 256, "M", 512, 512, 512, 512, "M", 512, 512, 512, 512, "M"],
}


def make_features(config):
layers = []
in_channels = 3
for value in config:
if value == "M":
layers.append(nn.MaxPool2d(kernel_size=2, stride=2))
continue
layers.extend([
nn.Conv2d(in_channels, value, kernel_size=3, padding=1),
nn.ReLU(inplace=True),
])
in_channels = value
return nn.Sequential(*layers)


class VGG(nn.Module):
def __init__(self, model_name="vgg16", num_classes=1000):
super().__init__()
self.features = make_features(CONFIGS[model_name])
self.avgpool = nn.AdaptiveAvgPool2d((7, 7))
self.classifier = nn.Sequential(
nn.Dropout(p=0.5),
nn.Linear(512 * 7 * 7, 2048),
nn.ReLU(inplace=True),
nn.Dropout(p=0.5),
nn.Linear(2048, 2048),
nn.ReLU(inplace=True),
nn.Linear(2048, num_classes),
)

def forward(self, x):
x = self.avgpool(self.features(x))
x = torch.flatten(x, start_dim=1)
return self.classifier(x)

原版 VGG 的两个隐藏全连接层通常是 4096 维,这里缩到 2048 是为了在普通显卡上训练。num_classes 必须和数据集类别数相同,花卉数据集要设置成 5

训练时替换模型

上一篇 AlexNet 的 train.py 可以继续使用,只需把模型导入和实例化部分替换为:

1
2
3
from model import VGG

model = VGG(model_name="vgg16", num_classes=5).to(device)

VGG 比这里的 AlexNet 更深,建议先使用较小的 batch size,例如 816。训练循环本身没有变化:清空梯度、前向传播、计算损失、反向传播、更新参数,然后在验证集上切换到 eval()

VGG 的配置列表

1
2
3
4
5
for value in CONFIGS["vgg16"]:
if value == "M":
print("MaxPool2d")
else:
print(f"Conv2d(..., out_channels={value}, kernel_size=3)")

这样写的好处是结构和代码分开了。想比较 VGG13 和 VGG16 时,只需要换 model_name,不用手动复制一长串卷积层。

常见问题

DataLoader 迭代器没有 next 方法

在 Python 3 中使用:

1
2
test_data_iter = iter(validate_loader)
test_image, test_label = next(test_data_iter)

显存不足

VGG 的参数量和中间特征图都比较大。降低 batch size 是第一步;如果仍然不够,再缩小全连接层、使用更小的输入,或者换用更轻量的模型。只降低 epoch 不能解决单个 batch 的显存峰值。

这次练习里最值得保留的并不是 VGG 的具体版本,而是配置列表的思路。网络结构变深以后,手写每一层很容易漏参数,数据驱动的构建方式更适合反复试验。