用 PyTorch 复现 VGG
这篇接着 AlexNet 的花卉分类练习,换成 VGG 网络。VGG 的代码重点不在某个复杂算子,而在于用配置列表批量生成卷积和池化层。列表中的数字表示卷积核数量,M 表示最大池化层,所以 VGG11、VGG13、VGG16、VGG19 可以共用同一套构建函数。
代码仍然使用 flower_data/train 和 flower_data/val,类别数为 5。数据划分和预测预处理可以沿用上一篇,下面只保留 VGG 自己的部分。
定义网络:model.py
1 | import torch |
原版 VGG 的两个隐藏全连接层通常是 4096 维,这里缩到 2048 是为了在普通显卡上训练。num_classes 必须和数据集类别数相同,花卉数据集要设置成 5。
训练时替换模型
上一篇 AlexNet 的 train.py 可以继续使用,只需把模型导入和实例化部分替换为:
1 | from model import VGG |
VGG 比这里的 AlexNet 更深,建议先使用较小的 batch size,例如 8 或 16。训练循环本身没有变化:清空梯度、前向传播、计算损失、反向传播、更新参数,然后在验证集上切换到 eval()。
VGG 的配置列表
1 | for value in CONFIGS["vgg16"]: |
这样写的好处是结构和代码分开了。想比较 VGG13 和 VGG16 时,只需要换 model_name,不用手动复制一长串卷积层。
常见问题
DataLoader 迭代器没有 next 方法
在 Python 3 中使用:
1 | test_data_iter = iter(validate_loader) |
显存不足
VGG 的参数量和中间特征图都比较大。降低 batch size 是第一步;如果仍然不够,再缩小全连接层、使用更小的输入,或者换用更轻量的模型。只降低 epoch 不能解决单个 batch 的显存峰值。
这次练习里最值得保留的并不是 VGG 的具体版本,而是配置列表的思路。网络结构变深以后,手写每一层很容易漏参数,数据驱动的构建方式更适合反复试验。
本博客所有文章除特别声明外,均采用 CC BY-NC-SA 4.0 许可协议。转载请注明来源 Study-HYK!




