新闻详情

新闻详情

首页 / 资讯中心 / 详情

deep-learning-for-image-processing 实战:使用 PyTorch 训练与部署 Swin Transformer 图像分类模型

发布时间:2026/9/30 6:39:24来源:尧图网络
deep-learning-for-image-processing 实战:使用 PyTorch 训练与部署 Swin Transformer 图像分类模型
示例工程【免费下载链接】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 中pytorch_classification/swin_transformer模块的配套文档与源码完整讲解从数据集准备、预训练权重加载、模型微调训练、单张图片预测到自定义数据集的整套 Swin Transformer 分类流程。读完本文你将能够直接复现花分类数据集上的 Swin-Tiny 训练与推理掌握--data-path、--weights、--freeze-layers、num_classes等关键参数的用法并理解窗口自注意力W-MSA/SW-MSA、Patch Merging 等核心组件在源码中的实际实现。一、模块结构与文档定位pytorch_classification/swin_transformer目录是仓库中 Swin Transformer 图像分类的完整实现包含以下关键文件model.pySwin Transformer 网络结构完整实现并提供了多种规格的预训练模型构建函数train.py训练入口脚本负责数据加载、模型创建、优化器配置与训练循环predict.py单张图片分类预测脚本my_dataset.py自定义数据集类MyDataSet与collate_fnutils.py数据集划分自动生成class_indices.json、训练/验证单轮逻辑create_confusion_matrix.py基于验证集绘制混淆矩阵并输出 Precision / Recall / Specificityselect_incorrect_samples.py将验证集中预测错误的样本及其真实/预测标签写入record.txt。配套文档即本目录下的README.md给出的使用流程可以归纳为九步下载数据集 → 设置--data-path→ 下载预训练权重 → 设置--weights→ 训练 → 导入相同模型 → 设置model_weight_path→ 设置img_path→ 预测。下文将按此主线逐步展开并结合源码说明每一步背后的实现细节。二、第一步数据集准备花分类数据集文档默认使用花分类数据集 flower_photos代码中所有路径与类别数均围绕该数据集设计。解压后其目录结构为一个类别对应一个文件夹flower_photos/ ├── daisy/ # 雏菊 ├── dandelion/ # 蒲公英 ├── roses/ # 玫瑰 ├── sunflowers/# 向日葵 └── tulips/ # 郁金香数据集的下载来源已在原文档中给出官方 flower_photos.tgz 数据包地址若无法访问可改用文档附带的百度云镜像链接与提取码。解压后需记录文件夹的绝对路径供后续--data-path参数使用。关于数据集的划分本模块并不依赖外部划分脚本而是由 utils.py 中的read_split_data函数在训练时自动完成遍历根目录下所有子文件夹作为类别按名称排序后生成类别名 → 数字索引的映射以val_rate0.2为比例使用random.seed(0)固定随机种子保证结果可复现在每个类别内随机采样约 20% 的样本作为验证集其余作为训练集自动将类别映射写入当前目录的class_indices.json这正是文档中训练过程中会自动生成 class_indices.json 文件的由来仅支持.jpg、.JPG、.png、.PNG四种图片后缀。从源码可以看出训练脚本与预测脚本通过共享这份class_indices.json保持类别顺序一致因此训练完成后切勿手动修改或删除该文件。三、第二步与第三步设置数据集路径并加载预训练权重3.1 在 train.py 中设置 --data-path训练脚本 使用argparse解析命令行参数其中数据集路径的默认值为/data/flower_photos需要改成你本机解压后的绝对路径python train.py --data-path /your/path/to/flower_photos也可直接修改 train.py 中的默认值。脚本内部通过read_split_data(args.data_path)返回训练/验证的图片路径与标签列表再交给自定义的MyDataSet见 my_dataset.py加载。3.2 下载预训练权重并设置 --weightsmodel.py中每个模型构建函数都注明了对应该规格的 ImageNet 预训练权重下载地址例如swin_tiny_patch4_window7_224、swin_base_patch4_window12_384等。本模块默认采用Swin-Tinywindow7 / 224 输入from model import swin_tiny_patch4_window7_224 as create_model下载好对应权重后将 train.py 中的--weights参数设置为权重文件路径python train.py --data-path /your/path/to/flower_photos \ --weights ./swin_tiny_patch4_window7_224.pth如果暂不加载预训练权重可将--weights置为空字符串。3.3 权重加载的源码细节训练脚本加载预训练权重时做了两件关键处理见 train.py从权重文件中取出model键对应的权重字典删除所有包含head的键即分类头层因为 ImageNet 预训练模型的 1000 类分类头与花分类的 5 类输出不匹配以strictFalse方式载入剩余权重。weights_dict torch.load(args.weights, map_locationdevice)[model] for k in list(weights_dict.keys()): if head in k: del weights_dict[k] print(model.load_state_dict(weights_dict, strictFalse))这样主干网络Patch Embedding、各 Stage 的 Swin Transformer Block、Patch Merging、最终 LayerNorm都使用 ImageNet 预训练初始化而分类头从头训练从而显著降低训练难度并加快收敛。3.4 冻结主干只训分类头若资源有限或希望快速验证流程可将--freeze-layers设为True。此时 train.py 会遍历模型参数将除head之外的所有参数requires_grad_(False)冻结仅训练分类头python train.py --data-path /your/path/to/flower_photos \ --weights ./swin_tiny_patch4_window7_224.pth \ --freeze-layers True四、核心参数与训练配置解析训练脚本 中所有可调参数及默认值如下参数默认值说明--num_classes5分类类别数需与数据集文件夹数量一致--epochs10训练轮数--batch-size8批大小过大时需留意显存--lr0.0001初始学习率--data-path/data/flower_photos数据集根目录需修改为本机路径--weights./swin_tiny_patch4_window7_224.pth预训练权重路径空串表示不加载--freeze-layersFalse是否冻结除 head 外的所有层--devicecuda:0训练设备无 GPU 时自动回退到 CPU训练相关实现要点数据增强train.py训练集使用RandomResizedCrop(224)RandomHorizontalFlip验证集使用Resize(224*1.143)CenterCrop(224)即先等比放大再中心裁剪避免直接缩放造成形变两者均使用 ImageNet 统计的均值[0.485, 0.456, 0.406]与标准差[0.229, 0.224, 0.225]做标准化。数据加载器train.pyDataLoader 的num_workers取min(os.cpu_count(), batch_size, 8)训练集shuffleTrue、验证集shuffleFalse并使用自定义collate_fn将图片堆叠为 batch。优化器使用 AdamWlr0.0001weight_decay5E-2与 Swin Transformer 官方训练策略一致。损失函数与指标train_one_epoch与evaluate见 utils.py均使用CrossEntropyLoss并统计每轮准确率训练过程中还会对 loss 做isfinite检查出现非有限值会提前终止并打印警告。日志与权重保存每轮通过SummaryWriter将 train_loss / train_acc / val_loss / val_acc / learning_rate 写入 TensorBoard每个 epoch 结束都会将模型权重保存为./weights/model-{epoch}.pth训练前会自动创建weights目录见 train.py。启动训练的完整命令python train.py \ --data-path /your/path/to/flower_photos \ --weights ./swin_tiny_patch4_window7_224.pth \ --num_classes 5 \ --epochs 10 \ --batch-size 8 \ --lr 0.0001五、Swin Transformer 网络结构源码解读要让训练与预测真正跑得明白有必要理解 model.py 中 Swin Transformer 的核心设计。该实现来自微软 Swin-Transformer 官方仓库的 PyTorch 移植主体包括以下组件5.1 Patch Embedding图像切块嵌入PatchEmbedmodel.py使用一个kernel_sizepatch_size, stridepatch_size的卷积默认 patch_size4将 224×224 图像切分为 56×56 个 4×4 的 patch并把每个 patch 映射为embed_dimTiny 为 96维向量。源码对 H、W 不是 patch_size 整数倍的情况做了 padding最后flatten并transpose为[B, HW, C]序列形式方便后续 Transformer 处理。5.2 窗口自注意力 W-MSA 与相对位置偏置WindowAttentionmodel.py实现窗口内的多头自注意力通过qkv线性层一次性生成 Q/K/V再 reshape 并按头拆分注意力分数除以head_dim ** -0.5进行缩放相对位置偏置初始化一张大小为(2*Mh-1)*(2*Mw-1)的可学习偏置表配合预计算的relative_position_index为窗口内任意两个 token 的配对位置提供可学习的偏置项默认窗口 7×7支持传入 mask 用于移位窗口注意力SW-MSA。5.3 Swin Transformer BlockW-MSA 与 SW-MSA 交替SwinTransformerBlockmodel.py是核心计算单元由LayerNorm → 窗口注意力 → DropPath → LayerNorm → MLP构成残差结构。在BasicLayer中偶数下标 Block 使用shift_size0W-MSA奇数下标 Block 使用shift_sizewindow_size//2SW-MSA交替排列shift_size0 if (i % 2 0) else self.shift_sizeSW-MSA 的实现包含四个步骤见 model.pycyclic shift用torch.roll将特征图沿 H、W 方向循环平移shift_sizewindow_partition将特征图划分为互不重叠的窗口model.py窗口内自注意力并叠加create_mask生成的(-100.0/0.0)注意力掩码屏蔽跨窗口区域的信息泄露model.pywindow_reverse 反向 cyclic shift还原特征图并去除 padding。此外Block 中还实现了DropPath随机深度训练时以概率丢弃残差分支测试时恒等通过model.py整个网络按torch.linspace(0, drop_path_rate, sum(depths))线性递增的方式分配各 Block 的丢弃概率默认drop_path_rate0.1。5.4 Patch Merging 与层级结构PatchMergingmodel.py在相邻 Stage 之间将 2×2 邻域间隔采样得到四个子图在通道维拼接后经 LayerNorm 和Linear(4*C, 2*C)降维实现分辨率减半、通道数翻倍的层级式下采样。SwinTransformer主类model.py按depths与num_heads配置堆叠多个BasicLayer最后一个 Stage 之后接 LayerNorm、AdaptiveAvgPool1d(1)与线性分类头head。以默认 Swin-Tiny 为例配置为embed_dim96, depths(2,2,6,2), num_heads(3,6,12,24), window_size7。5.5 可选模型规格一览model.py 提供了 8 个预训练模型构建函数可按需替换train.py/predict.py中的create_model函数名输入尺寸embed_dimdepthsnum_heads预训练数据swin_tiny_patch4_window7_22422496(2,2,6,2)(3,6,12,24)ImageNet-1Kswin_small_patch4_window7_22422496(2,2,18,2)(3,6,12,24)ImageNet-1Kswin_base_patch4_window7_224224128(2,2,18,2)(4,8,16,32)ImageNet-1Kswin_base_patch4_window12_384384128(2,2,18,2)(4,8,16,32)ImageNet-1Kswin_base_patch4_window7_224_in22k224128(2,2,18,2)(4,8,16,32)ImageNet-22Kswin_base_patch4_window12_384_in22k384128(2,2,18,2)(4,8,16,32)ImageNet-22Kswin_large_patch4_window7_224_in22k224192(2,2,18,2)(6,12,24,48)ImageNet-22Kswin_large_patch4_window12_384_in22k384192(2,2,18,2)(6,12,24,48)ImageNet-22K需要注意使用 22K 预训练权重或 384 输入规格时需在 train.py 中同步修改img_size与导入的create_model而create_confusion_matrix.py与select_incorrect_samples.py默认使用的就是swin_base_patch4_window12_384_in22k输入尺寸 384与训练脚本默认的 Swin-Tiny 224 不同若要在同一套权重上做评估请保持模型规格一致。六、使用 predict.py 进行单张图片预测训练完成后按文档步骤使用 predict.py 对单张图片进行分类导入与训练相同的模型from model import swin_tiny_patch4_window7_224 as create_model并确保num_classes5与训练时一致设置权重路径model_weight_path默认指向./weights/model-9.pth改成训练产出默认保存在weights文件夹下的权重文件路径设置图片路径img_path默认值为../tulip.jpg改成待预测图片的绝对路径。model_weight_path ./weights/model-9.pth # 改为训练好的权重 img_path ../tulip.jpg # 改为待预测图片预测流程的源码细节预处理与训练验证集完全一致Resize(int(224 * 1.14))→CenterCrop(224)→ToTensor→ 相同均值的Normalize保证输入分布与训练时对齐加载类别映射读取训练时自动生成的class_indices.json得到索引 → 类别名的字典推理model.eval()后置于torch.no_grad()上下文将单张图片unsqueeze(0)扩展 batch 维度输入模型对输出做softmax得到各类别概率argmax取概率最大的类别索引结果展示终端打印每个类别的名称与概率保留 3 位有效数字并用matplotlib显示图片与标题class: xxx prob: x.xxx。七、进阶混淆矩阵与错误样本分析除训练与预测外模块还提供了两个验证模型质量的辅助脚本均直接读取验证集可与训练产出的权重配合使用。7.1 混淆矩阵 create_confusion_matrix.py该脚本create_confusion_matrix.py默认使用swin_base_patch4_window12_384_in22k模型与 384 输入尺寸在验证集上统计混淆矩阵并输出指标总体准确率sum(TP) / sum(all)每个类别输出Precision精确率、Recall召回率、Specificity特异度使用matplotlib绘制带数值标注与 colorbar 的混淆矩阵热力图横轴 True Labels、纵轴 Predicted Labels。运行前需设置--data-path数据集根目录、--weights权重路径与--num_classes并依赖prettytable库格式化输出表格。7.2 错误样本挑选 select_incorrect_samples.py该脚本select_incorrect_samples.py逐样本比较预测类别与真实标签把预测错误的样本写入record.txt每行格式为图片绝对路径 TrueLabel:真实类别 PredictLabel:预测类别通过这份记录可以快速定位模型容易混淆的类别与图片辅助分析数据标注问题或类间相似性是排查分类错误的实用工具。八、使用自定义数据集文档最后一步明确指出如果要使用自己的数据集请按照花分类数据集的文件结构进行摆放一个类别对应一个文件夹并将训练与预测脚本中的num_classes设置成自己数据的类别数。即my_dataset/ ├── cat/ ├── dog/ └── bird/操作要点将--data-path指向该目录的绝对路径将--num_classes改为实际类别数如 3 类则设为 3同步修改 predict.py 中的create_model(num_classes...)保证预测时输出维度与训练一致训练后重新生成class_indices.jsonread_split_data会在启动时自动生成并覆盖预测脚本会读取最新映射。MyDataSetmy_dataset.py对图片有 RGB 校验非 RGB 模式的图片会抛出ValueError因此自定义数据集的图片请统一为 RGB 格式JPG/PNG 即可无需额外标注文件类别由文件夹名决定。九、快速上手命令汇总按文档九步流程整理的完整命令序列# 1. 下载并解压 flower_photos 数据集到本机 # 2. 下载 swin_tiny_patch4_window7_224.pth 预训练权重到当前目录 # 3. 启动训练自动生成 class_indices.json 与 weights/model-*.pth python train.py \ --data-path /your/path/to/flower_photos \ --weights ./swin_tiny_patch4_window7_224.pth \ --num_classes 5 \ --epochs 10 \ --batch-size 8 # 4. 修改 predict.py 中的 model_weight_path 与 img_path 后执行预测 python predict.py # 可选在验证集上绘制混淆矩阵 python create_confusion_matrix.py --data-path /your/path/to/flower_photos --weights ./weights/model-9.pth # 可选筛选预测错误的样本 python select_incorrect_samples.py --data-path /your/path/to/flower_photos --weights ./weights/model-9.pth若想了解数据集的通用下载与划分方法可参考仓库中的 data_set/README.md 及其 split_data.py 划分脚本本模块的训练流程亦可与仓库内 vision_transformer、ConvNeXt 等分类模块横向对比理解 Transformer 系与卷积系模型的训练配置差异。赞分享示例工程【免费下载链接】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点击查看免费下载相关推荐终极指南如何用YOLOv8 AI自瞄系统提升FPS游戏瞄准精度终极指南如何用YOLOv8 AI自瞄系统提升FPS游戏瞄准精度 想要在FPS游戏中拥有职业选手般的瞄准能力吗基于YOLOv8深度学习的AI自瞄系统为你带来革人工智能计算机视觉游戏开发图像分类模型部署容器pytorch-image-models与Kubernetes部署图像分类模型部署容器pytorch image models与Kubernetes部署 痛点与解决方案 你是否还在为图像分类模型的部署流程复杂而烦恼本文将详人工智能计算机视觉深度学习预训练终极指南如何将电视盒子改造为高性能Linux服务器终极指南如何将电视盒子改造为高性能Linux服务器 想象一下你手中那个闲置的电视盒子突然变成了一个功能强大的Linux服务器可以运行Docker容器、搭建嵌入式开发工具构建工具操作系统上一篇Electron视频编辑器时间轴与特效处理实现下一篇MMDrawerController终极指南了解侧边栏导航技术的未来发展趋势创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

