☰
PyTorch 轻量级图像分类实战:用 ShuffleNetV2 训练花分类数据集并完成推理(deep-learning-for-image-processing)
2026/10/1 6:59:05 网站建设 项目流程
  • 示例工程

【免费下载链接】deep-learning-for-image-processing

deep learning for image processing including classification and object-detection etc.

项目地址:https://gitcode.com/gh_mirrors/de/deep-learning-for-image-processing
点击查看免费下载

导读

本篇技术指南围绕开源仓库 deep-learning-for-image-processing 中Test7_shufflenet模块展开,完整讲解如何基于 PyTorch 使用轻量级网络 ShuffleNetV2 在花分类数据集上进行训练、验证与单张图片预测。读完本文,你将掌握:数据集下载与划分、预训练权重下载与载入、train.py训练脚本的全部参数含义、predict.py推理脚本的调用流程,以及如何将这套流程迁移到自己的自定义数据集上。文章同时结合仓库源码,深入剖析 ShuffleNetV2 的channel shuffle算子与InvertedResidual模块的底层实现,让实操与原理互为印证。

模块与文件总览

pytorch_classification/Test7_shufflenet/目录下共包含 6 个文件,分工如下(与 pytorch_classification/README.md 的描述一致):

文件作用
model.pyShuffleNetV2 模型定义,含channel_shuffle、InvertedResidual和 4 个不同宽度系数的模型构建函数
train.py训练入口脚本,负责数据加载、预训练权重载入、优化器与学习率调度、训练与验证
predict.py单张图片预测脚本,输出各类别概率与最高概率类别
my_dataset.py自定义Dataset封装与collate_fn实现
utils.py数据划分、单轮训练、验证评估等工具函数
class_indices.json训练过程中自动生成的类别索引映射文件(当前为花分类 5 类)

一、数据准备:下载并整理花分类数据集

本模块默认使用 TensorFlow 官方发布的flower_photos花分类数据集(共 3670 张样本,5 个类别:daisy、dandelion、roses、sunflowers、tulips)。数据集下载地址、备用网盘地址均已在 Test7_shufflenet/README.md 中给出,解压后得到flower_photos文件夹。

推荐的数据组织方式有两种:

方式一:使用仓库提供的划分脚本。按 data_set/README.md 的说明,在data_set目录下创建flower_data文件夹,将下载解压后的flower_photos放入其中,然后执行 split_data.py:

├── flower_data ├── flower_photos(解压的数据集文件夹,3670个样本) ├── train(生成的训练集,3306个样本) └── val(生成的验证集,364个样本)

split_data.py的核心逻辑是:按split_rate = 0.1的比例在每个类别文件夹内随机采样样本复制进val,其余复制进train,并使用random.seed(0)保证划分结果可复现(split_data.py)。

方式二:交给训练脚本自行划分。train.py依赖的 utils.py 中的read_split_data函数会按val_rate=0.2的默认比例在内存中随机划分训练集/验证集,并自动生成class_indices.json(utils.py)。它遍历数据集根目录下每个子文件夹,将文件夹名作为类别名、按字母序编号,支持的图片后缀为[".jpg", ".JPG", ".png", ".PNG"](utils.py)。因此只要数据集按“一个类别一个文件夹”的结构摆放即可直接使用。

二、下载预训练权重

ShuffleNetV2 的官方预训练权重(基于 ImageNet 1000 类训练)下载地址,均已注释在 model.py 中每个模型构建函数的 docstring 里:

  • shufflenet_v2_x0_5:权重文件shufflenetv2_x0.5-f707e7126e.pth
  • shufflenet_v2_x1_0:权重文件shufflenetv2_x1-5666bf0f80.pth
  • shufflenet_v2_x1_5:权重文件shufflenetv2_x1_5-3c479a10.pth
  • shufflenet_v2_x2_0:权重文件shufflenetv2_x2_0-8be3c8ee.pth

建议将下载好的权重文件放在Test7_shufflenet目录下,与train.py保持同级,便于默认路径直接命中。训练脚本默认使用的是shufflenet_v2_x1_0对应的权重(train.py),如果选用其他宽度系数版本,请同步修改脚本中导入的模型函数。

三、配置训练脚本train.py

train.py使用argparse定义的全部命令行参数如下(train.py):

