ARTICLE DETAIL

资讯详情

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

SwinTransformer水果分类实战:中小数据集的图像分类完整方案

SwinTransformer水果分类实战:中小数据集的图像分类完整方案 简介面向深度学习图像分类初学者的SwinTransformer水果6分类实战项目包含完整项目说明书与水果数据集可帮助用户快速掌握基于PyTorch的训练、评估与可视化全流程。项目代码模块清晰由主训练脚本、数据处理模块与训练评估工具组成支持命令行配置数据路径、批次大小、学习率等参数并自动记录损失值、准确率、精确率、召回率、特异度和F1分数六类指标动态保存最佳模型权重与训练曲线。训练集采用随机裁剪、水平翻转和颜色增强等策略增强泛化能力验证集统一归一化至224×224分辨率用户可灵活修改模型类或数据增强函数扩展至其他分类任务。压缩包为7z格式整体约116.34MB以Python脚本、项目说明书和按类别分目录的水果图像数据集为主。已有32人学习适合从零搭建分类系统或研究SwinTransformer应用的学习者也可作为课程设计与毕业设计的参考基线。1. 为什么用SwinTransformer做水果6分类一个本地能跑完的Transformer图像分类落地样本手头只有两三万张水果图、六七个类别很多人第一反应是拿ResNet50或EfficientNet直接训觉得Transformer是给ImageNet那种千万级数据准备的。可实际跑下来会发现SwinTransformer在中小规模图像分类任务上一点不虚CNN甚至把验证集准确率从91%推到96%以上。这个项目标题的核心不是“水果分类”本身而是把2021年之后公认最能打的层次化视觉Transformer用一套完整可复现的工程流程装进一个普通显卡就能跑完的实战项目里数据集、项目说明书、训练脚本、推理脚本一次性串齐。适合两类人一是刚入门Transformer图像分类、想拿真实数据集练手的学生二是已经在用CNN做分类、想评估SwinTransformer值不值得替换的生产环境工程师。全文围绕怎么组织数据、怎么写训练代码、参数怎么调、坑在哪里展开跟着做就能复现出八九十分的效果。2. 拆解SwinTransformer的模型核心它凭什么取代CNN做图像分类2.1 窗口自注意力与移位窗口局部建模与跨窗口信息交换SwinTransformer和ViT最大的区别是它不一开始就在全局算自注意力而是把图像切成不重叠的窗口在每个窗口内部算注意力。以swin_tiny_patch4_window7_224为例输入224×224的图像先按4×4的patch切块得到56×56个patch每个patch是4×4×348维的向量然后进入4个Stage每个Stage的特征图尺寸依次是56×56、28×28、14×14、7×7通道数依次是96、192、384、768。窗口大小为7×7所以在第一个Stage里56×56的特征图被分成8×864个窗口每个窗口内部做自注意力计算量从全局注意力的O(N²)降到了O(W×H×窗内patch数)。这个设计让它比ViT省显存也更好迁移到分辨率更高的输入上。但只在固定窗口内算注意力会损失全局视野所以SwinTransformer在偶数层使用移位窗口Shifted Window把窗口的划分边界偏移半个窗口大小让相邻窗口的patch有机会出现在同一个窗口里从而完成跨窗口信息交换。用代码看更直接# 配置一个SwinTransformer窗口大小和移位机制都封装在模型内部 import timm model timm.create_model( swin_tiny_patch4_window7_224, # tiny版本参数约28M pretrainedTrue, num_classes6 # 改成自己的水果类别数 ) print(model.patch_embed) # 4x4 patch切块 print(model.layers[0].blocks[0]) # 第1层WindowAttention print(model.layers[0].blocks[1]) # 第2层Shifted Window Attention print(model.norm) # 最后的LayerNorm这段代码用timm库把整个SwinTransformer的骨架打印出来能看到模型内部确实交替排列了WindowAttention和Shifted Window Attention。实际训练时不需要手动实现窗口划分因为模型已经把这些逻辑封装成层了但理解这个结构有助于判断什么时候该换更大的窗口、什么时候该改patch大小。关于移位窗口的实现细节常见做法是在注意力模块里维护一个attn_mask用来屏蔽移位后跨边界窗口的无效区域。打印模型的blocks[1].attn_mask就能看到这个mask的形状与窗口偏移量相关。如果你后续想改窗口大小记得同时调整这个mask的生成逻辑否则训练时会看到loss发飘、注意力泄漏到无关区域。2.2 层次化下采样为什么小数据集也需要4个StageSwinTransformer保留了CNN式的金字塔结构每个Stage结尾通过PatchMerging把分辨率减半、通道翻倍。这个设计对水果分类这种中小规模数据集非常关键。第一金字塔结构让模型前几层专注于局部纹理苹果表面的斑点、橙子的气孔后几层逐渐聚合出全局形状圆形轮廓、果梗位置这正好匹配水果分类需要的多尺度特征。第二层次化结构天然支持输入多尺度比如训练时用224微调时可以拉大到384不需要改动模型结构。在水果6分类场景里最常用的配置是swin_tiny_patch4_window7_224tiny版本参数量28M左右比ResNet50的25M略大但实际推理速度相当因为窗口注意力在224分辨率下计算量更低。如果你的显卡显存小于6G可以考虑swin_tiny_patch4_window7_224加上梯度累积如果只有4G显存建议直接用swin_tiny_patch4_window7_224配合batch_size8或者干脆换用Swin-S。不需要一上来就上base版本水果6分类的数据量和特征复杂度撑不起base的容量容易过拟合。2.3 分类头与预训练权重的正确打开方式SwinTransformer默认的分类头是一个全连接层接在最后一个Stage的LayerNorm之后。用timm创建模型时会把特征维度从768直接映射到6类输出。这里有个常见误操作有人喜欢在模型后面再接一个自己的FC层然后把原模型的head去掉结果两层全连接之间缺了BatchNorm或者Dropout训练时很快过拟合。我一般习惯直接用timm的num_classes参数让它替换掉内置head然后在head前面加一个Global Average Poolingtimm默认已经包含和Dropouttimm里默认是0.0需要手动设置。# 重新设置dropout rate防止小数据集过拟合 model timm.create_model( swin_tiny_patch4_window7_224, pretrainedTrue, num_classes6, drop_rate0.3, # 分类头前的dropout drop_path_rate0.2 # 每个block的stochastic depth比例 )drop_rate控制分类头前的随机失活比例drop_path_rate控制每个Transformer block被随机丢弃的概率这两个参数在数据量只有两三万张时非常关键。经验值是drop_rate设0.2到0.4之间、drop_path_rate设0.1到0.3之间如果你的数据集在验证集上表现良好但训练集准确率远高于验证集说明过拟合了把drop_path_rate往0.3调。用预训练权重时注意一个容易踩的坑timm从 huggingface 下载Swin预训练权重文件大约110MB第一次运行需要网络连接。如果网络受限timm会报“Unable to load weights”之类的错误解决办法是手动下载权重文件放到~/.cache/torch/hub/checkpoints/目录下。后续模型加载就不会再访问网络。3. 准备水果6分类数据集目录结构、划分脚本与预处理参数3.1 构建符合ImageNet风格的目录结构水果6分类项目里的数据集一般按train和val两个文件夹组织每个类别一个子目录。训练时用torchvision.datasets.ImageFolder直接读取这种结构不需要额外写Dataset类。6类水果可以从Felex/Fruit-Image-Dataset这类开源水果图集里选一个苹果、香蕉、樱桃、葡萄、橙子、草莓的常见组合或者用自己拍摄的商品果图。但要注意一个现实问题公开水果数据集的图片尺寸非常不统一有的只有320×240有的是1080×1920竖拍图背景也五花八门所以数据集预处理的核心是全流程归一化到模型输入尺寸。# 推荐的数据集目录结构示例 dataset/ ├── train/ │ ├── apple/ # 共3000张 │ ├── banana/ # 共2800张 │ ├── cherry/ # 共2600张 │ ├── grape/ # 共2900张 │ ├── orange/ # 共3100张 │ └── strawberry/ # 共3000张 ├── val/ │ ├── apple/ # 每类500张 │ ├── banana/ │ ├── cherry/ │ ├── grape/ │ ├── orange/ │ └── strawberry/ └── test/ ├── apple/ # 每类300张用于最终评估 ├── banana/ ├── cherry/ ├── grape/ ├── orange/ └── strawberry/目录结构本身没有魔法关键是类别的划分逻辑。我一般会把所有图片按照类别先集中到一个总目录然后用脚本按固定比例划分而不是手动拖拽文件。手动划分很容易在验证集里混入同一棵树同一批次的照片导致验证分数虚高部署时翻车。3.2 写一个数据划分脚本保证随机性并让分布可追溯手头数据如果是零散的文件夹最稳妥的做法是先按类别汇总再做一个可复现的随机划分。写一个划分脚本固定随机种子让后续实验可以复现。# split_dataset.py import os import random import shutil from collections import defaultdict random.seed(42) SRC_ROOT raw_fruit_images # 原始图片按类别分子目录 DST_ROOT fruit6_dataset # 输出根目录 TRAIN_RATIO, VAL_RATIO, TEST_RATIO 0.6, 0.2, 0.2 # 收集所有类别及其文件 class_files defaultdict(list) for class_name in os.listdir(SRC_ROOT): class_dir os.path.join(SRC_ROOT, class_name) if not os.path.isdir(class_dir): continue for img_name in os.listdir(class_dir): if img_name.lower().endswith((.jpg, .jpeg, .png)): class_files[class_name].append(img_name) # 按比例划分并复制文件 for class_name, files in class_files.items(): random.shuffle(files) n_train int(len(files) * TRAIN_RATIO) n_val int(len(files) * VAL_RATIO) for split_name, split_files in [ (train, files[:n_train]), (val, files[n_train:n_train n_val]), (test, files[n_train n_val:]) ]: dest_dir os.path.join(DST_ROOT, split_name, class_name) os.makedirs(dest_dir, exist_okTrue) for img_name in split_files: src_path os.path.join(SRC_ROOT, class_name, img_name) dst_path os.path.join(dest_dir, img_name) shutil.copy2(src_path, dst_path) print(f{split_name}/{class_name}: {img_name})逻辑说明脚本先按类别收集所有图片文件名用同一个随机种子shuffle再按6:2:2切成三段复制到新目录。实际使用时把SRC_ROOT改成自己的原始路径即可。参数说明TRAIN_RATIO_VAL_RATIO_TEST_RATIO三者的和必须为1.0。如果原始数据只有1万张可以把比例改成0.7/0.15/0.15保证训练集有7000张以上。水果分类样本量不足时验证集太小会导致准确率波动大每轮可能差2-3个百分点所以验证集至少留1500张。划分完建议检查一下每个目录的文件计数防止某些类别图片数量过少导致抽不到足够样本。3.3 预处理参数与数据增强给SwinTransformer喂对格式SwinTransformer的预处理和CNN模型有细微差别。CNN时代流行先resize到256再中心裁剪到224Swin官方更推荐直接resize到224或者稍微放大一点再裁剪。关键是一定要用ImageNet的均值和标准差做标准化因为预训练模型是在ImageNet统计量上训练的不匹配的均值和标准差会让Swin模型在训练初期震荡严重。# transforms.py from torchvision import transforms train_transforms transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomResizedCrop(224, scale(0.6, 1.0)), transforms.RandomHorizontalFlip(), transforms.RandomRotation(15), transforms.ColorJitter(brightness0.2, contrast0.3, saturation0.3), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) val_transforms transforms.Compose([ transforms.Resize((224, 224)), # 验证集不做增强直接resize transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])逻辑说明训练集先resize到256再用随机裁剪放大到224相当于引入尺度和位置的随机性ColorJitter的亮度、对比度、饱和度扰动对水果分类尤其有效因为同一类水果在不同光照条件下颜色差异很大。验证集和测试集只做resize和标准化不做任何随机增强保证指标稳定可对比。参数说明RandomResizedCrop的scale设0.6到1.0是因为水果分类时主体通常比较大如果scale下限设到0.08有些图会裁到只剩背景反而变成噪声。RandomRotation角度15度足够模拟摆放角度变化超过30度会让果柄和形状失真。实际训练时如果发现验证集波动大优先检查ColorJitter是否太强把颜色特征破坏掉可以先把saturation调到0.2试一轮。4. 训练实战PyTorch加timm的训练脚本与关键参数调节4.1 最小可复现训练脚本从数据加载到保存权重训练脚本用PyTorch加timm实现完整覆盖数据加载、模型构建、训练循环、验证和权重保存。这个脚本在单卡GPU上可以训练出能直接部署的模型。# train_swin.py import torch import torch.nn as nn import timm from torch.utils.data import DataLoader from torchvision.datasets import ImageFolder from torchvision import transforms # ---------- 配置区 ---------- BATCH_SIZE 32 EPOCHS 30 LEARNING_RATE 2e-5 NUM_CLASSES 6 DEVICE cuda if torch.cuda.is_available() else cpu DATA_ROOT fruit6_dataset # ---------- 数据加载 ---------- train_transforms transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomResizedCrop(224, scale(0.6, 1.0)), transforms.RandomHorizontalFlip(), transforms.RandomRotation(15), transforms.ColorJitter(brightness0.2, contrast0.3, saturation0.3), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) val_transforms transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) train_ds ImageFolder(f{DATA_ROOT}/train, transformtrain_transforms) val_ds ImageFolder(f{DATA_ROOT}/val, transformval_transforms) train_dl DataLoader(train_ds, batch_sizeBATCH_SIZE, shuffleTrue, num_workers4, pin_memoryTrue) val_dl DataLoader(val_ds, batch_sizeBATCH_SIZE, shuffleFalse, num_workers4, pin_memoryTrue) # ---------- 模型 ---------- model timm.create_model( swin_tiny_patch4_window7_224, pretrainedTrue, num_classesNUM_CLASSES, drop_rate0.3, drop_path_rate0.2 ) model.to(DEVICE) # ---------- 优化器与loss ---------- optimizer torch.optim.AdamW(model.parameters(), lrLEARNING_RATE, weight_decay0.05) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_maxEPOCHS) criterion nn.CrossEntropyLoss(label_smoothing0.1) # ---------- 训练循环 ---------- best_acc 0.0 for epoch in range(EPOCHS): model.train() train_loss, train_correct, train_total 0.0, 0, 0 for images, labels in train_dl: images, labels images.to(DEVICE), labels.to(DEVICE) outputs model(images) loss criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() batch_acc (outputs.argmax(1) labels).sum().item() train_correct batch_acc train_total labels.size(0) train_loss loss.item() * labels.size(0) scheduler.step() # 验证 model.eval() val_correct, val_total 0, 0 val_loss 0.0 with torch.no_grad(): for images, labels in val_dl: images, labels images.to(DEVICE), labels.to(DEVICE) outputs model(images) loss criterion(outputs, labels) val_loss loss.item() * labels.size(0) val_correct (outputs.argmax(1) labels).sum().item() val_total labels.size(0) val_acc val_correct / val_total train_acc train_correct / train_total print(fEpoch {epoch1}/{EPOCHS} | train_loss {train_loss/train_total:.4f} f| train_acc {train_acc:.4f} | val_loss {val_loss/val_total:.4f} f| val_acc {val_acc:.4f} | lr {scheduler.get_last_lr()[0]:.2e}) if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_swin_fruit6.pth) print(f-- saved best model, val_acc{best_acc:.4f})逻辑说明整个脚本是标准的训练循环。每次迭代读一个batch的图片前向算输出用CrossEntropyLoss计算损失反向传播更新参数再在每个epoch结束后跑一遍验证集。val_acc超过历史最好值时保存权重所以最终得到的best_swin_fruit6.pth就是验证集上表现最好的模型。参数说明learning_rate设2e-5而不是常见的1e-3是因为SwinTransformer在ImageNet上预训练过微调时学习率太大会破坏已经学好的特征。AdamW的weight_decay设0.05是Swin官方配置不要换成0.01或0.0001否则特征更新幅度过大。CosineAnnealingLR配合Swin的微调节奏非常合适总周期30轮足够覆盖从小学习率升到峰值再跌回谷底的过程。label_smoothing设0.1可以让模型不追求给训练集打满100%的置信度减少过拟合。4.2 影响SwinTransformer训练效果的6个关键参数同一套模型代码参数不同结果可能差出5个点以上这几个参数是我反复对比后确认优先级最高的。参数推荐值调整方向影响learning_rate2e-5 ~ 5e-5数据量大可以调高到1e-4微调阶段最敏感过大直接掉点weight_decay0.05过拟合时往0.1加防过拟合的核心drop_path_rate0.2数据少设0.3每类少于2000张就加大抑制过拟合最有效batch_size32显存允许就64小数据集不用刻意加大影响BN统计量和梯度稳定性label_smoothing0.1类别高度相似时0.2让分类边界更平滑T_max余弦周期等于总epoch数改成总epoch的一半可以退火两次决定学习率衰减节奏batch_size在SwinTransformer上比在CNN上更影响验证指标因为Transformer的注意力输出对输入分布更敏感。显存够的情况下尽量用32以上如果显卡只有4Gbatch_size降到8的同时需要把learning_rate按比例降到5e-6否则参数更新步长太大。4.3 训练日志与验证集指标配置从loss曲线判断模型状态训练日志里最值得盯的是train_acc和val_acc的差值。正常情况下第5轮前后train_acc应该超过90%val_acc在85%左右两者差距在5个点以内都是健康的。第10轮以后如果train_acc到了98%而val_acc还停在92%说明进入了过拟合区间这时调drop_path_rate比调weight_decay更有效。如果val_acc从某轮开始不再提升而train_acc还在上涨不要急着提前停止。Swin的微调经常出现平台期后面几轮会因为余弦退火把学习率降下来而突然再涨1-2个点。我一般会训练完30轮再看整体曲线而不是每5轮做一次早停。另外要正确看待验证集指标水果6分类的验证集准确率如果达到95%以上说明模型已经能区分青苹果和绿梨这种近色果实但如果验证集是直接从训练数据里随机抽的实际部署时的准确率会低2-3个点因为真实果品会有新的拍摄角度和光照条件。所以最终评估要拿test目录里的图片做一次盲测不要让模型在验证集上反复调参调成“验证集过拟合怪”。5. SwinTransformer训练水果分类的避坑记录过拟合、AMP数值不稳与数据集陷阱5.1 青苹果和绿梨总是互相混验证集指标高但实际预测乱现象苹果类别里带了青苹果梨类别里也有青绿品种验证集准确率到了95%但抽几张真实场景的青苹果图模型一下判成梨。原因SwinTransformer在前几个Stage特别关注局部纹理和颜色青苹果和绿梨的颜色分布几乎重叠单靠颜色特征无法区分。加上训练数据里这两类的拍摄背景相似度过高其实特征都被背景给带偏了。解决排查训练集里同时含有青苹果和绿梨的图片单独抽出来看是否有大量绿背景重复图。一是给这两个类别加更多背景差异大的样本二是增强里把ColorJitter的色相扰动打开迫使模型关注果皮的质感差异苹果有蜡质反光梨有砂粒感。如果数据实在不够可以考虑把这两个类合并成“青果”类牺牲细粒度换鲁棒性。5.2 训练时loss下降很正常但验证集一直卡在70%上不去现象epoch 1-3训练loss从2.0降到0.6val_acc却一直在0.65到0.70之间横跳连基础水平都达不到。原因最常见的是预处理不匹配。如果忘了用ImageNet的mean/std做Normalize或者Resize之后没有接CenterCrop直接喂给模型Swin的patch_embed会把输入切得和预训练时完全不同模型前几层的统计量全部失效。另一个常见原因是drop_rate设太高比如0.5以上分类头直接被随机丢弃打残。解决先去掉所有数据增强只保留Resize(224)Normalize跑10轮看验证集能不能到85%。能到说明模型和数据没问题再逐步加回增强。卡在70%就先检查Normalize参数再用上面说的最小预处理做冒烟测试。5.3 混合精度训练出现NaNloss变成inf现象打开AMPtorch.cuda.amp训练前几轮正常第5轮突然loss变成NaN权重也污染了。原因SwinTransformer的窗口注意力里有一个attn_mask与logits相加的操作fp16下attn_mask里的-100整数和logits的float16相加时发生溢出。timm的模型在纯fp32下没问题但AMP下经常在深层Stage触发这个问题。解决把优化器换成torch.optim.AdamW并设置foreachFalse或者在训练循环里把关键张量强制转fp32。更省事的方案是关掉AMPSwinTransformer在水果小数据集上训练速度本来就不慢30轮单卡也就1-2小时没必要为一点速度引入数值风险。如果一定要用AMP先跑10轮验证loss稳定再继续。5.4 数据划分时把同一拍摄批次的照片同时分进训练和验证现象训练时val_acc轻松到98%训练集和验证集几乎无差别一到测试集上立刻掉到88%。原因数据集里的图片往往是一张原图多角度裁剪出来的划分脚本按文件名随机切分时同一来源的照片被分到了不同集合验证集实际上和训练集高度重叠指标虚高。这是水果分类最隐蔽的坑。解决划分脚本改成按“图片拍摄批次”或“原始文件夹”切分而不是按单张文件名切分。最稳妥的做法是每个类别下先按子目录如拍摄日期或来源分好再把这些子目录整体划分到train和val中。脚本里可以加一个filter同一目录下的文件只允许进入同一个集合。防止这类问题比调参更值得花时间直接关系模型上线后的真实表现。5.5 权重文件加载时报shape不匹配现象torch.load加载在另一台机器上训练的权重报错提示expected shape和actual shape不一致。原因常见的几个触发点一是创建模型时的drop_rate或drop_path_rate跟训练时不同这不会影响shape但会影响行为二是timm版本不一致导致分类头名称不同三是模型创建时pretrained参数没有与训练时保持一致加载的state_dict里多了或少了head的键。解决加载时先检查state_dict的key列表重点看head相关的键。把创建模型的参数整理成一份config文件写进项目说明书训练和推理严格共用同一个config。自己写加载代码时不要直接用model.load_state_dict(torch.load(path))改为先创建一个相同的模型再load_state_dict否则容易出现key不匹配的隐性问题。代码里加一句model.eval()再推理防止BN和Dropout在推理时引入随机性。6. 单张推理的正确姿势与本地验证习惯训练完拿到了best_swin_fruit6.pth最后一公里是用它做真实的推理验证。很多人直接在训练脚本里加一段预测代码然后发现结果和验证差异很大这里的问题几乎都出在预处理不一致。# infer.py import torch import timm from PIL import Image from torchvision import transforms # 与训练完全相同的标准化参数 val_transforms transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) # 创建与训练一致的模型结构 model timm.create_model( swin_tiny_patch4_window7_224, pretrainedFalse, num_classes6, drop_rate0.0, # 推理时不需要dropout drop_path_rate0.0 ) model.load_state_dict(torch.load(best_swin_fruit6.pth, map_locationcpu)) model.eval() class_names [apple, banana, cherry, grape, orange, strawberry] img_path test_fruit.jpg img Image.open(img_path).convert(RGB) input_tensor val_transforms(img).unsqueeze(0) with torch.no_grad(): logits model(input_tensor) probs torch.softmax(logits, dim1) top_prob, top_idx torch.max(probs, dim1) print(f预测类别: {class_names[top_idx.item()]} f置信度: {top_prob.item():.4f})逻辑说明推理时drop_rate和drop_path_rate必须设为0否则模型仍处于有随机丢弃的状态多次推理结果不一致。用softmax把logits转成概率后取top1同时可以顺手把top2打印出来检查是否有两个类别概率接近的情况这种情况往往对应着前文提到的青苹果绿梨混淆。参数说明map_locationcpu让权重在无GPU的机器上也能加载方便部署到服务器或嵌入式设备。推理时不需要用pretrainedTrue因为加载的是自己训练的权重如果设置pretrainedTrue反而会先下载一个预训练模型再被覆盖浪费时间。最后分享一个我自己的习惯每次训练完除了看测试集准确率一定挑10张来自项目训练集之外的真实水果照片做盲测记录预测结果和置信度存成一个文本文件。这样做的目的是防止模型只适配了数据集的拍摄风格而不是真正学到了果品特征。实战里我踩过最狠的一次模型在测试集上99%准确率但拿超市随手拍的一张苹果图置信度最高的居然是香蕉问题就出在训练集全部是纯白背景棚拍图。这一点在你跟着做这个SwinTransformer水果6分类项目时一定要尽早确认。希望帮到你。本文还有配套的精品资源点击获取
返回列表