用SwinTransformer实现水果图像分类:一份带数据集的完整实战指南

发布时间:2026/10/2 9:47:45
用SwinTransformer实现水果图像分类:一份带数据集的完整实战指南 简介基于Swin Transformer的水果6分类图像分类实战项目面向正在学习PyTorch、希望掌握完整图像分类训练流程的深度学习入门者及课程实践者。整套资源为7z压缩包大小约116.34MB现有32人学习借鉴。项目实现了从数据加载到模型评估的完整闭环通过ImageDataset类完成标准化预处理训练集采用随机裁剪、水平翻转、颜色增强等动态策略验证集仅做调整与归一化统一至224×224分辨率训练过程支持GPU加速自动记录损失值、准确率、精确率、召回率、特异度与F1分数六类指标并逐轮生成验证评估报告、动态保存最佳.pth模型至checkpoints目录同时输出含六项指标对比的训练曲线和详细日志方便观察过拟合与欠拟合。配套项目说明书模块化结构清晰可通过命令行配置数据路径、批次大小与学习率也能方便地调整网络结构、自定义数据增强策略和评估指标适合作为水果六分类实战参考与拓展基础。1. 用 SwinTransformer 做水果分类一份能直接跑的实战项目与数据集水果图像分类在工业质检、无人零售、农业分拣这些场景里反复出现但多数人第一次实践 Transformer 图像分类时要么卡在数据集格式上要么被模型训练时间劝退。这份基于 SwinTransformer 的水果 6 分类实战项目核心价值是把「数据准备 → 模型搭建 → 训练评估 → 部署转换」整条链路打包好了自带整理过的水果数据集图像分类的 train/val/test 目录划分明确配合项目说明书里的参数配置用 Swin Transformer Tiny 这种能在单张消费级显卡上跑起来的规模就能拿到一个可用的 6 分类模型。对刚接触 Transformer 图像分类的工程师或者需要快速做分类原型验证的算法岗来说这份资源省掉的是自己苦找数据集、调预训练权重、查窗口注意力参数的时间。2. 数据集与文件结构吃透 SwinTransformer 分类的输入组织方式2.1 数据集的目录划分与类别映射这份资源里的水果数据集不是一张张散图而是按 ImageFolder 格式组织好的目录结构。我先说结论这种目录式组织方式是 PyTorch 训练图像分类模型最省事的输入格式因为它直接对应torchvision.datasets.ImageFolder的加载逻辑。拿到资源后解压数据集目录你会看到类似下面的结构fruit6/ ├── train/ │ ├── apple/ │ │ ├── 001.jpg │ │ └── 002.jpg │ ├── banana/ │ ├── grape/ │ ├── orange/ │ ├── pear/ │ └── watermelon/ ├── val/ │ └── (同上6 个子目录) └── test/ └── (同上6 个子目录)代码里通常会用ImageFolder直接加载它会自动按子目录名字典序映射类别索引。这里有个关键边界apple 是 0、banana 是 1还是按别的顺序取决于目录名的字典序。如果分类结果打印出来和你的预期对不上先检查这里。2.2 标签映射与样本均衡性检查拿到数据后第一件事不是直接开训而是先统计每个类别的样本数量。我一般会写个快速脚本做这个检查避免训练到一半才发现某个类只有几十张图导致验证集指标失真。import os from collections import Counter dataset_root fruit6 def count_samples(root): category_count {} for split in [train, val, test]: split_dir os.path.join(root, split) if not os.path.exists(split_dir): continue num_files sum(len(files) for _, _, files in os.walk(split_dir)) category_count[split] num_files return category_count for split in [train, val, test]: split_dir os.path.join(dataset_root, split) labels os.listdir(split_dir) counts {} for label in labels: label_dir os.path.join(split_dir, label) counts[label] len(os.listdir(label_dir)) print(split, counts)这段脚本做的事情是遍历 train、val、test 三个目录把每个类别下的图片数量打印出来。os.walk递归统计总文件数用于快速核对第二个循环则精确到每个类别。运行后重点看两个数字一是 train/val/test 的比例是否接近 8:1:1二是每个类别样本数是否相差过大。如果某个类别的训练样本数只有另一个类别的三分之一就需要考虑类别加权采样或者数据增强策略否则 Swin Transformer 的注意力机制会偏向样本多的类别。2.3 数据增强策略SwinTransformer 对数据量的要求Swin Transformer 虽然是层级式结构局部窗口注意力比 ViT 的全局注意力对数据量的要求低一些但它的归纳偏置依然不如 CNN 强。这意味着当训练数据量低于几千张时数据增强不是可选项而是必需品。项目说明书里通常会给一组增强策略核心组合是from torchvision import transforms train_transform transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.8, 1.0)), transforms.RandomHorizontalFlip(), transforms.RandomRotation(15), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) val_transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])这里有个细节训练用的RandomResizedCrop的scale参数从 0.8 开始而不是默认的 0.08。因为水果图像大多主体占据画面中心且面积较大如果 scale 下限太低会裁到大量背景区域反而增加学习难度。ColorJitter 的三个通道参数同时调整亮度、对比度、饱和度目的是模拟不同光照环境下的水果拍摄效果。Normalize用的 mean 和 std 是 ImageNet 预训练权重对应的统计值换数据集时通常沿用这套值不需要重新计算。3. 模型原理与关键配置读懂 SwinTransformer 的窗口注意力机制3.1 从 ViT 到 Swin为什么图像分类要用窗口注意力在你把这份实战项目跑通之前有必要先理解 Swin Transformer 在架构上到底改了什么。ViT 的做法是把图像切成固定大小的 patch然后对所有 patch 做全局自注意力这个操作的计算复杂度是序列长度的平方。当输入图像分辨率升高比如从 224 提升到 384patch 数量翻倍计算量直接变成四倍显存压力很大。Swin Transformer 的解决方案是引入分层结构和移动窗口。它把 patch 分组到不重叠的窗口内做局部自注意力每个窗口内的计算量只跟窗口大小有关跟整张图的 patch 总数无关计算复杂度因此从二次方降为线性。同时通过 Shifted Window 机制在不同层之间移动窗口划分让相邻窗口之间的信息可以交互弥补了局部注意力丢失全局感受野的问题。这个改动使得 Swin 在图像分类任务上既能享受 Transformer 的表达能力又能保持接近 CNN 的计算效率。3.2 Swin-Tiny 架构参数与预训练权重加载这份项目里用的骨干网络是 Swin-Tiny这是 Swin Transformer 系列里最小的变体参数量在 2800 万左右和 ResNet50 相当。它的具体配置如下表参数Swin-Tiny 配置Patch Size4x4Patch Embedding 输出维度 96Stage 1 深度2 个 Swin Transformer BlockStage 2 深度2 个 BlockStage 3 深度6 个 BlockStage 4 深度2 个 Block多头注意力头数[3, 6, 12, 24]窗口大小 Win. Size7x7输入分辨率224x224代码中加载预训练权重时需要注意分类头尺寸不匹配的问题。Swin-Tiny 在 ImageNet 上预训练时最后的分类头是 1000 类这里要替换成 6 类。常见的做法是import torch import torch.nn as nn from swin_transformer import SwinTransformer # 项目包含的模型定义 model SwinTransformer( embed_dim96, depths[2, 2, 6, 2], num_heads[3, 6, 12, 24], window_size7, num_classes6, ) # 加载预训练权重忽略分类头 checkpoint torch.load(swin_tiny_patch4_window7_224.pth, map_locationcpu) checkpoint checkpoint.get(model, checkpoint) new_state_dict {k: v for k, v in checkpoint.items() if head not in k} model.load_state_dict(new_state_dict, strictFalse) model.head nn.Linear(model.embed_dim * 8, 6)逐行说明embed_dim96是 Swin-Tiny 的 Patch Embedding 输出通道数后面的 Stage 每下采样一次通道翻倍最终进入分类头前是 96×8 的 768 维特征。depths决定每个 Stage 里堆叠多少个 Transformer BlockSwin-Tiny 的 2-2-6-2 结构里最重的计算量集中在第三阶段。加载预训练权重时用strictFalse并过滤掉head键是因为本地模型的分类头输出维度是 6与预训练的 1000 维不一致strictFalse允许缺失不匹配的键存在其余主干权重正常加载。3.3 窗口大小与输入分辨率的关系Swin Transformer 的窗口大小参数window_size7是针对 224×224 输入设计的。这里有一个非常容易踩的边界输入图像的尺寸必须能被patch_size × window_size的倍数整除。224 除以 4patch后得到 56×56 的特征网格56 能被 7 整除所以窗口划分是完整的。如果你把输入分辨率从 224 改成 384那特征网格变成 96×9696 除以 7 有余数窗口划分时就需要 padding。多数开源实现会自动处理这种情况但性能会受影响。如果训练脚本里输入的尺寸不是 224你可以用下面这段代码检查窗口划分是否对齐import math patch_size 4 window_size 7 input_size 224 grid_size input_size // patch_size print(Feature grid:, grid_size) assert grid_size % window_size 0, \ fgrid_size {grid_size} 无法被 window_size {window_size} 整除 # 输出Feature grid: 56 # 56 % 7 0窗口划分正常这段检查逻辑的意义在于Swin Transformer 在构建window_partition函数时默认窗口数量是(H // window_size) × (W // window_size)如果 H 和 W 不是 window_size 的整数倍剩余像素会被丢弃。数据集里的图像虽然会被增强管线统一 resize 成 224但如果你后续要换更高分辨率微调务必先用这个断言确认整除关系。4. 训练实战从数据加载到模型收敛的完整流程4.1 数据加载器与训练超参数配置把上一章的数据集和模型组合起来就进入实际训练流程了。数据加载器负责把磁盘上的图像按批次喂给 GPU这里需要设置num_workers和pin_memory来加速数据读取。训练超参数则决定了模型能不能收敛到理想精度。from torch.utils.data import DataLoader from torchvision.datasets import ImageFolder train_dataset ImageFolder(fruit6/train, transformtrain_transform) val_dataset ImageFolder(fruit6/val, transformval_transform) train_loader DataLoader( train_dataset, batch_size64, shuffleTrue, num_workers8, pin_memoryTrue, drop_lastTrue, ) val_loader DataLoader( val_dataset, batch_size64, shuffleFalse, num_workers8, pin_memoryTrue, ) print(类别映射:, train_dataset.class_to_idx) print(训练集样本数:, len(train_dataset))ImageFolder返回的class_to_idx字典展示了类别名到索引的映射关系这个映射在推理阶段需要保存下来。drop_lastTrue在训练集样本数不能被 batch_size 整除时丢弃最后不足一个 batch 的数据避免 BatchNorm 层因 batch 过小导致统计量抖动。如果显存有限batch_size 降到 32 时需要同步调整学习率经验法则是 batch_size 减半时学习率也减半。4.2 优化器选择AdamW 与余弦退火调度器Swin Transformer 的训练建议使用 AdamW 优化器权重衰减系数设为 0.01 或 0.05。相比标准 AdamAdamW 把权重衰减从梯度更新中解耦对 Transformer 这种大参数量模型效果更稳定。学习率策略上项目说明书里通常建议使用余弦退火调度器并配合线性 warmup。import torch.optim as optim optimizer optim.AdamW( model.parameters(), lr1e-4, weight_decay0.05, betas(0.9, 0.999), ) total_epochs 100 warmup_epochs 5 def lr_lambda(epoch): if epoch warmup_epochs: return (epoch 1) / warmup_epochs else: progress (epoch - warmup_epochs) / (total_epochs - warmup_epochs) return 0.5 * (1.0 math.cos(math.pi * progress)) scheduler optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)lr1e-4是 Swin-Tiny 在自定义数据集上的常见起点ImageNet 预训练权重已经收敛到了不错的特征空间微调时学习率不宜过大。warmup_epochs5让学习率在前 5 个 epoch 从 0 线性升到1e-4这段热身期的作用是让模型权重在预训练基础上逐步适应新数据的分布避免一开始就用大步长把预训练学到的特征破坏掉。余弦退火部分则让学习率按照余弦曲线从峰值降到接近 0这比固定学习率的收敛精度更高。4.3 训练主循环与验证指标计算训练主循环是整份代码的核心它把数据加载、前向传播、反向传播、学习率调度串起来。每一轮训练结束后跑一次验证集记录 top-1 准确率最终保存验证集准确率最高的模型权重。import torch.nn as nn from tqdm import tqdm criterion nn.CrossEntropyLoss() device torch.device(cuda if torch.cuda.is_available() else cpu) best_acc 0.0 model.to(device) for epoch in range(total_epochs): model.train() running_loss 0.0 for images, labels in tqdm(train_loader): images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() running_loss loss.item() * images.size(0) epoch_loss running_loss / len(train_dataset) model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.to(device), labels.to(device) outputs model(images) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() val_acc correct / total print(fEpoch {epoch1}/{total_epochs}, Loss: {epoch_loss:.4f}, Val Acc: {val_acc:.4f}) if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_model.pth) print(f保存最优模型验证准确率 {val_acc:.4f})这里criterion是交叉熵损失它内部做了 log_softmax 和负对数似然的计算所以模型输出的 logits 不需要额外过 softmax。optimizer.zero_grad()必须在每次反向传播前清空梯度否则 PyTorch 会累加梯度导致参数更新方向错误。torch.no_grad()在验证阶段阻断梯度计算减少显存占用并加速推理。实际训练时Swin-Tiny 在 6 分类水果数据集上100 个 epoch 大约需要 10~15 分钟单张 RTX 3090 级别显卡验证准确率通常能到 95% 以上。4.4 训练过程中的监控指标与损失曲线解读训练时除了看验证集准确率还要关注训练损失的下降趋势。正常情况是训练损失前几个 epoch 快速下降然后逐渐变缓验证准确率同步上升。如果出现训练损失下降但验证准确率不升反降说明过拟合了这时候应该增大 weight_decay 或增强数据增强强度。如果训练损失在某几个 epoch 不降反升先检查学习率是不是太大。Swin Transformer 对学习率比 CNN 更敏感5e-5到1e-4通常是安全区间。还有一个容易遇到的问题loss 变成 NaN这通常是梯度爆炸解决方案是降低学习率或增加梯度裁剪nn.utils.clip_grad_norm_(model.parameters(), 5.0)。5. 踩坑与排查SwinTransformer 实战中常见的四个问题5.1 预训练权重加载报错size mismatch for head现象是训练脚本一启动就报RuntimeError: Error(s) in loading state_dict for SwinTransformer: size mismatch for head.weight模型卡在加载权重这一步。原因是本地模型的分类头是 6 维输出而 ImageNet 预训练权重的分类头是 1000 维输出load_state_dict默认要求键名和形状完全一致。解决方法是加载时过滤掉 head 键state_dict torch.load(swin_tiny_patch4_window7_224.pth, map_locationcpu) filtered {k: v for k, v in state_dict.items() if head not in k} model.load_state_dict(filtered, strictFalse)如果不用strictFalse哪怕只有一个键不匹配程序都会直接终止。加入过滤逻辑后主干权重正常加载随机初始化的分类头从头训练这符合迁移学习的标准做法。5.2 验证集准确率在 90% 左右卡住上不去了现象是训练到 40~50 个 epoch 后验证准确率稳定在 90%~93% 之间但 train loss 还在下降再往上调学习率也不管用。原因是这属于细粒度分类问题苹果的红富士和黄元帅、梨的多个品种外观差异太细微仅靠整图分类的 Swin-Tiny 难以区分。解决的突破口不在模型结构而在数据层面# 提升分辨率从 224 提升到 384让模型看到更多纹理细节 train_transform transforms.Compose([ transforms.Resize((384, 384)), transforms.RandomHorizontalFlip(p0.5), transforms.RandomRotation(degrees10), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ])把输入分辨率从 224 提升到 384特征网格从 56×56 变为 96×96Swin Transformer 的窗口数量从 64 变为 144模型能捕捉的局部细节更多。代价是显存占用约为原来的约 3 倍如果显存不够可以减小 batch_size。这个方法在水果细分类任务上通常能带来 2~3 个百分点的提升。5.3 显存不足CUDA out of memory现象是训练脚本运行到第二个 epoch 时报CUDA out of memory第一个 epoch 没事第二个 epoch 跑一半炸了。原因通常是 batch_size 设置过大或者输入分辨率调整后显存需求超了还有一种可能是验证阶段没有关梯度。解决方法是先调 batch_size同时确保验证时用torch.no_grad()# 显存不足时优先调整 batch_size train_loader DataLoader( train_dataset, batch_size32, # 原来 64现在减半 shuffleTrue, num_workers8, pin_memoryTrue, drop_lastTrue, ) # 验证阶段强制关闭梯度计算 with torch.no_grad(): val_loss criterion(outputs, labels)显存的占用大头是激活值batch_size 减半后显存需求约减半但训练时间会相应延长。如果 batch_size 已经很小还是显存不足升级做法是启用梯度累积optimizer.step()每 4 个 batch 执行一次等效于维持较大的 batch。5.4 推理时预测结果全部输出同一个类别现象是训练过程一切正常验证准确率 95%但加载保存的权重对单张图片做推理时所有图片都输出同一个类别。原因大概率是数据预处理不对齐。训练时用的是RandomResizedCropColorJitter推理时如果忘记套用相同的 Resize 和 Normalize输入图像的分布跟训练时不一致模型输出会退化。# 推理时的预处理管线必须与验证集保持一致 inference_transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ])另一个原因是模型没有调到 eval 模式或整体卡在训练状态导致 BatchNorm 层使用训练统计量。推理前记得调用model.eval()并检查输入图像的通道顺序PyTorch 默认是 RGB如果数据源是 OpenCV 读入的 BGR需要先转换。6. 模型部署与进阶ONNX 导出和 Grad-CAM 可视化验证模型训练完成后还有一个动作值得做导出 ONNX 格式做部署验证以及用 Grad-CAM 可视化确认模型关注的区域确实落在水果上而不是背景。这两个操作一个解决「模型能不能上线」的问题一个解决「模型学到的特征对不对」的问题。ONNX 导出用 PyTorch 的torch.onnx.export就可以完成核心是固定输入尺寸和 batch 维的语义import torch.onnx model.eval() dummy_input torch.randn(1, 3, 224, 224).to(device) torch.onnx.export( model, dummy_input, fruit6_classifier.onnx, input_names[input], output_names[output], opset_version12, dynamic_axes{input: {0: batch_size}, output: {0: batch_size}}, ) print(ONNX 导出完成)dynamic_axes允许 ONNX 模型在推理时接收可变 batch但如果你确定线上场景每次只处理单张图片去掉 dynamic_axes 反而能获得更好的优化。导出后用 ONNX Runtime 做一次推理和 PyTorch 的输出对齐验证import onnxruntime as ort import numpy as np ort_session ort.InferenceSession(fruit6_classifier.onnx) ort_inputs {ort_session.get_inputs()[0].name: dummy_input.cpu().numpy()} ort_outputs ort_session.run(None, ort_inputs)[0] with torch.no_grad(): torch_outputs model(dummy_input).cpu().numpy() print(最大输出差异:, np.abs(ort_outputs - torch_outputs).max())PyTorch 的模型默认对 logits 不做 softmaxONNX Runtime 的输出也一样两者直接比较即可。差异在1e-5量级以下说明导出无误。这个验证动作虽然只花两分钟但它能提前暴露算子兼容问题避免部署环境里才发现推理结果不对。Grad-CAM 可视化则用来回答「模型凭什么把它分对」这个问题具体是取最后一个 Transformer Block 的输出特征图结合分类头的梯度计算热度图from grad_cam import GradCAM, visualize_cam cam GradCAM(modelmodel, target_layermodel.layers[-1]) input_tensor load_inference_image(test.jpg) # 形状 [1, 3, 224, 224] cam_map cam.generate(input_tensor, target_class0) # 0 对应 apple visualize_cam(input_tensor, cam_map, save_pathattention_map.jpg)注意这里target_layer选择最后一层 Transformer Block 输出的特征图分辨率是 7×7需要上采样到 224×224 才能和原图叠加显示。如果可视化结果显示模型高亮的区域不是水果本体而是桌面、叶子等背景说明训练数据里这些背景和类别强相关得回到数据清洗环节去修。从那以后我每次训练图像分类模型都强制跑一遍 Grad-CAM 验证先看正确分类样本的热度图再挑几个错分样本看它到底在盯哪里。这个习惯帮我排查过好几次「验证集 96% 但真实场景一测就翻车」的问题希望帮到你。本文还有配套的精品资源点击获取