参数默认值含义与建议
--num_classes5分类类别数,花数据集为 5,自定义数据集需改成自己的类别数
--epochs30训练轮数
--batch-size16批大小
--lr0.01初始学习率
--lrf0.1学习率衰减的下限比例(配合余弦退火调度)
--data-path/data/flower_photos解压后数据集根目录的绝对路径,必须修改
--weights./shufflenetv2_x1.pth预训练权重路径
--freeze-layersFalse是否冻结除全连接层外的全部权重(迁移学习常用)
--devicecuda:0训练设备,支持cuda:0、0,1、cpu

启动训练的最小命令:

python train.py --data-path /你的路径/flower_photos --weights ./shufflenetv2_x1.pth

3.1 训练前处理与数据加载

main函数首先根据torch.cuda.is_available()自动选择设备(train.py),随后定义训练/验证两套数据增强:

  • 训练:RandomResizedCrop(224)+RandomHorizontalFlip()+ToTensor()+ 按 ImageNet 均值方差归一化
  • 验证:Resize(256)+CenterCrop(224)+ToTensor()+ 同样的归一化(train.py)

DataLoader的num_workers取min(os.cpu_count(), batch_size, 8)三者最小值,并开启pin_memory(train.py)。数据读取由 my_dataset.py 中的MyDataSet完成:它会校验每张图片必须为 RGB 模式,非 RGB 图片直接抛出ValueError(my_dataset.py);collate_fn通过torch.stack与torch.as_tensor将 batch 打包为张量(my_dataset.py)。

3.2 预训练权重载入与层冻结

权重载入时做了一个关键的过滤处理:只加载“参数元素个数与模型对应层一致”的键,再以strict=False方式载入(train.py)。这样即使预训练权重里 1000 类的fc层与你设置的num_classes不一致,也能顺利加载其余层的特征提取参数,实现“替换分类头做迁移学习”。

若设置--freeze-layers True,则除名字含fc的全连接层外,其余参数全部requires_grad_(False),此时训练只更新最后的分类层(train.py),适合小数据集快速微调。

3.3 优化器、余弦退火调度与训练循环

优化器使用带动量的 SGD:SGD(lr, momentum=0.9, weight_decay=4E-5)(train.py)。学习率调度采用余弦退火(Cosine Annealing)LambdaLR,lr_lambda的公式为:

((1 + cos(epoch * π / epochs)) / 2) * (1 - lrf) + lrf

即学习率从初始lr平滑衰减到lr * lrf(train.py)。