IEEE 802.3cm 400G多模光纤标准解读:SR4.2与SR8物理层参数及设计指南 2026/9/30 7:36:48

IEEE 802.3cm 400G多模光纤标准解读:SR4.2与SR8物理层参数及设计指南

简介:IEEE Std 802.3cm-2020 是 IEEE 发布的以太网修订标准,聚焦多模光纤上 400Gb/s 的物理层与管理参数,面向光模块研发、数据中心网络架构及高速以太网测试工程师。标准新增 Clause 150,定义了 400GBASE-SR8 与 400GBASE-SR4.2 …

阅读更多 →
计算机三级网络技术备考:IP地址规划与路由设计核心知识框架 2026/9/30 7:36:47

计算机三级网络技术备考:IP地址规划与路由设计核心知识框架

简介:这份计算机三级网络技术备考资料面向准备全国计算机等级考试三级网络技术科目的考生,尤其适合需要系统梳理网络原理与工程实践的中高级学习者。资料以PDF文档形式呈现,共1个文件,压缩包约3.35MB,内容围绕网络系统…

阅读更多 →
基于Django+Python的新能源汽车数据分析系统开发实战 2026/9/30 7:36:47

基于Django+Python的新能源汽车数据分析系统开发实战

做毕业设计最怕的不是“难”,而是项目做完你自己都说不清它到底解决了什么问题。这几年我带过的毕设里,凡是做得顺、答辩不被老师追着问、最后还能拿出一套完整作品的,基本都服从同一个规律:选题落点小、数据可获取、技术能闭环。…

