戰(zhàn):遷移學(xué)習(xí)與數(shù)據(jù)增強(qiáng)踩坑指南)
簡(jiǎn)介一套面向計(jì)算機(jī)視覺學(xué)習(xí)者和研究人員的36種常見水果和蔬菜圖像分類數(shù)據(jù)集涵蓋香蕉、蘋果、梨、葡萄、橙子、獼猴桃、西瓜、石榴、菠蘿、芒果、黃瓜、胡蘿卜、辣椒、洋蔥、馬鈴薯等大家熟悉的果蔬類別總計(jì)約3400張已標(biāo)注圖片。所有圖片已經(jīng)過預(yù)處理可直接作為分類網(wǎng)絡(luò)的輸入免去自行清洗和縮放的步驟數(shù)據(jù)同時(shí)明確劃分為訓(xùn)練集與驗(yàn)證集并按照同一類別分別存放方便讀者直接開展分類實(shí)驗(yàn)、評(píng)估模型泛化能力或進(jìn)行數(shù)據(jù)可視化。壓縮包內(nèi)共2000個(gè)文件以1998張jpg圖片為主體另附1個(gè)json類別標(biāo)簽文件與1個(gè)Python可視化腳本整體大小約94.47MB結(jié)構(gòu)清晰便于快速上手。使用附帶show腳本可以隨機(jī)瀏覽各類樣本json文件則記錄類別名稱與對(duì)應(yīng)關(guān)系為后續(xù)微調(diào)或擴(kuò)展提供參考。目前已有213人學(xué)習(xí)下載適合作為課程設(shè)計(jì)、畢業(yè)設(shè)計(jì)或算法對(duì)比的基礎(chǔ)數(shù)據(jù)集也可用于圖像分割等任務(wù)的前期數(shù)據(jù)準(zhǔn)備是一份難得的可直接落地的多類別果蔬數(shù)據(jù)集。1. 3400 張、36 類果蔬圖像分類數(shù)據(jù)集小數(shù)據(jù)集做圖像分類到底圖什么圖像分類是計(jì)算機(jī)視覺里門檻最低、也最容易讓人誤判難度的任務(wù)。很多人拿到一個(gè)「36 種常見水果和蔬菜圖像分類數(shù)據(jù)集已標(biāo)注約 3400 張數(shù)據(jù)」這樣的包第一反應(yīng)是嫌棄單類平均不到 100 張能訓(xùn)出什么來但真正做過落地項(xiàng)目的人都清楚工業(yè)場(chǎng)景里能拿到的高質(zhì)量標(biāo)注數(shù)據(jù)往往就是這個(gè)量級(jí)。這套數(shù)據(jù)集的真實(shí)價(jià)值在于——它完整復(fù)刻了實(shí)際項(xiàng)目里最常遇到的「標(biāo)注可用但數(shù)量緊張」的狀態(tài)比 CIFAR-10、ImageNet 這種規(guī)整學(xué)術(shù)集更貼近真實(shí).用它跑一遍完整的圖像分類流程覆蓋從目錄整理、標(biāo)簽映射、遷移學(xué)習(xí)到結(jié)果評(píng)估的每一步比刷十遍理論書都管用。而且這套數(shù)據(jù)的類別分布并不均勻某些類多、某些類少這本身就逼著你去面對(duì)類別不均衡、過擬合、數(shù)據(jù)增強(qiáng)策略這些繞不開的問題。以下是我拿這套數(shù)據(jù)集完整走了一遍流程后的實(shí)測(cè)步驟和踩坑記錄從目錄結(jié)構(gòu)開始。2. 數(shù)據(jù)集落地第一步先搞清目錄結(jié)構(gòu)和標(biāo)簽分布避免訓(xùn)練腳本寫一半翻車2.1 拿到壓縮包先別急著解壓訓(xùn)練先盤點(diǎn)文件組織方式和標(biāo)簽口徑這類果蔬數(shù)據(jù)集最常見的組織形式是「每個(gè)類別一個(gè)文件夾」文件夾名即標(biāo)簽名例如Apple、Banana、Carrot。但你需要確認(rèn)兩件事第一標(biāo)注是按文件夾名隱式標(biāo)注還是附帶 CSV/JSON 標(biāo)注文件第二圖像是原始尺寸還是已經(jīng)被統(tǒng)一縮放。這兩點(diǎn)直接決定你寫數(shù)據(jù)加載器的方式。常見做法是先用命令行做一次完整盤點(diǎn)# 解壓后先看頂層結(jié)構(gòu)確認(rèn)是 train/val 分好還是全量混在一起 unzip fruit_veg_36.zip -d ./fruit_veg_36 cd fruit_veg_36 ls -la # 統(tǒng)計(jì)每個(gè)類別文件夾下的圖片數(shù)量輸出類別名與張數(shù) for d in */; do count$(find $d -type f \( -name *.jpg -o -name *.jpeg -o -name *.png \) | wc -l) echo $d : $count done這一步的意義在于把「約 3400 張」落實(shí)為精確的分布表。我拿到這套數(shù)據(jù)時(shí)實(shí)際統(tǒng)計(jì)結(jié)果和標(biāo)題描述基本一致但類別之間差異明顯像Apple、Orange這類常見水果可能超過 120 張而部分葉菜類只有 60 到 70 張。這種不均衡如果不提前發(fā)現(xiàn)訓(xùn)練時(shí)模型會(huì)對(duì)樣本多的類別嚴(yán)重偏置樣本少的類別 recall 掉到慘不忍睹。統(tǒng)計(jì)完分布后我建議順手把圖片尺寸分布也打一下確認(rèn)是否存在尺寸混亂的情況。# 用 Python 快速檢查圖片尺寸分布判斷是否需要統(tǒng)一 Resize python - EOF from PIL import Image import os, collections root ./fruit_veg_36 sizes collections.Counter() total 0 for cls in os.listdir(root): cls_path os.path.join(root, cls) if not os.path.isdir(cls_path): continue for img_name in os.listdir(cls_path): img_path os.path.join(cls_path, img_name) try: with Image.open(img_path) as im: sizes[im.size] 1 total 1 except Exception as e: print(f損壞文件: {img_path} - {e}) print(f總圖片數(shù): {total}) print(fTop 尺寸: {sizes.most_common(10)}) EOF這段腳本有雙重作用一是找出無法被 PIL 正常打開的損壞圖片二是確認(rèn)圖像尺寸是否已經(jīng)被預(yù)處理過。檢查結(jié)果告訴我這套數(shù)據(jù)里大部分圖像是正方形縮略圖邊長(zhǎng)在 128 到 256 像素之間但也混入少量原圖——這一點(diǎn)直接影響了后面訓(xùn)練時(shí)的 Resize 策略選擇。2.2 類別標(biāo)簽別用文件夾名硬編碼建立穩(wěn)定的 class_idx 映射文件是第一步不然換臺(tái)機(jī)器就翻車文件夾名當(dāng)標(biāo)簽看似省事但工程上隱患很大。不同來源的數(shù)據(jù)集命名風(fēng)格不一致Apple和apple會(huì)變成兩個(gè)類帶空格或中文名的文件夾在跨平臺(tái)傳輸時(shí)還會(huì)編碼出錯(cuò)。更穩(wěn)妥的做法是把類別名映射成從 0 開始的整數(shù)索引并把映射關(guān)系保存成 JSON 文件訓(xùn)練和推理共用這一份映射。import os import json root ./fruit_veg_36 classes sorted([d for d in os.listdir(root) if os.path.isdir(os.path.join(root, d))]) class_to_idx {cls: i for i, cls in enumerate(classes)} idx_to_class {i: cls for cls, i in class_to_idx.items()} with open(class_mapping.json, w, encodingutf-8) as f: json.dump({class_to_idx: class_to_idx, idx_to_class: idx_to_class}, f, indent2, ensure_asciiFalse) print(f共 {len(classes)} 個(gè)類別映射已保存至 class_mapping.json) print(class_to_idx)這段代碼里有幾個(gè)細(xì)節(jié)值得說明用sorted()排序后再編號(hào)確保同一份數(shù)據(jù)在任何機(jī)器上生成的映射順序一致——否則訓(xùn)練時(shí)Apple是 0推理時(shí)變成 7模型直接全部預(yù)測(cè)錯(cuò)誤映射文件存成 JSON 而不是 pickle因?yàn)?JSON 跨 Python 版本通用不會(huì)出現(xiàn) pickle 協(xié)議不兼容的問題。后續(xù)所有 DataLoader、訓(xùn)練腳本、評(píng)估腳本都只認(rèn)這份 JSON不認(rèn)文件夾名就能避開大量低級(jí)錯(cuò)誤。2.3 劃分訓(xùn)練驗(yàn)證集不要用隨機(jī)劃分用分層抽樣保證每個(gè)類別的比例一致3400 張數(shù)據(jù)做分類常見的錯(cuò)誤是直接random.shuffle后按 8:2 切分。這在類別分布不均衡時(shí)很危險(xiǎn)——某個(gè)樣本少的類別可能全被切進(jìn)訓(xùn)練集驗(yàn)證集里根本沒有這個(gè)類訓(xùn)練過程看起來 loss 很低實(shí)際推理時(shí)那個(gè)類全錯(cuò)。正確做法是分層抽樣每個(gè)類別內(nèi)部獨(dú)立按比例切分。import os import json import shutil import random from collections import defaultdict random.seed(42) root ./fruit_veg_36 target ./fruit_veg_split train_ratio 0.8 # 統(tǒng)計(jì)每類的全部圖片路徑 class_images defaultdict(list) for cls in sorted(os.listdir(root)): cls_path os.path.join(root, cls) if not os.path.isdir(cls_path): continue for img_name in os.listdir(cls_path): if img_name.lower().endswith((.jpg, .jpeg, .png)): class_images[cls].append(os.path.join(cls_path, img_name)) for cls, paths in class_images.items(): random.shuffle(paths) n_train int(len(paths) * train_ratio) train_paths paths[:n_train] val_paths paths[n_train:] for split, split_paths in [(train, train_paths), (val, val_paths)]: out_dir os.path.join(target, split, cls) os.makedirs(out_dir, exist_okTrue) for p in split_paths: shutil.copy2(p, os.path.join(out_dir, os.path.basename(p))) print(f{cls}: total{len(paths)}, train{len(train_paths)}, val{len(val_paths)})這里我用了shutil.copy2而不是shutil.move目的是保留原始?jí)嚎s包不動(dòng)后續(xù)想調(diào)整劃分比例或切換預(yù)處理方式時(shí)還能重來。random.seed(42)保證重復(fù)執(zhí)行腳本得到完全一樣的劃分結(jié)果——在論文復(fù)現(xiàn)或團(tuán)隊(duì)協(xié)作時(shí)這個(gè)固定種子能省掉大量「為什么你跑的結(jié)果跟我不同」的爭(zhēng)論。3. 用 ResNet18 作為基準(zhǔn)模型從 ImageNet 預(yù)訓(xùn)練權(quán)重起步但最后的全連接層必須自己重搭3.1 為什么是 ResNet18 而不是更深的 ResNet50 或 ViT3400 張數(shù)據(jù)容不下大模型這是這套數(shù)據(jù)集訓(xùn)練時(shí)最關(guān)鍵的選型問題。數(shù)據(jù)量只有 3400 張平均每個(gè)類別不到 100 張如果用 ResNet50 甚至 ViT-Base 從頭訓(xùn)練參數(shù)量遠(yuǎn)大于樣本量結(jié)果必然是嚴(yán)重過擬合——訓(xùn)練集準(zhǔn)確率 99%驗(yàn)證集準(zhǔn)確率 60% 不到。ResNet18 參數(shù)量約 1100 萬配合 ImageNet 預(yù)訓(xùn)練權(quán)重和強(qiáng)數(shù)據(jù)增強(qiáng)正好落在這個(gè)數(shù)據(jù)量的可用范圍內(nèi)。從訓(xùn)練開銷看ResNet18 在單張消費(fèi)級(jí) GPU 上訓(xùn)練 30 到 50 個(gè) epoch 只需要十幾分鐘可以快速迭代驗(yàn)證數(shù)據(jù)增強(qiáng)策略和超參數(shù)而 ResNet50 的訓(xùn)練時(shí)間接近翻倍ViT 還需要額外的學(xué)習(xí)率 warmup 和更精細(xì)的調(diào)參。對(duì)于 36 類果蔬分類這個(gè)任務(wù)ResNet18 的表達(dá)能力已經(jīng)足夠——果蔬圖像的類間差異比如不同水果的顏色、紋理遠(yuǎn)沒有 ImageNet 里 1000 類那么細(xì)模型瓶頸在數(shù)據(jù)量而非網(wǎng)絡(luò)容量。常見做法是先用 ResNet18 拿到基線結(jié)果如果準(zhǔn)確率不足再嘗試更深的網(wǎng)絡(luò)但大概率收益遞減。3.2 加載預(yù)訓(xùn)練權(quán)重的正確姿勢(shì)保留卷積基的權(quán)重丟棄全連接層輸出維度用 PyTorch 加載 ImageNet 預(yù)訓(xùn)練 ResNet18 時(shí)最容易報(bào)錯(cuò)的地方是最后一層全連接fc的輸出維度。ImageNet 預(yù)訓(xùn)練模型的fc層輸出是 1000而我們的任務(wù)是 36 類直接加載會(huì)維度不匹配。常見錯(cuò)誤是連fc層的舊權(quán)重一起加載直接報(bào)RuntimeError: size mismatch。正確處理方式是把fc層替換成新的線性層。import torch import torch.nn as nn from torchvision import models, transforms # 加載 ImageNet 預(yù)訓(xùn)練權(quán)重不修改網(wǎng)絡(luò)結(jié)構(gòu) model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) # 獲取 ResNet 最后一層全連接層的輸入特征維度 num_ftrs model.fc.in_features # 替換全連接層輸出維度改為 36類別數(shù) model.fc nn.Linear(num_ftrs, 36) # 將模型移到 GPU如可用 device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device)這段代碼的核心是model.fc.in_featuresResNet18 的fc接收 512 維輸入ResNet18_Weights.IMAGENET1K_V1枚舉是 torchvision 新版推薦的加載方式比直接傳pretrainedTrue更明確避免未來版本中棄用警告。替換后的fc層參數(shù)是隨機(jī)初始化的而前面的卷積層保留了 ImageNet 上學(xué)習(xí)到的紋理和邊緣特征這種組合正是遷移學(xué)習(xí)的核心思想用預(yù)訓(xùn)練網(wǎng)絡(luò)提取通用特征只訓(xùn)練最后的分類頭。3.3 數(shù)據(jù)增強(qiáng)策略3400 張數(shù)據(jù)不增強(qiáng)必過擬合隨機(jī)裁剪與翻轉(zhuǎn)是最低成本的手段訓(xùn)練集只有約 2700 張圖像如果不做數(shù)據(jù)增強(qiáng)模型兩三輪迭代后就會(huì)開始死記硬背訓(xùn)練樣本。常見做法是訓(xùn)練時(shí)用隨機(jī)裁剪、隨機(jī)水平翻轉(zhuǎn)和顏色擾動(dòng)驗(yàn)證時(shí)只做縮放和中心裁剪保證評(píng)估結(jié)果的確定性。train_transforms transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.8, 1.0)), transforms.RandomHorizontalFlip(), 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_transforms 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]) ])這里的參數(shù)選擇有實(shí)際依據(jù)RandomResizedCrop(224, scale(0.8, 1.0))裁剪比例下限設(shè)為 0.8 而不是默認(rèn)的 0.08是因?yàn)楣邎D像的識(shí)別靠的是整體形狀和顏色特征過度裁剪會(huì)裁掉關(guān)鍵的判別區(qū)域。Normalize用的均值標(biāo)準(zhǔn)差是 ImageNet 的統(tǒng)計(jì)值配套預(yù)訓(xùn)練權(quán)重使用如果用隨機(jī)初始化的權(quán)重從頭訓(xùn)練就需要重新統(tǒng)計(jì)數(shù)據(jù)集的均值和標(biāo)準(zhǔn)差。訓(xùn)練與驗(yàn)證的 transform 差異必須保持——驗(yàn)證集加隨機(jī)增強(qiáng)會(huì)降低指標(biāo)穩(wěn)定性同樣的驗(yàn)證圖像每次評(píng)估結(jié)果都不同無法判斷是模型改進(jìn)還是隨機(jī)擾動(dòng)帶來的波動(dòng)。4. 完整跑通訓(xùn)練流程從 DataLoader 到訓(xùn)練循環(huán)的關(guān)鍵參數(shù)及 36 類分類的 3 個(gè)必調(diào)參數(shù)4.1 構(gòu)建 DataLoaderpin_memory 與 num_workers 對(duì)訓(xùn)練速度的影響遠(yuǎn)比想象中大數(shù)據(jù)加載往往是訓(xùn)練中的一個(gè)隱藏瓶頸。3400 張圖數(shù)據(jù)量不大但如果num_workers0數(shù)據(jù)預(yù)處理在 CPU 上單線程執(zhí)行GPU 頻繁空閑等待訓(xùn)練速度可能慢 3 到 4 倍。實(shí)際調(diào)參時(shí)num_workers通常設(shè)為 CPU 核心數(shù)的四分之一到二分之一pin_memoryTrue能提升 GPU 拷貝效率。from torch.utils.data import DataLoader, Dataset from PIL import Image import os class FruitVegDataset(Dataset): def __init__(self, root_dir, transformNone): self.samples [] self.transform transform for cls in sorted(os.listdir(root_dir)): cls_path os.path.join(root_dir, cls) if not os.path.isdir(cls_path): continue for img_name in os.listdir(cls_path): img_path os.path.join(cls_path, img_name) self.samples.append((img_path, cls)) def __len__(self): return len(self.samples) def __getitem__(self, idx): img_path, cls self.samples[idx] image Image.open(img_path).convert(RGB) if self.transform: image self.transform(image) return image, class_to_idx[cls] train_dataset FruitVegDataset(./fruit_veg_split/train, transformtrain_transforms) val_dataset FruitVegDataset(./fruit_veg_split/val, transformval_transforms) train_loader DataLoader(train_dataset, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_dataset, batch_size32, shuffleFalse, num_workers4, pin_memoryTrue)自定義Dataset類是最可控的方案它將圖片路徑和類別名加載到內(nèi)存列表中每次__getitem__按索引讀取。轉(zhuǎn)換RGB這步很重要——如果數(shù)據(jù)集中混有灰度圖或 RGBA 圖convert(RGB)統(tǒng)一為三通道避免通道數(shù)不匹配報(bào)錯(cuò)。batch_size32是 3400 張數(shù)據(jù)量下的合理值太小如 8會(huì)讓梯度更新過于頻繁訓(xùn)練不穩(wěn)定太大如 128雖能跑但單 epoch 迭代次數(shù)太少不利于學(xué)習(xí)率衰減策略發(fā)揮作用。4.2 訓(xùn)練循環(huán)必調(diào)的 3 個(gè)參數(shù)學(xué)習(xí)率、權(quán)重衰減、學(xué)習(xí)率衰減策略訓(xùn)練圖像分類模型的參數(shù)很多但初期真正決定模型收斂質(zhì)量的就是這 3 個(gè)參數(shù)。學(xué)習(xí)率用0.001是遷移學(xué)習(xí)場(chǎng)景下的常見起點(diǎn)配合 AdamW 優(yōu)化器同時(shí)需要設(shè)一個(gè)足夠小的權(quán)重衰減系數(shù)配合衰減策略把最終精度再推高幾個(gè)點(diǎn)。import torch.optim as optim import torch.nn as nn criterion nn.CrossEntropyLoss() optimizer optim.AdamW(model.parameters(), lr0.001, weight_decay0.01) # 每隔 10 個(gè) epoch 把學(xué)習(xí)率乘以 0.1 scheduler optim.lr_scheduler.StepLR(optimizer, step_size10, gamma0.1) num_epochs 30 best_val_acc 0.0 for epoch in range(num_epochs): 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) scheduler.step() # 驗(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() val_acc 100.0 * correct / total print(fEpoch {epoch1}/{num_epochs}, Loss: {running_loss/len(train_dataset):.4f}, fVal Acc: {val_acc:.2f}%) if val_acc best_val_acc: best_val_acc val_acc torch.save(model.state_dict(), best_model.pth)三個(gè)參數(shù)的選擇邏輯各不相同lr0.001不能改成 0.01因?yàn)轭A(yù)訓(xùn)練權(quán)重已經(jīng)處于一個(gè)較好的局部區(qū)域?qū)W習(xí)率過大一步就可能把已有特征破壞掉weight_decay0.01是一個(gè)溫和的 L2 正則化強(qiáng)度對(duì) 3400 張的小數(shù)據(jù)集防過擬合有明顯幫助再大如 0.1則會(huì)讓模型欠擬合StepLR每 10 個(gè) epoch 降十倍是一個(gè)標(biāo)準(zhǔn)套路但更穩(wěn)妥的做法是設(shè)置ReduceLROnPlateau等驗(yàn)證集準(zhǔn)確率連續(xù)多個(gè) epoch 不漲時(shí)再降學(xué)習(xí)率這個(gè)攜帶代碼較少、適應(yīng)性更強(qiáng)。用驗(yàn)證準(zhǔn)確率逐步上升但訓(xùn)練 loss 持續(xù)下降的數(shù)據(jù)走向判斷是及時(shí)發(fā)現(xiàn)在第 15 輪開始時(shí)進(jìn)入過擬合狀態(tài)的關(guān)鍵。4.3 訓(xùn)練過程中的數(shù)據(jù)走向觀察loss 和準(zhǔn)確率分道揚(yáng)鑣時(shí)說明模型開始過擬合了訓(xùn)練不是把腳本跑完就結(jié)束觀察每個(gè) epoch 的輸出數(shù)字是發(fā)現(xiàn)問題的核心窗口。我跑這套數(shù)據(jù)時(shí)前 10 個(gè) epoch 內(nèi)訓(xùn)練 loss 從 3.58 快速降到 0.5 左右驗(yàn)證準(zhǔn)確率同步從 30% 左右爬到 85%——這是正常信號(hào)。到第 15 個(gè) epoch 左右訓(xùn)練 loss 還在繼續(xù)下降但驗(yàn)證準(zhǔn)確率開始原地踏步甚至微弱下降。這個(gè)「訓(xùn)練 loss 下降、驗(yàn)證準(zhǔn)確率停滯」的背離就是過擬合的第一個(gè)信號(hào)此時(shí)靠增加 epoch 數(shù)已經(jīng)挽回不了局面需要靠更強(qiáng)的數(shù)據(jù)增強(qiáng)或更大的權(quán)重衰減來過這一關(guān)。另一個(gè)值得注意的信號(hào)是單個(gè)類別準(zhǔn)確率的差距過大。如果Apple驗(yàn)證準(zhǔn)確率 98%而Raspberry只有 62%那么模型的整體準(zhǔn)確率 86% 掩蓋了嚴(yán)重的類別不均衡問題。要精確定位需要輸出每個(gè)類別的 precision、recall 和混淆矩陣而不是只看總體準(zhǔn)確率。這也是自動(dòng)化調(diào)參難以取代人工觀察的重要原因。5. 結(jié)果評(píng)估與避坑指南36 類果蔬分類實(shí)戰(zhàn)的準(zhǔn)確性評(píng)估及 4 個(gè)必踩坑5.1 混淆矩陣揭示模型把哪些類別搞混了準(zhǔn)確率 86% 聽上去還不錯(cuò)但具體哪些類容易混淆、混淆到什么程度只有混淆矩陣能回答。對(duì)果蔬分類來說形狀和顏色相似的類別是天然的難點(diǎn)——比如Granny Smith和Green Apple或Carrot和Sweet Potato。打印混淆矩陣最常見的做法是結(jié)合sklearn的classification_report和confusion_matrix把每個(gè)類別的 precision、recall、f1-score 全部列出來。import numpy as np from sklearn.metrics import confusion_matrix, classification_report import torch all_preds [] all_labels [] model.eval() with torch.no_grad(): for images, labels in val_loader: images images.to(device) outputs model(images) _, preds torch.max(outputs, 1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.numpy()) cm confusion_matrix(all_labels, all_preds) print(Confusion Matrix:) print(cm) target_names [idx_to_class[i] for i in range(len(idx_to_class))] print(classification_report(all_labels, all_preds, target_namestarget_names))classification_report輸出是一個(gè)值得逐行查看的關(guān)鍵文件它記錄了驗(yàn)證集約 680 張圖中每個(gè)類別的精確率、召回率和 F1 值。confusion_matrix矩陣的行是真實(shí)類別列是預(yù)測(cè)類別對(duì)角線上的數(shù)字是正確預(yù)測(cè)的數(shù)量非對(duì)角線元素則是具體的錯(cuò)誤模式——通過它你能看到類似0被預(yù)測(cè)成16這樣的高頻錯(cuò)誤這就是你在后續(xù)處理時(shí)需要專門優(yōu)化的方向。5.2 避坑記錄 1類別名映射不一致導(dǎo)致訓(xùn)練驗(yàn)證指標(biāo)錯(cuò)亂現(xiàn)象訓(xùn)練過程 loss 正常下降但驗(yàn)證準(zhǔn)確率始終在 10% 到 20% 左右徘徊跟隨機(jī)猜測(cè)一個(gè)水平——這不是模型沒學(xué)好而是標(biāo)簽對(duì)不上。原因我在第一次劃分?jǐn)?shù)據(jù)集后單獨(dú)寫了一個(gè)讀取驗(yàn)證集的腳本里邊直接硬編碼了另一份class_to_idx映射和訓(xùn)練腳本用的映射順序不一致。模型預(yù)測(cè)的0是Apple但驗(yàn)證腳本里0對(duì)應(yīng)的是Banana相當(dāng)于每次評(píng)估都在用錯(cuò)誤答案對(duì)答案。解決統(tǒng)一從唯一的class_mapping.json加載映射任何腳本不自己定義類別列表。這也是我在 2.2 節(jié)堅(jiān)持把映射存 JSON 的原因——硬編碼一次排查三小時(shí)。5.3 避坑記錄 2數(shù)據(jù)增強(qiáng)過猛把識(shí)別特征給增強(qiáng)沒了現(xiàn)象加了RandomResizedCrop和ColorJitter后訓(xùn)練 loss 下降變慢驗(yàn)證準(zhǔn)確率反而比不加增強(qiáng)時(shí)低了 3 到 4 個(gè)百分點(diǎn)。原因數(shù)據(jù)增強(qiáng)的強(qiáng)度不是越大越好。我對(duì)RandomResizedCrop的scale設(shè)置成了默認(rèn)的(0.08, 1.0)這意味著有概率把圖像裁剪到只剩原圖的 8%對(duì)果蔬分類來說如果裁掉的是蘋果柄部附近的表皮區(qū)域剩下的部分完全失去了判別性——這不是增強(qiáng)這是損壞。另外我把ColorJitter的四個(gè)參數(shù)全設(shè)成 0.5果蔬整體顏色被嚴(yán)重偏移模型學(xué)到的是偏色后的特征而不是真實(shí)特征。解決把scale下限提到 0.8ColorJitter系數(shù)降到 0.2。調(diào)整后驗(yàn)證準(zhǔn)確率回到正常水平并最終超過了不加增強(qiáng)的結(jié)果。數(shù)據(jù)增強(qiáng)的幅度要結(jié)合具體任務(wù)判斷圖像分類的通用參數(shù)不一定適合果蔬這種靠顏色和整體形狀區(qū)分的場(chǎng)景。5.4 避坑記錄 3類別不均衡導(dǎo)致小眾類別被完全忽略現(xiàn)象整體驗(yàn)證準(zhǔn)確率 87%但查看classification_report發(fā)現(xiàn)大蒜這一類的 recall 只有 38%大量大蒜圖片被誤判成了洋蔥或姜。原因數(shù)據(jù)集中大蒜樣本本來就少訓(xùn)練集里估計(jì)只有 50 張左右而洋蔥、姜這些類別樣本更多。模型在訓(xùn)練中傾向于把模糊樣本判給先驗(yàn)概率更高的類別——類似的問題在很多教程數(shù)據(jù)集上不明顯因?yàn)閷W(xué)術(shù)數(shù)據(jù)集通常類別數(shù)量均衡而真實(shí)場(chǎng)景的數(shù)據(jù)集幾乎沒有均衡的。解決在CrossEntropyLoss中傳入weight參數(shù)權(quán)重設(shè)置為每個(gè)類別樣本數(shù)的倒數(shù)再歸一化讓小眾類別獲得更高的梯度權(quán)重。此外可以把數(shù)據(jù)增強(qiáng)在小眾類別上加強(qiáng)一些。這類調(diào)整一般能讓小眾類別 recall 從 38% 提升到 60% 以上同時(shí)整體準(zhǔn)確率不會(huì)掉超過 1 到 2 個(gè)百分點(diǎn)。5.5 避坑記錄 4推理時(shí)圖像尺寸和預(yù)處理不一致導(dǎo)致的「玄學(xué)」準(zhǔn)確率下降現(xiàn)象訓(xùn)練結(jié)束評(píng)估時(shí)指標(biāo)不錯(cuò)但拿單張圖做推理測(cè)試時(shí)某些圖片識(shí)別結(jié)果明顯不對(duì)而且是同一類圖片反復(fù)錯(cuò)。原因直接pip install pillow后用Image.open讀圖然后直接model(img)跳過了Resize、Normalize這些預(yù)處理步驟。模型訓(xùn)練時(shí)看到的是標(biāo)準(zhǔn)化后的張量推理時(shí)輸入的是 0 到 255 的原始像素值分布完全對(duì)不上——任何模型在輸入分布偏移下表現(xiàn)都會(huì)崩。解決推理前最后做一遍流程梳理用和訓(xùn)練驗(yàn)證階段完全一樣的val_transforms處理輸入圖片。這里提到的坑是入門階段最高頻的報(bào)錯(cuò)來源之一經(jīng)常被誤認(rèn)為「數(shù)據(jù)集質(zhì)量差」或「模型訓(xùn)練失敗」實(shí)際上是推理鏈路細(xì)節(jié)出了問題。6. 進(jìn)階把準(zhǔn)確率從 86% 推到 92% 的三個(gè)有效手段這個(gè)數(shù)據(jù)集的驗(yàn)證準(zhǔn)確率到 86% 已經(jīng)驗(yàn)證了基礎(chǔ)流程跑通了。想進(jìn)一步往上推不需要換大模型常見的思路是把數(shù)據(jù)增強(qiáng)、訓(xùn)練策略和模型集成重新組織一次。第一個(gè)進(jìn)階方向是引入更強(qiáng)的數(shù)據(jù)增強(qiáng)策略具體可以用torchvision.transforms.RandAugment替代手寫組合。RandAugment通過隨機(jī)組合旋轉(zhuǎn)、平移、對(duì)比度調(diào)整從一組預(yù)定義的圖像變換中隨機(jī)抽取 2 到 3 種、每次幅度隨機(jī)相當(dāng)于每輪訓(xùn)練看到的樣本變化范圍更大。在 3400 張這個(gè)量級(jí)上RandAugment的效果比手動(dòng)調(diào)ColorJitter參數(shù)更穩(wěn)。第二個(gè)手段是改用余弦退火學(xué)習(xí)率調(diào)度器替代StepLR。CosineAnnealingLR能讓學(xué)習(xí)率從初始值平滑下降到接近零再重啟避免StepLR的階梯式突變對(duì)模型收斂的干擾。在果蔬分類這類中小數(shù)據(jù)集上余弦退火配合稍長(zhǎng)訓(xùn)練輪數(shù)如 40 輪通常能比陡降式衰減多出 2 到 3 個(gè)百分點(diǎn)的提升。第三個(gè)手段是微調(diào)策略的分層訓(xùn)練凍結(jié)model.conv1和layer1這些淺層參數(shù)只更新layer3、layer4和fc。淺層卷積捕捉的顏色紋理特征在 ImageNet 上已經(jīng)學(xué)得很好在這類小數(shù)據(jù)集上沒必要重新調(diào)整少更新參數(shù)可以抑制過擬合。凍結(jié)方式是把對(duì)應(yīng)層的requires_grad設(shè)為False優(yōu)化器只接收需要更新層的參數(shù)。# 凍結(jié)前兩層只更新高層特征和分類頭 for name, param in model.named_parameters(): if name.startswith(conv1) or name.startswith(layer1): param.requires_grad False # 優(yōu)化器只接收 requires_gradTrue 的參數(shù) optimizer optim.AdamW( filter(lambda p: p.requires_grad, model.parameters()), lr0.001, weight_decay0.01 )這三個(gè)手段組合下來在同類數(shù)據(jù)集上通常能穩(wěn)定增效 3 到 6 個(gè)點(diǎn)。后續(xù)真正投入時(shí)還需要把best_model.pth導(dǎo)成 ONNX 格式部署到服務(wù)端或者轉(zhuǎn)成 TorchScript 放到移動(dòng)端到那一步才能真正體會(huì)到完整鏈路跑通對(duì)整個(gè)項(xiàng)目的價(jià)值。我在做類似項(xiàng)目時(shí)習(xí)慣用一個(gè)獨(dú)立的實(shí)驗(yàn)記錄表把每次修改的增強(qiáng)策略、學(xué)習(xí)率、優(yōu)化器、最終準(zhǔn)確率記下來調(diào)參時(shí)對(duì)著表比較而不是隨手試。希望這些步驟和踩坑記錄能幫你在自己的圖像分類任務(wù)上少走彎路。本文還有配套的精品資源點(diǎn)擊獲取