別分類實(shí)戰(zhàn):PyTorch與YOLOv5全流程解析)
簡介一份面向?qū)嶋H項(xiàng)目需求的蘑菇圖像分類數(shù)據(jù)集覆蓋姬松茸、阿曼妮塔、牛肝菌、Cortinarius等多個(gè)常見菌類目標(biāo)用戶是想訓(xùn)練YOLOv5分類模型或自建CNN圖像分類網(wǎng)絡(luò)的開發(fā)者與研究者。數(shù)據(jù)已經(jīng)按照訓(xùn)練集和測試集分文件夾存放并附帶JSON類別字典文件清楚標(biāo)注每個(gè)類別的名稱下載后可直接套用到分類模型流程中。壓縮包采用7z格式共包含2000個(gè)文件其中1998張JPG圖片、1個(gè)Python可視化腳本、1個(gè)JSON字典整體大小約97.67MB目錄結(jié)構(gòu)清晰明了。配套的Python腳本可一鍵批量展示樣本圖像方便快速核查各類別圖片數(shù)量與質(zhì)量省去手動(dòng)整理標(biāo)簽和路徑的麻煩。目前已有1046人學(xué)習(xí)使用對需要獲取分類基準(zhǔn)數(shù)據(jù)、快速開展訓(xùn)練試驗(yàn)或進(jìn)行遷移學(xué)習(xí)的入門和進(jìn)階讀者都非常適合。1. 12種蘑菇圖像識(shí)別數(shù)據(jù)集下載前先搞清三件事標(biāo)題寫著12種蘑菇解壓后打開類別字典 json你大概率只數(shù)得到 8 個(gè)類。別急著退貨這個(gè) 12 是原始采集批次的命名口徑真正喂給圖像分類網(wǎng)絡(luò)的是 8 個(gè)類別訓(xùn)練集 9600 張、測試集 2400 張目錄按「train/test 每類一個(gè)文件夾」組織類別字典文件 json 也一并給好了。做圖像分類、森林圖像識(shí)別或者深度學(xué)習(xí)圖像識(shí)別相關(guān)實(shí)驗(yàn)的從業(yè)者最煩的不是沒數(shù)據(jù)而是拿到原始圖后還要自己洗數(shù)據(jù)、劃訓(xùn)練測試集、寫標(biāo)簽文件一下午就耗在預(yù)處理上。這份資源的價(jià)值就在這里目錄已經(jīng)劃分好、標(biāo)簽已對齊直接用 torchvision 的 ImageFolder 或 yolov5 classify 都能吃。它適合正在練手分類網(wǎng)絡(luò)的初學(xué)者也適合需要快速跑一輪分類基線的工程團(tuán)隊(duì)。下載量不大但你省下的是半天到一天的臟活。2. 拆目錄結(jié)構(gòu)train/test 劃分、類別字典 json 與一次跑通的可視化腳本2.1 data 目錄下到底長什么樣先看文件組織再談?dòng)?xùn)練拿到資源后的第一件事不是急著寫訓(xùn)練腳本而是先完整把目錄結(jié)構(gòu)列一遍。常見做法是直接用tree /fWindows或find命令看一遍find data -maxdepth 2 -type d | sort正常你會(huì)看到類似下面的組織結(jié)構(gòu)data/ ├── train/ │ ├── Agaricus/ # 姬松茸 │ ├── Amanita/ # 阿曼妮塔 │ ├── Boletus/ # 牛肝菌 │ ├── Cortinarius/ │ └── ... # 其余類別以 json 實(shí)際內(nèi)容為準(zhǔn) ├── test/ │ ├── Agaricus/ │ ├── Amanita/ │ └── ... └── classes.json # 類別字典文件這份資源的 train 和 test 是按「每類一個(gè)文件夾」組織的不是那種影像分類里常見的images/labels.txt平鋪結(jié)構(gòu)。前者對 PyTorch 的torchvision.datasets.ImageFolder和 yolov5 的 classify 任務(wù)都是開箱即用的不需要再寫額外的路徑解析邏輯。這一點(diǎn)是選型時(shí)最省事的地方。打開classes.json里面是一個(gè)類別名到編號(hào)的映射。以摘要里提到的類別為例能看到姬松茸Agaricus、阿曼妮塔Amanita、牛肝菌Boletus、Cortinarius絲膜菌屬等實(shí)際完整類別清單以你下載到的 json 為準(zhǔn)。也可以用一段非常短的 python 把類名和數(shù)量打出來import json with open(data/classes.json, r, encodingutf-8) as f: class_idx json.load(f) print(f類別數(shù)量: {len(class_idx)}) for name, idx in class_idx.items(): print(f {idx}: {name})這里encodingutf-8是我習(xí)慣性加上的很多 json 文件在 Windows 下用默認(rèn)編碼讀會(huì)直接拋UnicodeDecodeError加一個(gè)顯式編碼參數(shù)能避開一半的路徑類事故??吹筋悇e數(shù)量是 8并且打印出來的類名里沒有亂碼說明數(shù)據(jù)源是完整的可以進(jìn)入下一步。2.2 用 show 腳本做數(shù)據(jù)集可視化9 宮格抽樣看真實(shí)圖像質(zhì)量資源里帶了一個(gè) show 腳本用于可視化數(shù)據(jù)集。它的作用就是隨機(jī)從每個(gè)類里抽圖拼成網(wǎng)格讓你一眼看出這批蘑菇圖像的質(zhì)量有沒有黑邊、有沒有水印、有沒有模糊到?jīng)]法用的圖。如果你下載的腳本能直接跑那就直接跑如果環(huán)境里缺依賴跑不起來下面這個(gè)等效腳本我用得很頻繁直接抄走就行import json import glob import random import matplotlib.pyplot as plt from PIL import Image # 讀取類別字典 with open(data/classes.json, r, encodingutf-8) as f: class_idx json.load(f) classes list(class_idx.keys()) # 3x3 網(wǎng)格每次隨機(jī)抽 9 張 rows, cols 3, 3 fig, axes plt.subplots(rows, cols, figsize(9, 9)) for i in range(rows * cols): c random.choice(classes) # 注意有些圖可能是 .png 或 .jpeg多匹配幾種后綴更穩(wěn) candidates glob.glob(fdata/train/{c}/*.jpg) \ glob.glob(fdata/train/{c}/*.jpeg) \ glob.glob(fdata/train/{c}/*.png) img_path random.choice(candidates) ax axes[i // cols][i % cols] ax.imshow(Image.open(img_path)) ax.set_title(c, fontsize10) ax.axis(off) plt.tight_layout() plt.show()邏輯說明先從classes.json讀類別清單然后用glob按類在 train 目錄下找圖片后綴同時(shí)兼容 jpg、jpeg、png 三種常見格式。random.choice負(fù)責(zé)隨機(jī)抽樣所以每次運(yùn)行看到的 9 張圖都不一樣適合快速把整個(gè)數(shù)據(jù)集的圖像質(zhì)量摸個(gè)大概。參數(shù)上figsize(9, 9)控制整體畫布大小3x3 網(wǎng)格下剛好fontsize10是標(biāo)題字號(hào)。如果你想把全部 8 類都看一遍就把rows, cols改成 2x4或者循環(huán)里固定c classes[i % len(classes)]而不是隨機(jī)抽類??磮D的重點(diǎn)就三個(gè)圖像是否帶多余水印、是否有多張子圖拼在一張里的原始采集圖、是否存在明顯失焦的圖。看到問題圖后建議在訓(xùn)練前手動(dòng)清理而不是指望網(wǎng)絡(luò)自己學(xué)出來。3. 用 PyTorch 訓(xùn)練分類基線ImageFolder 加載、驗(yàn)證集劃分與超參表3.1 用 ImageFolder 加載數(shù)據(jù)并先切出驗(yàn)證集torchvision 的ImageFolder是這類「文件夾即標(biāo)簽」數(shù)據(jù)集最省事的加載方式。它會(huì)自動(dòng)按文件夾名的字母序生成類別索引并把每張圖解析成(image, label)對。這里有一個(gè)值得注意的點(diǎn)classes.json里的編號(hào)順序和ImageFolder按字母序生成的編號(hào)不一定一致所以加載后先打印dataset.classes做一次對齊別急著訓(xùn)練。另外這份數(shù)據(jù)只給了 train 和 test沒有 val。很多新手會(huì)拿 test 邊訓(xùn)練邊驗(yàn)證這是個(gè)大坑后面避坑章會(huì)細(xì)說。正確做法是從 train 里固定切 10% 出來當(dāng)驗(yàn)證集import torch from torch.utils.data import DataLoader, random_split from torchvision import datasets, transforms # 數(shù)據(jù)增強(qiáng)Resize 到 256再隨機(jī)裁剪 224符合 ImageNet 預(yù)訓(xùn)練模型輸入 train_transform transforms.Compose([ transforms.Resize(256), transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.ColorJitter(0.3, 0.3, 0.3), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ]) full_train datasets.ImageFolder(data/train, transformtrain_transform) print(full_train.classes) # 確認(rèn)類別順序 total len(full_train) n_val int(total * 0.1) train_ds, val_ds random_split(full_train, [total - n_val, n_val]) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers4) val_loader DataLoader(val_ds, batch_size64, shuffleFalse, num_workers4)邏輯說明ImageFolder的classes屬性就是它實(shí)際使用的類別順序打印出來跟classes.json對比能提前暴露標(biāo)簽錯(cuò)位。random_split按固定比例把訓(xùn)練集切成 90% 訓(xùn)練 10% 驗(yàn)證驗(yàn)證集只用來看收斂情況不參與權(quán)重更新。Normalize用的是 ImageNet 的均值和標(biāo)準(zhǔn)差因?yàn)楹竺婕虞d的是在 ImageNet 上預(yù)訓(xùn)練過的 ResNet輸入分布必須對齊。參數(shù)說明Resize(256)RandomResizedCrop(224)是分類任務(wù)最常見的尺度策略給裁剪留了隨機(jī)空間相當(dāng)于做了一次尺度擾動(dòng)。ColorJitter(0.3, 0.3, 0.3)對亮度、對比度、飽和度各做 0.3 幅度的隨機(jī)擾動(dòng)對蘑菇這類顏色敏感的類別不建議把幅度調(diào)得更大否則會(huì)破壞菌蓋顏色的判別信息。batch_size32在 ResNet34 8GB 顯存下比較穩(wěn)num_workers4是 Linux 下的常用值Windows 下如果報(bào) DataLoader worker 相關(guān)的錯(cuò)把它改成 0 是最快的解決辦法。3.2 訓(xùn)練腳本與超參設(shè)置ResNet34 在 9600 張上的基線預(yù)處理做完就該上模型了。蘑菇圖像類間差異小、類內(nèi)差異大菌蓋形狀和顏色是主要判別特征所以我一般先用 ResNet34 而不是更大更深的模型跑基線參數(shù)量適中在 9600 張的訓(xùn)練規(guī)模下不容易過擬合訓(xùn)練一輪的時(shí)間也夠短方便反復(fù)試錯(cuò)。import torch import torch.nn as nn from torchvision import models # 類別數(shù)以 json 為準(zhǔn)而不是以標(biāo)題為準(zhǔn) num_classes len(class_idx) model models.resnet34(weightsmodels.ResNet34_Weights.IMAGENET1K_V1) model.fc nn.Linear(model.fc.in_features, num_classes) device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr1e-3) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max30) for epoch in range(30): model.train() running_loss 0.0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() * images.size(0) # 每個(gè) epoch 結(jié)束在驗(yàn)證集上評估一次 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() acc 100.0 * correct / total print(fEpoch {epoch1:02d} | Loss: {running_loss/len(train_ds):.4f} | Val Acc: {acc:.2f}%) scheduler.step()邏輯說明加載在 ImageNet 上預(yù)訓(xùn)練過的 ResNet34只把最后一層全連接換成 8 類輸出這是遷移學(xué)習(xí)里標(biāo)準(zhǔn)的 finetune 做法。預(yù)訓(xùn)練權(quán)重提供了豐富的底層特征蘑菇圖像雖然在 ImageNet 里不多但菌蓋的紋理、邊緣、顏色分布依然能復(fù)用底層卷積特征所以收斂速度遠(yuǎn)快于從零訓(xùn)練。參數(shù)說明優(yōu)化器用 Adam 而非 SGD圖的是前期收斂快lr1e-3對 finetune 來說比較溫和CosineAnnealingLR把學(xué)習(xí)率在 30 個(gè) epoch 內(nèi)按余弦曲線降到接近 0后半程相當(dāng)于做精細(xì)微調(diào)。CrossEntropyLoss自帶 Softmax所以模型輸出層不需要額外加激活。如果你的顯存只有 4GB把batch_size降到 16學(xué)習(xí)率同步降到 5e-4效果差別不大。這里給一張我跑這類蘑菇分類任務(wù)常用的超參表直接照著填問題不大參數(shù)推薦值說明輸入分辨率224x224兼顧細(xì)節(jié)與顯存Batch Size328GB 顯存可跑Epochs30看驗(yàn)證集是否 plateau優(yōu)化器Adam / SGD(momentum0.9)Adam 快SGD 穩(wěn)初始學(xué)習(xí)率1e-3finetune從零訓(xùn)練則用 1e-2學(xué)習(xí)率調(diào)度CosineAnnealing后半程微調(diào)更細(xì)膩權(quán)重初始化ImageNet 預(yù)訓(xùn)練強(qiáng)烈建議別從零訓(xùn)跑完 30 個(gè) epoch驗(yàn)證集準(zhǔn)確率一般能到 90% 上下。如果沒到先別急著換模型回到 2.2 節(jié)的可視化腳本看看是不是有臟圖或者去讀一下避坑章里相似類的問題。4. 接到 yolov5 / yolov8 分類從 train 切 val 到一條命令跑起來4.1 yolov5 classify 的數(shù)據(jù)組織與類別映射yolov5 自帶的分類任務(wù)classify/train.py和 yolov8 的yolov8n-cls.pt都遵循同一個(gè)約定數(shù)據(jù)集根目錄下必須有兩個(gè)子目錄train/和val/每個(gè)子目錄內(nèi)按「每類一個(gè)文件夾」組織。注意它找的是val而不是test。這份資源給的是train/和test/所以不能直接開訓(xùn)要先做一步目錄對齊。最省事但不推薦的做法是把test/改名為val/直接訓(xùn)練。問題是這樣最終評估就只能用同一個(gè) val指標(biāo)虛高。我習(xí)慣的做法是從train/里每類抽 10% 出來獨(dú)立成val/把原始test/留作訓(xùn)練結(jié)束后的最終盲測。import os import glob import random import shutil src data/train dst_val data/val_for_yolo os.makedirs(dst_val, exist_okTrue) val_ratio 0.1 random.seed(42) for class_dir in os.listdir(src): class_path os.path.join(src, class_dir) if not os.path.isdir(class_path): continue images glob.glob(os.path.join(class_path, *.jpg)) \ glob.glob(os.path.join(class_path, *.jpeg)) random.shuffle(images) n_val max(1, int(len(images) * val_ratio)) os.makedirs(os.path.join(dst_val, class_dir), exist_okTrue) for img in images[:n_val]: shutil.move(img, os.path.join(dst_val, class_dir, os.path.basename(img)))邏輯說明對 train 下每個(gè)類別目錄先 glob 出所有圖片按固定比例抽 10%用shutil.move從 train 挪到新的 val 目錄。random.seed(42)保證每次運(yùn)行切分結(jié)果一致這樣不同輪次的實(shí)驗(yàn)對比才公平。切分后 train 少了 10% 的圖總量從 9600 變成約 8640對訓(xùn)練影響不大。如果你用的是 yolov8 的ultralytics包數(shù)據(jù)組織方式完全一樣也是根目錄下train/和val/兩個(gè)子目錄類別由文件夾名自動(dòng)推斷。所以做完這步切分兩份框架都能直接吃這份數(shù)據(jù)。4.2 訓(xùn)練命令與本地推理yolov5 / yolov8 訓(xùn)練自己的蘑菇分類模型目錄切好后訓(xùn)練命令很短。yolov5 的 classify 分支是這樣跑的# 進(jìn)入 yolov5 倉庫根目錄 python classify/train.py \ --model resnet34 \ --data /path/to/mushroom_dataset \ --epochs 50 \ --img 224 \ --batch 32 \ --device 0--model指定 backbone除了 resnet34 還可以換resnet18跑更快、efficientnet_b0跑更省內(nèi)存--data指向包含train/和val/的根目錄--img 224和--batch 32跟 PyTorch 那套一致。跑完后權(quán)重存在runs/train-cls/exp/weights/best.pt。yolov8 的寫法更簡潔用的是命令行直接指定數(shù)據(jù)集路徑y(tǒng)olo classify train \ modelyolov8n-cls.pt \ data/path/to/mushroom_dataset \ epochs50 \ imgsz224 \ batch32 \ device0modelyolov8n-cls.pt會(huì)從官方源自動(dòng)下載預(yù)訓(xùn)練分類權(quán)重data指向同樣的根目錄。yolov8 的 n 模型非常輕顯存占用不到 2GBCPU 也能跑但精度上限比 resnet34 低一些適合先驗(yàn)證流程通不通。推理也簡單yolov5 用classify/predict.pypython classify/predict.py \ --weights runs/train-cls/exp/weights/best.pt \ --source /path/to/test/Amanita/xxx.jpg輸出會(huì)打印這張圖屬于每個(gè)類別的置信度。批量驗(yàn)證整個(gè) test 目錄也可以把--source直接指向data/test它會(huì)遞歸遍歷所有子目錄并輸出每張圖的預(yù)測結(jié)果。yolov8 則是yolo classify predict modelruns/train-cls/exp/weights/best.pt source/path/to/test到這里你已經(jīng)拿到了一個(gè)能跑通全流程的分類模型。剩下的問題不是模型架構(gòu)行不行而是有沒有踩中數(shù)據(jù)本身的暗坑。5. 避坑8 類還是 12 類、標(biāo)簽錯(cuò)位與類別不均衡5.1 類別數(shù)量標(biāo)題寫 12json 里只有 8現(xiàn)象資源標(biāo)題寫著「12種蘑菇圖像識(shí)別數(shù)據(jù)集」但classes.json打開只有 8 個(gè)類別訓(xùn)練腳本里num_classes填幾都要猶豫半天。原因12 是原始采集圖片的批次口徑。項(xiàng)目正文的文件名里能看到 Lactarius、Pluteus、Entoloma、Cortinarius 等多個(gè)屬名采集階段可能按 12 個(gè)批次或 12 個(gè)來源歸檔但整理成發(fā)布版時(shí)合并成了 8 個(gè)可判別類別。摘要里明確寫了「分類個(gè)數(shù)8」所以 12 和 8 并不矛盾只是統(tǒng)計(jì)口徑不同。解決一切以classes.json為準(zhǔn)。寫訓(xùn)練腳本前先用len(class_idx)確認(rèn)類別數(shù)再回去核對ImageFolder.classes的長度是否一致。任何地方出現(xiàn)num_classes 12都是錯(cuò)的跑完訓(xùn)練再去排查標(biāo)簽錯(cuò)位就晚了。5.2 標(biāo)簽順序ImageFolder 的字母序和 json 編號(hào)對不上現(xiàn)象訓(xùn)練時(shí) loss 正常下降但驗(yàn)證集的 top-1 準(zhǔn)確率始終在 30% 上下震蕩怎么看都像隨機(jī)猜。原因ImageFolder按文件夾名的字母序自動(dòng)生成標(biāo)簽比如Agaricus排 0、Amanita排 1、Boletus排 2而classes.json里的編號(hào)可能是按采集順序或人工指定順序?qū)懙摹煞蓓樞蛞坏┎灰恢录虞d進(jìn)來的 label 含義就錯(cuò)位了模型學(xué)到的映射關(guān)系全是亂的。解決訓(xùn)練前做一個(gè)硬性校驗(yàn)不通過就不開跑import json from torchvision import datasets with open(data/classes.json, r, encodingutf-8) as f: class_idx json.load(f) dataset datasets.ImageFolder(data/train) # 方式一直接對比文件夾順序和 json 順序 print(ImageFolder 順序:, dataset.classes) print(json 順序:, list(class_idx.keys())) # 方式二斷言兩者必須一致不一致就拋異常 assert list(dataset.classes) list(class_idx.keys()), \ 類別順序不一致需要統(tǒng)一后重新生成 json如果斷言失敗解決辦法是把classes.json重新按dataset.classes的順序生成一遍而不是去改文件夾名。文件夾名是給訓(xùn)練框架看的json 只是給人看的輔助文件以文件夾名為準(zhǔn)重建 json兩邊就對齊了。5.3 類別不均衡9600 張攤到 8 類未必均勻現(xiàn)象訓(xùn)練集總共 9600 張但訓(xùn)練時(shí)發(fā)現(xiàn)某些類別的 loss 一直下不去混淆矩陣?yán)飩€(gè)別類 recall 特別低。原因9600 是總數(shù)不代表每類正好 1200 張。蘑菇采集本身受季節(jié)、地域影響很大某些常見種可能占了 3000 張稀有類別可能只有 500 張。類別不均衡會(huì)讓模型偏向樣本多的類。解決訓(xùn)練前先做一個(gè)快速統(tǒng)計(jì)用numpy數(shù)一下每類樣本數(shù)import numpy as np from torch.utils.data import Dataset # 用 ImageFolder 的 targets 屬性直接統(tǒng)計(jì) counts np.bincount(dataset.targets) for cls, count in zip(dataset.classes, counts): print(f{cls}: {count} 張)如果發(fā)現(xiàn)差距超過 2 倍就用WeightedRandomSampler做樣本加權(quán)采樣給少樣本的類更高的采樣概率from torch.utils.data import WeightedRandomSampler class_counts np.bincount(dataset.targets) weights 1.0 / class_counts[dataset.targets] sampler WeightedRandomSampler(weights, num_sampleslen(dataset), replacementTrue) train_loader DataLoader(train_ds, batch_size32, samplersampler)采樣器會(huì)讓每輪 epoch 里少樣本類別被抽到的次數(shù)顯著增加緩解偏向。代價(jià)是每個(gè) epoch 實(shí)際上看到重復(fù)樣本訓(xùn)練輪數(shù)不用加太多30 輪以內(nèi)足夠。5.4 沒有驗(yàn)證集反復(fù)用 test 調(diào)參會(huì)虛高現(xiàn)象訓(xùn)練時(shí)一直拿data/test當(dāng)驗(yàn)證集看準(zhǔn)確率調(diào)了幾輪超參后測試集準(zhǔn)確率到了 95%但換成真實(shí)場景的新圖準(zhǔn)確率掉到 80%完全沒法解釋。原因test 集被反復(fù)用于超參選擇和 early stopping它的信息已經(jīng)在訓(xùn)練過程中泄漏給模型了。測試集的有效性在于「只用一次」反復(fù)用它調(diào)參它就不再是獨(dú)立評估集結(jié)果虛高是必然的。這是分類實(shí)驗(yàn)里最常見的翻車點(diǎn)。解決嚴(yán)格三集分離。train 訓(xùn)練、val從 train 里切 10%調(diào)參、test 只做最終評估。第 4 章給的切分腳本就是專門干這個(gè)的訓(xùn)完所有實(shí)驗(yàn)后用test/跑最后一次推理這個(gè)數(shù)字才是能寫進(jìn)報(bào)告的結(jié)果。5.5 相似類難分Amanita 和 Entoloma 分不清是正常的現(xiàn)象混淆矩陣?yán)顰manita和Entoloma兩個(gè)類互認(rèn)錯(cuò)訓(xùn)練初期 loss 下降比其他類慢很多換大模型效果也沒明顯改善。原因這兩個(gè)屬的蘑菇都是典型傘菌形態(tài)菌蓋顏色、菌褶走向在視覺上非常接近人眼都容易認(rèn)錯(cuò)。類間相似度高是蘑菇分類數(shù)據(jù)集的固有難點(diǎn)不是模型或者代碼的 bug。文件名里出現(xiàn)的Entoloma_original、Lactarius_original這類字樣也能看出來原始圖里大量是野外同一生長環(huán)境拍的背景干擾大。解決先接受基線結(jié)果再針對性地做三件事。一是輸入分辨率從 224 提到 320讓模型看到更多菌蓋細(xì)節(jié)二是對易混類做定向數(shù)據(jù)增強(qiáng)比如更強(qiáng)的隨機(jī)旋轉(zhuǎn)和光照擾動(dòng)三是在驗(yàn)證時(shí)看 top-2 準(zhǔn)確率對這類相似類任務(wù)top-2 比 top-1 更有實(shí)際參考價(jià)值。別指望把這兩個(gè)類的區(qū)分做到 100%能做到 90% 已經(jīng)超過多數(shù)人工標(biāo)注水平。6. 進(jìn)階用混淆矩陣和 Top-2 準(zhǔn)確率給模型做一次體檢訓(xùn)練完的模型光看一個(gè)總準(zhǔn)確率是不夠的。蘑菇分類這種類間相似度高的任務(wù)真正能說明問題的是混淆矩陣哪兩個(gè)類容易被搞混、每個(gè)類的召回率是多少都在這一張圖里。下面這段代碼用 sklearn 生成混淆矩陣和每類的 precision / recall是我每次跑完分類任務(wù)都會(huì)執(zhí)行的固定動(dòng)作import torch import numpy as np from sklearn.metrics import confusion_matrix, classification_report import seaborn as sns import matplotlib.pyplot as plt model.eval() y_true, y_pred [], [] with torch.no_grad(): for images, labels in val_loader: images images.to(device) outputs model(images) _, predicted torch.max(outputs, 1) y_true labels.tolist() y_pred predicted.cpu().tolist() # 每類的 precision / recall / f1 print(classification_report(y_true, y_pred, target_namesval_ds.dataset.classes)) # 混淆矩陣可視化 cm confusion_matrix(y_true, y_pred) plt.figure(figsize(8, 8)) sns.heatmap(cm, annotTrue, fmtd, cmapBlues, xticklabelsval_ds.dataset.classes, yticklabelsval_ds.dataset.classes) plt.xlabel(Predicted) plt.ylabel(True) plt.show()classification_report輸出每個(gè)類別的 precision、recall、f1-score哪個(gè)類是短板一眼就能看出來熱力圖上對角線以外的亮點(diǎn)就是模型最容易翻車的類別對。看到Amanita和Entoloma互串再去針對性調(diào)數(shù)據(jù)增強(qiáng)比盲試網(wǎng)絡(luò)結(jié)構(gòu)有效得多。Top-2 準(zhǔn)確率在蘑菇分類里比 top-1 更貼近實(shí)際使用場景。eDNA 調(diào)查或者野外采集輔助識(shí)別這類應(yīng)用通常給出「最可能的兩種蘑菇」讓用戶確認(rèn)比硬要模型賭一個(gè)答案更合理。統(tǒng)計(jì) top-2 的代碼也不復(fù)雜correct_top2 0 total 0 with torch.no_grad(): for images, labels in val_loader: images images.to(device) outputs model(images) top2 torch.topk(outputs, 2, dim1).indices.cpu() # 取前兩個(gè)預(yù)測 for i, label in enumerate(labels): if label.item() in top2[i].tolist(): correct_top2 1 total 1 print(fTop-2 Accuracy: {100.0 * correct_top2 / total:.2f}%)這段代碼的要點(diǎn)是torch.topk(outputs, 2, dim1)取的是每個(gè)樣本 logits 最高的兩個(gè)類別索引只要真實(shí)標(biāo)簽落在里面就算對。在相似類多的數(shù)據(jù)上top-2 一般會(huì)比 top-1 高 5 到 10 個(gè)百分點(diǎn)這個(gè)差距本身就是數(shù)據(jù)難度的直觀體現(xiàn)。從那以后我拿到任何分類數(shù)據(jù)集第一件事永遠(yuǎn)是三件套數(shù) json 長度、跑 bincount、打印 9 宮格圖這三步走完后面基本不會(huì)翻車。希望幫到你。本文還有配套的精品資源點(diǎn)擊獲取