ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyTorch 轻量级图像分类实战:用 ShuffleNetV2 训练花分类数据集并完成推理(deep-learning-for-image-processing)

PyTorch 轻量级图像分类实战:用 ShuffleNetV2 训练花分类数据集并完成推理(deep-learning-for-image-processing) 示例工程【免费下载链接】deep-learning-for-image-processingdeep 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_rate0.2的默认比例在内存中随机划分训练集/验证集并自动生成class_indices.jsonutils.py。它遍历数据集根目录下每个子文件夹将文件夹名作为类别名、按字母序编号支持的图片后缀为[.jpg, .JPG, .png, .PNG]utils.py。因此只要数据集按“一个类别一个文件夹”的结构摆放即可直接使用。二、下载预训练权重ShuffleNetV2 的官方预训练权重基于 ImageNet 1000 类训练下载地址均已注释在 model.py 中每个模型构建函数的 docstring 里shufflenet_v2_x0_5权重文件shufflenetv2_x0.5-f707e7126e.pthshufflenet_v2_x1_0权重文件shufflenetv2_x1-5666bf0f80.pthshufflenet_v2_x1_5权重文件shufflenetv2_x1_5-3c479a10.pthshufflenet_v2_x2_0权重文件shufflenetv2_x2_0-8be3c8ee.pth建议将下载好的权重文件放在Test7_shufflenet目录下与train.py保持同级便于默认路径直接命中。训练脚本默认使用的是shufflenet_v2_x1_0对应的权重train.py如果选用其他宽度系数版本请同步修改脚本中导入的模型函数。三、配置训练脚本train.pytrain.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.pth3.1 训练前处理与数据加载main函数首先根据torch.cuda.is_available()自动选择设备train.py随后定义训练/验证两套数据增强训练RandomResizedCrop(224)RandomHorizontalFlip()ToTensor() 按 ImageNet 均值方差归一化验证Resize(256)CenterCrop(224)ToTensor() 同样的归一化train.pyDataLoader的num_workers取min(os.cpu_count(), batch_size, 8)三者最小值并开启pin_memorytrain.py。数据读取由 my_dataset.py 中的MyDataSet完成它会校验每张图片必须为 RGB 模式非 RGB 图片直接抛出ValueErrormy_dataset.pycollate_fn通过torch.stack与torch.as_tensor将 batch 打包为张量my_dataset.py。3.2 预训练权重载入与层冻结权重载入时做了一个关键的过滤处理只加载“参数元素个数与模型对应层一致”的键再以strictFalse方式载入train.py。这样即使预训练权重里 1000 类的fc层与你设置的num_classes不一致也能顺利加载其余层的特征提取参数实现“替换分类头做迁移学习”。若设置--freeze-layers True则除名字含fc的全连接层外其余参数全部requires_grad_(False)此时训练只更新最后的分类层train.py适合小数据集快速微调。3.3 优化器、余弦退火调度与训练循环优化器使用带动量的 SGDSGD(lr, momentum0.9, weight_decay4E-5)train.py。学习率调度采用余弦退火Cosine AnnealingLambdaLRlr_lambda的公式为((1 cos(epoch * π / epochs)) / 2) * (1 - lrf) lrf即学习率从初始lr平滑衰减到lr * lrftrain.py。每个 epoch 依次执行train_one_epoch内部使用CrossEntropyLoss计算损失、反向传播并更新参数同时用滑动平均方式记录mean_loss若 loss 出现非有限值会打印警告并提前终止utils.pyscheduler.step()更新学习率evaluate在验证集上统计argmax预测与标签一致的样本占比得到 accuracyutils.py将 loss、accuracy、learning_rate 写入 TensorBoard可用tensorboard --logdirruns查看地址http://localhost:6006/并把模型权重保存到./weights/model-{epoch}.pthtrain.py。训练过程中会自动生成class_indices.json由read_split_data在数据加载阶段写入它记录“数字索引 → 类别名”的映射供预测脚本使用。四、配置预测脚本predict.py训练完成后按以下步骤修改 predict.py导入与训练一致的模型脚本默认from model import shufflenet_v2_x1_0num_classes必须与训练时一致默认 5predict.py设置权重路径将model_weight_path改为训练好的权重文件路径默认保存在weights文件夹下例如./weights/model-29.pth载入后调用model.eval()predict.py设置预测图片路径将img_path改成待预测图片的绝对路径predict.py。预测时的数据预处理必须与验证集保持一致Resize(256)CenterCrop(224) 相同归一化参数否则结果会有偏差。图片经过unsqueeze增加 batch 维后前向传播在no_grad上下文下取softmax得到各类别概率再argmax得到预测类别predict.py最终在终端打印每个类别的概率并用 matplotlib 展示图片与标题。运行方式python predict.py五、迁移到自定义数据集要把这套流程用于自己的数据只需三步目录结构按照花分类数据集的组织方式一个类别对应一个文件夹例如my_dataset/ ├── class1/ # 该类别的所有图片 ├── class2/ └── class3/修改类别数将train.py与predict.py中的num_classes改成你的类别总数修改路径分别设置好--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 InvertedResidualstride 1 与 stride 2 的双分支设计InvertedResidual模块根据stride分为两种结构model.pystride 1输入沿通道维chunk(2, dim1)切成两半一半直接恒等映射shortcut另一半经过“1×1 卷积 → 3×3 深度卷积 → 1×1 卷积”的轻量分支最后两半拼接并做 channel shufflemodel.py。该分支完全无相加操作计算量更低stride 2无恒等分支两条分支都对输入做处理分支1 是“3×3 深度卷积下采样 1×1 卷积”分支2 是“1×1 卷积 3×3 深度卷积下采样 1×1 卷积”最后拼接并 channel shufflemodel.py。其中“深度卷积”通过nn.Conv2d(groupsinput_c)实现model.py每组通道只在自己的通道内做卷积是大幅降低计算量的核心手段。模块构造函数中还包含两个重要断言output_c必须为偶数用于等分且 stride1 时输入通道数必须等于branch_features 1model.py保证通道拼接维度匹配。6.3 四个宽度系数版本整个网络由conv13×3stride 2maxpoolstage2/3/4三个堆叠阶段 conv51×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指向的文件不存在时会抛出FileNotFoundErrortrain.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-processingdeep learning for image processing including classification and object-detection etc.项目地址https://gitcode.com/gh_mirrors/de/deep-learning-for-image-processing点击查看免费下载相关推荐基于 PyTorch 训练与预测 RegNet 图像分类网络花分类数据集实战指南deep-learning-for-image-processing 项目基于 PyTorch 训练与预测 RegNet 图像分类网络花分类数据集实战指南deep learning for image processing 项目示例工程EfficientNetV2 图像分类实战基于 PyTorch 的花卉识别训练与推理全流程deep-learning-for-image-processingEfficientNetV2 图像分类实战基于 PyTorch 的花卉识别训练与推理全流程deep learning for image processin示例工程deep-learning-for-image-processing 实战使用 PyTorch 训练与部署 Swin Transformer 图像分类模型deep learning for image processing 实战使用 PyTorch 训练与部署 Swin Transformer 图像分类模型 本示例工程创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表