阅读更多 →
根分区磁盘空间告急?从诊断清理到LVM扩容全攻略 2026/9/30 7:36:46

根分区磁盘空间告急?从诊断清理到LVM扩容全攻略

挂载根的磁盘空间太小,这次咱们一次性解决只要跑过Linux服务器的人,基本都被“挂载根”的分区容量告警折磨过。df -h一敲,红字跳出来,根分区使用率冲到95%以上,紧接着就是服务无响应、日志写不进去、SSH卡到怀疑人生。…

阅读更多 →
Linux终端复用神器tmux:告别窗口多开焦虑,配置实战全解析 2026/9/30 7:36:46

Linux终端复用神器tmux:告别窗口多开焦虑,配置实战全解析

告别“窗口多开”焦虑:Linux 终端神器 tmux,让你的效率翻倍(附超全实战配置)在 Linux 下干活时间久了,特别是天天泡在终端里的人,基本都会碰到这么几个场景:SSH 连到服务器,跑着一个…

阅读更多 →
基于Django的证券分析系统开发实战:数据采集到K线展示全解析 2026/9/30 7:36:39

基于Django的证券分析系统开发实战:数据采集到K线展示全解析

去年帮一个学弟远程调试这套基于Django的证券分析系统时,我第一次认真审视"毕设全套源码"这类项目的水有多深。他拿到手的源码压缩包超过1GB,解压后光模型迁移文件就有几十个,数据库却是空的,依赖装了三遍还是报错&…

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞 ✉