用 PyTorch 复现 AlexNet
这篇把 AlexNet 用到五分类花卉数据集上。流程和前面的 LeNet Demo 相同,只是多了数据集划分,模型也更深:准备 ImageFolder 目录、定义网络、训练并保存最优权重,最后预测单张图片。
我在全连接层使用 2048 个节点,而不是原论文的 4096 个节点。这样能减少本地显存占用,所以它是适合练习的 AlexNet 变体,不是逐参数复刻论文。
准备数据集
数据集下载地址:http://download.tensorflow.org/example_images/flower_photos.tgz
原始目录中每个子目录代表一个类别:
1 | flower_data/ |
下面的脚本把 10% 图片随机划到验证集。它会重新创建 train 和 val 目录,运行前不要在这两个目录里放其他文件。
1 | import random |
定义 AlexNet:model.py
1 | import torch |
AdaptiveAvgPool2d((6, 6)) 明确了全连接层需要的输入尺寸。这样即使前面的输入尺寸略有变化,只要特征图不小于目标尺寸,分类器仍能收到固定长度的向量。
训练:train.py
1 | import json |
ImageFolder 会按目录名排序并生成类别编号,因此预测脚本不能自己手写另一套编号。训练时把映射保存成 JSON,后面直接读取最稳妥。
预测:predict.py
1 | import json |
我遇到过的问题
DataLoader 迭代器没有 next 方法
旧教程里经常出现:
1 | images, labels = data_iter.next() |
Python 3 中直接使用内置函数:
1 | images, labels = next(data_iter) |
CUDA 显存不足
先减小 batch size,例如从 32 改成 16 或 8。单个 batch 仍然放不下时,再缩小全连接层、输入尺寸或模型通道数。减少 epoch 只会缩短总训练时间,不能解决一次前向传播造成的显存不足。
这次实践让我真正理解了 AlexNet 结构图之外的部分:数据目录、类别映射、训练和预测预处理是否一致,往往比某一层参数更容易让结果出错。
本博客所有文章除特别声明外,均采用 CC BY-NC-SA 4.0 许可协议。转载请注明来源 Study-HYK!