每个 epoch 依次执行:

  1. train_one_epoch:内部使用CrossEntropyLoss计算损失、反向传播并更新参数,同时用滑动平均方式记录mean_loss;若 loss 出现非有限值会打印警告并提前终止(utils.py);
  2. scheduler.step()更新学习率;
  3. evaluate:在验证集上统计argmax预测与标签一致的样本占比得到 accuracy(utils.py);
  4. 将 loss、accuracy、learning_rate 写入 TensorBoard(可用tensorboard --logdir=runs查看,地址http://localhost:6006/),并把模型权重保存到./weights/model-{epoch}.pth(train.py)。

训练过程中会自动生成class_indices.json(由read_split_data在数据加载阶段写入),它记录“数字索引 → 类别名”的映射,供预测脚本使用。

四、配置预测脚本predict.py

训练完成后,按以下步骤修改 predict.py:

  1. 导入与训练一致的模型:脚本默认from model import shufflenet_v2_x1_0,num_classes必须与训练时一致(默认 5)(predict.py);
  2. 设置权重路径:将model_weight_path改为训练好的权重文件路径(默认保存在weights文件夹下,例如./weights/model-29.pth),载入后调用model.eval()(predict.py);
  3. 设置预测图片路径:将img_path改成待预测图片的绝对路径(predict.py)。

预测时的数据预处理必须与验证集保持一致(Resize(256)+CenterCrop(224)+ 相同归一化参数),否则结果会有偏差。图片经过unsqueeze增加 batch 维后前向传播,在no_grad上下文下取softmax得到各类别概率,再argmax得到预测类别(predict.py),最终在终端打印每个类别的概率,并用 matplotlib 展示图片与标题。

运行方式:

python predict.py

五、迁移到自定义数据集

要把这套流程用于自己的数据,只需三步:

  1. 目录结构:按照花分类数据集的组织方式,一个类别对应一个文件夹,例如:
my_dataset/ ├── class1/ # 该类别的所有图片 ├── class2/ └── class3/
  1. 修改类别数:将train.py与predict.py中的num_classes改成你的类别总数;
  2. 修改路径:分别设置好--data-path与--weights(训练时),以及model_weight_path与img_path(预测时)。

注意:若自定义数据集类别数与 1000 不一致,预训练权重中fc层的参数会被过滤掉不载入,这属于预期行为;训练脚本会自动生成新的class_indices.json供预测脚本读取。

六、ShuffleNetV2 核心原理:从源码看轻量设计

为了让训练更有的放矢,这里结合 model.py 剖析 ShuffleNetV2 的两个关键设计。

6.1 channel shuffle:通道重排算子

channel_shuffle函数将特征图先 reshape 成[batch_size, groups, channels_per_group, height, width],交换中间两维后contiguous()再 flatten 回[batch_size, -1, height, width](model.py)。它的作用是打破分组卷积中“组与组之间信息隔离”的局限,让不同分组的通道信息充分交互,是 ShuffleNet 系列在分组卷积基础上保证精度的关键操作。

6.2 InvertedResidual:stride 1 与 stride 2 的双分支设计

InvertedResidual模块根据stride分为两种结构(model.py):

  • stride = 1:输入沿通道维chunk(2, dim=1)切成两半,一半直接恒等映射(shortcut),另一半经过“1×1 卷积 → 3×3 深度卷积 → 1×1 卷积”的轻量分支,最后两半拼接并做 channel shuffle(model.py)。该分支完全无相加操作,计算量更低;
  • stride = 2:无恒等分支,两条分支都对输入做处理:分支1 是“3×3 深度卷积(下采样)+ 1×1 卷积”,分支2 是“1×1 卷积 + 3×3 深度卷积(下采样)+ 1×1 卷积”,最后拼接并 channel shuffle(model.py)。

其中“深度卷积”通过nn.Conv2d(groups=input_c)实现(model.py),每组通道只在自己的通道内做卷积,是大幅降低计算量的核心手段。模块构造函数中还包含两个重要断言:output_c必须为偶数(用于等分),且 stride=1 时输入通道数必须等于branch_features << 1(model.py),保证通道拼接维度匹配。

6.3 四个宽度系数版本

整个网络由conv1(3×3,stride 2)+maxpool+stage2/3/4三个堆叠阶段 +conv5(1×1)+ 全局平均池化 +fc组成(model.py)。三个阶段的重复次数固定为[4, 8, 4],区别仅在每阶段的输出通道数(model.py):

模型函数各阶段输出通道[24, s2, s3, s4, 1024/2048]特点
shufflenet_v2_x0_5[24, 48, 96, 192, 1024]最轻量,适合移动端/嵌入式场景
shufflenet_v2_x1_0[24, 116, 232, 464, 1024]默认版本,速度与精度平衡
shufflenet_v2_x1_5[24, 176, 352, 704, 1024]更宽的 1.5 倍通道
shufflenet_v2_x2_0[24, 244, 488, 976, 2048]最宽的 2.0 倍通道,精度更高但计算量更大

在算力受限的环境(如只有 CPU)下,x0_5或x1_0都能以较低资源开销完成花分类这类轻量任务的训练与推理;这也是本模块作为“轻量级图像分类”示例的意义所在。

七、常见问题与排查要点

  • --data-path报路径不存在:read_split_data会断言数据集根目录存在(utils.py),请确认传入的是解压后flower_photos文件夹的绝对路径,且其下每个子文件夹直接存放图片。
  • 权重文件找不到:train.py在args.weights指向的文件不存在时会抛出FileNotFoundError(train.py),请核对权重文件名与下载地址注释中的名称一致。
  • 预测时class_indices.json缺失:该文件在训练阶段自动生成,若直接跑预测需确保当前目录存在该文件(predict.py),仓库已随附一份花分类的class_indices.json可直接使用。
  • 自定义数据集非 RGB 图片:MyDataSet会拒绝非 RGB 模式图片并报错(my_dataset.py),请预先统一图片格式。
  • 训练只更新分类头:小数据量场景可加--freeze-layers True先冻结骨干,待收敛后再解冻整体微调;追求更好精度则保持默认全量训练。

结语

至此,从数据下载、脚本参数配置、预训练权重载入,到训练、验证、预测的完整闭环已经打通,同时通过对channel shuffle与InvertedResidual源码的剖析,你也理解了 ShuffleNetV2 在保持轻量的同时维持精度的方法论。以本模块为模板,替换 model.py 中的模型函数即可无缝切换到仓库中的其他分类网络(如 ResNet、MobileNet、EfficientNet 等),整个训练/预测框架无需改动,可直接复用于你自己的图像分类任务。

  • 示例工程

【免费下载链接】deep-learning-for-image-processing

deep learning for image processing including classification and object-detection etc.

项目地址:https://gitcode.com/gh_mirrors/de/deep-learning-for-image-processing
点击查看免费下载

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询