習(xí)實戰(zhàn):CNN草莓腐爛圖像分類模型訓(xùn)練與部署)
簡介這份資源面向具備一定Python基礎(chǔ)、希望入門計算機(jī)視覺與深度學(xué)習(xí)實戰(zhàn)的開發(fā)者與學(xué)習(xí)者圍繞草莓腐爛識別這一具體場景提供從數(shù)據(jù)準(zhǔn)備到模型訓(xùn)練再到可視化交互的完整代碼方案。壓縮包共523個文件以517張jpg圖片構(gòu)成數(shù)據(jù)集主體另含3個py腳本與3個txt說明文件整體約48.45MB體積輕便便于本地運行。代碼基于PyTorch環(huán)境搭建依次運行數(shù)據(jù)集文本生成、模型訓(xùn)練與PyQt界面三個腳本即可完成流程訓(xùn)練前對圖片做了短邊補(bǔ)灰邊轉(zhuǎn)正方形及旋轉(zhuǎn)角度等預(yù)處理用于擴(kuò)增增強(qiáng)數(shù)據(jù)集訓(xùn)練完成后模型會保存至本地。已有99人學(xué)習(xí)關(guān)注適合作為圖像分類入門練手項目幫助讀者理解數(shù)據(jù)增強(qiáng)、標(biāo)簽生成、訓(xùn)練驗證劃分與界面部署的完整鏈路并可直接替換數(shù)據(jù)集遷移到其他二分類識別任務(wù)中。1. 草莓爛沒爛為什么人眼會看走眼而 CNN 能兜住做過草莓分揀的人都知道最難的不是把明顯長毛的果子挑出來而是那種「看著還行、捏著已經(jīng)軟了」的早期腐爛。人工分揀在流水線上每小時要過幾千顆果子眼睛疲勞之后漏檢率會陡增而一顆爛果混進(jìn)包裝盒整盒的貨架期都會被拖垮。這個標(biāo)題要解決的就是這件事用 Python 加深度學(xué)習(xí)訓(xùn)練一個能判斷草莓是否腐爛的圖像分類模型配套一份圖片數(shù)據(jù)集讓整套流程可以在本地跑通。它適合三類人想找一個完整深度學(xué)習(xí)實戰(zhàn)項目練手的 Python 學(xué)習(xí)者、做農(nóng)產(chǎn)品分揀設(shè)備或質(zhì)檢系統(tǒng)的工程師、以及手里已經(jīng)攢了一批草莓照片但不知道怎么用起來的從業(yè)者。核心鏈路并不復(fù)雜——數(shù)據(jù)集整理、CNN 模型搭建、訓(xùn)練調(diào)參、推理驗證四步走完就能得到一個可用的二分類器。真正決定成敗的是數(shù)據(jù)質(zhì)量和你對過擬合的警惕程度而不是模型有多深。2. 數(shù)據(jù)集先過一遍手草莓腐爛圖片的清洗、劃分與增強(qiáng)2.1 先搞清楚你手里的是什么樣的圖片數(shù)據(jù)集拿到一個草莓腐爛圖片數(shù)據(jù)集第一件事不是急著寫模型而是把目錄結(jié)構(gòu)和樣本分布摸清楚。常見的組織方式是按類別分文件夾比如fresh/和rotten/兩個目錄每個目錄下是若干張 jpg 或 png。你需要確認(rèn)三件事類別是否平衡、圖片尺寸是否統(tǒng)一、有沒有混入明顯不屬于草莓的圖。類別不平衡是這類數(shù)據(jù)集最常見的問題——新鮮草莓的照片往往比腐爛的多得多因為拍新鮮果子容易拍腐爛果子需要等它真的壞掉。如果新鮮和腐爛的比例超過 3:1訓(xùn)練出來的模型會傾向于把一切都判成新鮮準(zhǔn)確率看著高實際召回率慘不忍睹。我一般會先跑一段統(tǒng)計腳本把每個類別的圖片數(shù)量、尺寸分布、文件格式都打印出來。這一步花不了五分鐘但能幫你避開后面幾個小時的無效訓(xùn)練。import os from PIL import Image from collections import Counter data_dir strawberry_dataset for split in [train, val, test]: split_path os.path.join(data_dir, split) if not os.path.exists(split_path): continue for cls in os.listdir(split_path): cls_path os.path.join(split_path, cls) if not os.path.isdir(cls_path): continue sizes [] formats Counter() for fname in os.listdir(cls_path): fpath os.path.join(cls_path, fname) try: with Image.open(fpath) as img: sizes.append(img.size) formats[img.format] 1 except Exception as e: print(f壞圖: {fpath}, 原因: {e}) size_counter Counter(sizes) print(f{split}/{cls}: 共{len(sizes)}張, f尺寸分布{size_counter.most_common(3)}, 格式{formats})這段腳本做了三件事遍歷每個 split 下的每個類別目錄、用 PIL 打開每張圖讀取尺寸和格式、統(tǒng)計尺寸分布和格式分布。如果發(fā)現(xiàn)某個類別里有大量非 jpg 格式或者尺寸差異極大比如混了 200x200 和 2000x2000 的圖說明數(shù)據(jù)集需要先做統(tǒng)一預(yù)處理。Image.open放在 try 里是為了捕獲損壞文件——數(shù)據(jù)集里偶爾會有下載不完整或者傳輸損壞的圖不處理的話訓(xùn)練時會在某個 batch 突然報錯排查起來很煩。2.2 劃分訓(xùn)練集、驗證集、測試集的比例怎么定如果數(shù)據(jù)集已經(jīng)幫你分好了 train/val/test那直接用就行。如果沒有分你需要自己切。草莓腐爛識別這種二分類任務(wù)樣本量通常在幾百到幾千張之間我一般按 7:1.5:1.5 來切也就是訓(xùn)練集 70%、驗證集 15%、測試集 15%。驗證集用來在訓(xùn)練過程中監(jiān)控過擬合測試集只在最后評估時用一次絕對不能拿測試集來調(diào)參否則你得到的準(zhǔn)確率是虛高的。切分的時候要注意一個坑如果數(shù)據(jù)集里的圖片是從視頻里抽幀來的相鄰幀之間幾乎一模一樣隨機(jī)切分會導(dǎo)致訓(xùn)練集和驗證集里出現(xiàn)近乎重復(fù)的圖驗證集準(zhǔn)確率會虛高。判斷方法很簡單——如果驗證集準(zhǔn)確率在第一輪就沖到 95% 以上大概率是數(shù)據(jù)泄漏了。遇到這種情況要按視頻來源分組切分同一段視頻的幀只能進(jìn)同一個 split。import splitfolders # 假設(shè)原始數(shù)據(jù)在 raw_data/ 下按類別分文件夾 splitfolders.ratio( raw_data, outputstrawberry_dataset, seed42, ratio(0.7, 0.15, 0.15), group_prefixNone )splitfolders.ratio的seed參數(shù)保證每次切分結(jié)果一致方便復(fù)現(xiàn)。ratio的順序?qū)?yīng) train/val/test。如果你的數(shù)據(jù)有分組屬性比如每張圖屬于某個采摘批次可以用group_prefix按前綴分組避免同組數(shù)據(jù)跨 split。2.3 數(shù)據(jù)增強(qiáng)讓幾百張圖發(fā)揮幾千張的效果草莓腐爛數(shù)據(jù)集通常不會太大幾百到一兩千張是常態(tài)。這種量級直接訓(xùn) CNN 很容易過擬合數(shù)據(jù)增強(qiáng)是必須的。對于草莓這種目標(biāo)水平翻轉(zhuǎn)、小角度旋轉(zhuǎn)、亮度微調(diào)是安全且有效的垂直翻轉(zhuǎn)要慎用因為草莓正常擺放都是果蒂朝上垂直翻轉(zhuǎn)會產(chǎn)生現(xiàn)實中不太可能出現(xiàn)的姿態(tài)反而引入噪聲。我一般用 torchvision 的 transforms 做在線增強(qiáng)訓(xùn)練時實時變換驗證和測試時只做 resize 和歸一化。歸一化的均值和標(biāo)準(zhǔn)差用 ImageNet 的[0.485, 0.456, 0.406]和[0.229, 0.224, 0.225]就行即使你的數(shù)據(jù)不是 ImageNet 分布用這套參數(shù)也不會出大問題而且方便直接加載預(yù)訓(xùn)練權(quán)重。from torchvision import transforms train_tf transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(p0.5), transforms.RandomRotation(degrees15), 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_tf transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])RandomRotation(15)表示正負(fù) 15 度內(nèi)隨機(jī)旋轉(zhuǎn)再大就可能把草莓轉(zhuǎn)出畫面。ColorJitter的亮度、對比度、飽和度各 0.2 的擾動幅度模擬不同光照條件下的拍攝差異。注意增強(qiáng)只在訓(xùn)練集上做驗證集和測試集必須用確定性的變換否則每次評估結(jié)果都在變沒法比較。提示如果你的數(shù)據(jù)集里腐爛樣本明顯偏少除了增強(qiáng)還可以考慮用WeightedRandomSampler給少數(shù)類更高的采樣權(quán)重這比簡單復(fù)制樣本更不容易過擬合。3. 用遷移學(xué)習(xí)搭一個能打的草莓腐爛分類器3.1 為什么選 ResNet18 而不是自己從零搭 CNN草莓腐爛識別的視覺特征其實不算復(fù)雜——顏色從鮮紅變暗、表面出現(xiàn)白色或灰色菌絲、果面凹陷。這些特征在淺層卷積里就能捕捉到不需要 ResNet50 甚至更深的網(wǎng)絡(luò)。ResNet18 參數(shù)量約 1100 萬在幾百到幾千張圖的量級上剛好夠用訓(xùn)練快顯存占用低普通筆記本的 GPU 甚至 CPU 都能跑。自己從零搭一個五六層的 CNN 也能做但收斂慢、對初始化和學(xué)習(xí)率敏感除非你是為了學(xué)習(xí) CNN 結(jié)構(gòu)否則沒必要。遷移學(xué)習(xí)的做法是加載 ImageNet 預(yù)訓(xùn)練權(quán)重把最后的全連接層換成二分類輸出。ImageNet 上訓(xùn)出來的淺層卷積核已經(jīng)能識別邊緣、紋理、顏色塊這些特征對草莓腐爛識別同樣有效。凍結(jié)前面的層只訓(xùn)分類頭是一種做法但更常見的是全部解凍一起微調(diào)只是給預(yù)訓(xùn)練層設(shè)一個更小的學(xué)習(xí)率。import torch import torch.nn as nn from torchvision import models def build_model(num_classes2, freeze_backboneFalse): model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) if freeze_backbone: for param in model.parameters(): param.requires_grad False in_features model.fc.in_features model.fc nn.Sequential( nn.Dropout(0.3), nn.Linear(in_features, num_classes) ) return model model build_model(num_classes2, freeze_backboneFalse) device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device)ResNet18_Weights.IMAGENET1K_V1是 torchvision 提供的預(yù)訓(xùn)練權(quán)重枚舉比舊版的pretrainedTrue更明確。freeze_backboneFalse表示全部參數(shù)都參與訓(xùn)練適合數(shù)據(jù)量在 1000 張以上的情況如果只有兩三百張可以先凍結(jié)主干只訓(xùn)分類頭訓(xùn)幾輪后再解凍微調(diào)。Dropout(0.3)加在全連接層前面是防止過擬合的常規(guī)操作草莓?dāng)?shù)據(jù)集小這個 dropout 很有必要。3.2 訓(xùn)練循環(huán)里必須盯住的三個量訓(xùn)練循環(huán)本身不復(fù)雜但有三個量你必須每輪都看訓(xùn)練損失、驗證損失、驗證準(zhǔn)確率。訓(xùn)練損失持續(xù)下降但驗證損失開始上升就是過擬合的典型信號這時候要么加增強(qiáng)、要么加 dropout、要么早停。驗證準(zhǔn)確率震蕩不升可能是學(xué)習(xí)率太大也可能是數(shù)據(jù)標(biāo)注有問題。我見過最隱蔽的坑是標(biāo)注錯誤——把幾張腐爛圖誤放進(jìn)了新鮮目錄模型怎么訓(xùn)都到不了高準(zhǔn)確率最后靠混淆矩陣才定位到。import torch.optim as optim from torch.utils.data import DataLoader from torchvision.datasets import ImageFolder train_ds ImageFolder(strawberry_dataset/train, transformtrain_tf) val_ds ImageFolder(strawberry_dataset/val, transformval_tf) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers4) val_loader DataLoader(val_ds, batch_size32, shuffleFalse, num_workers4) criterion nn.CrossEntropyLoss() optimizer optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-4) scheduler optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max20) best_acc 0.0 for epoch in range(20): model.train() running_loss 0.0 for imgs, labels in train_loader: imgs, labels imgs.to(device), labels.to(device) optimizer.zero_grad() outputs model(imgs) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() * imgs.size(0) scheduler.step() model.eval() correct, total 0, 0 val_loss 0.0 with torch.no_grad(): for imgs, labels in val_loader: imgs, labels imgs.to(device), labels.to(device) outputs model(imgs) loss criterion(outputs, labels) val_loss loss.item() * imgs.size(0) preds outputs.argmax(dim1) correct (preds labels).sum().item() total labels.size(0) acc correct / total print(fEpoch {epoch1}: train_loss{running_loss/len(train_ds):.4f}, fval_loss{val_loss/total:.4f}, val_acc{acc:.4f}) if acc best_acc: best_acc acc torch.save(model.state_dict(), best_strawberry.pth)AdamW比普通 Adam 多了正確的權(quán)重衰減實現(xiàn)weight_decay1e-4是常用起點。CosineAnnealingLR讓學(xué)習(xí)率按余弦曲線從 1e-4 降到接近 0比固定學(xué)習(xí)率更容易收斂到好的局部最優(yōu)。batch_size32在 224x224 輸入下對顯存要求適中如果顯存不夠就降到 16。保存best_strawberry.pth時只存state_dict加載時需要先實例化模型再load_state_dict這是 PyTorch 的標(biāo)準(zhǔn)做法。3.3 學(xué)習(xí)率和 batch size 的搭配經(jīng)驗學(xué)習(xí)率和 batch size 是聯(lián)動參數(shù)不能單獨調(diào)。經(jīng)驗法則是batch size 翻倍學(xué)習(xí)率也大致翻倍。如果你從 batch 32 換到 batch 64學(xué)習(xí)率可以從 1e-4 提到 2e-4。反過來如果顯存只夠 batch 8學(xué)習(xí)率要降到 2.5e-5 左右否則梯度更新太劇烈損失會震蕩。另一個常見問題是微調(diào)時學(xué)習(xí)率設(shè)太大把預(yù)訓(xùn)練權(quán)重「沖毀」了。預(yù)訓(xùn)練權(quán)重是 ImageNet 上花了大量算力學(xué)到的微調(diào)時應(yīng)該用較小的學(xué)習(xí)率1e-4 到 1e-5讓模型在原有特征基礎(chǔ)上做小幅調(diào)整。如果你發(fā)現(xiàn)訓(xùn)練前幾輪驗證準(zhǔn)確率反而下降大概率就是學(xué)習(xí)率太大把預(yù)訓(xùn)練特征破壞了。注意如果你用了freeze_backboneTrue先凍結(jié)訓(xùn)練解凍后一定要把學(xué)習(xí)率降一個數(shù)量級否則之前訓(xùn)好的分類頭會被大梯度帶偏。4. 推理、評估與踩坑記錄模型上線前必須過的幾道關(guān)4.1 用混淆矩陣和單張推理驗證模型真實水平準(zhǔn)確率這個指標(biāo)在類別不平衡時會騙人。假設(shè)測試集里 90% 是新鮮草莓模型把所有樣本都判成新鮮準(zhǔn)確率也有 90%但腐爛草莓一個都沒檢出來。所以評估時必須看混淆矩陣重點看腐爛類的召回率——也就是真正腐爛的草莓里有多少被模型找出來了。對于分揀場景漏檢一顆爛果的代價遠(yuǎn)大于把一顆好果誤判成爛果所以召回率比精確率更重要。from sklearn.metrics import confusion_matrix, classification_report import numpy as np model.load_state_dict(torch.load(best_strawberry.pth)) model.eval() all_preds, all_labels [], [] with torch.no_grad(): for imgs, labels in val_loader: imgs imgs.to(device) outputs model(imgs) preds outputs.argmax(dim1).cpu().numpy() all_preds.extend(preds) all_labels.extend(labels.numpy()) print(confusion_matrix(all_labels, all_preds)) print(classification_report(all_labels, all_preds, target_names[fresh, rotten]))confusion_matrix輸出的 2x2 矩陣對角線是正確分類右上角是「實際腐爛但判成新鮮」的漏檢數(shù)這個數(shù)要盡量小。classification_report會給出每個類別的精確率、召回率、F1重點看 rotten 那一行的 recall。如果 recall 低于 0.85說明模型對腐爛特征學(xué)得不夠需要檢查數(shù)據(jù)增強(qiáng)是否過度、或者腐爛樣本是否太少。單張推理的代碼也要寫一個方便實際使用時快速驗證from PIL import Image def predict(image_path, model, transform, device): img Image.open(image_path).convert(RGB) tensor transform(img).unsqueeze(0).to(device) model.eval() with torch.no_grad(): output model(tensor) prob torch.softmax(output, dim1) pred prob.argmax(dim1).item() label rotten if pred 1 else fresh confidence prob[0][pred].item() return label, confidence label, conf predict(test_strawberry.jpg, model, val_tf, device) print(f判定: {label}, 置信度: {conf:.4f})unsqueeze(0)是給單張圖加一個 batch 維度因為模型期望輸入是[N, C, H, W]。torch.softmax把 logits 轉(zhuǎn)成概率方便看置信度。實際部署時如果置信度低于某個閾值比如 0.7可以標(biāo)記為「不確定」交給人工復(fù)核而不是硬判。4.2 草莓腐爛識別最常見的五個坑坑一驗證集準(zhǔn)確率虛高測試集一塌糊涂?,F(xiàn)象是訓(xùn)練時驗證準(zhǔn)確率 98%換一批新圖測試只有 70%。原因是數(shù)據(jù)泄漏——訓(xùn)練集和驗證集里有同一顆草莓不同角度的照片或者同一段視頻抽的幀。解決辦法是按來源分組切分確保同一顆草莓、同一段視頻的圖只出現(xiàn)在一個 split 里??佣P桶驯尘爱?dāng)特征。現(xiàn)象是模型在訓(xùn)練集上表現(xiàn)很好但換一個背景拍攝的草莓圖就失效。原因是數(shù)據(jù)集里所有腐爛草莓都放在白色盤子上拍所有新鮮草莓都放在木桌上拍模型學(xué)的是盤子 vs 桌子不是草莓本身。解決辦法是統(tǒng)一背景或者用隨機(jī)裁剪讓模型更關(guān)注草莓區(qū)域??尤^擬合到訓(xùn)練集的噪聲?,F(xiàn)象是訓(xùn)練損失降到接近 0驗證損失卻持續(xù)上升。原因是模型太小、數(shù)據(jù)太少、或者訓(xùn)練輪數(shù)太多。解決辦法是加數(shù)據(jù)增強(qiáng)、加 dropout、用早停驗證損失連續(xù)幾輪不降就停??铀念悇e標(biāo)簽搞反?,F(xiàn)象是模型準(zhǔn)確率始終在 50% 附近徘徊怎么調(diào)都上不去。原因是ImageFolder按文件夾名排序fresh排在rotten前面索引 0 是 fresh、1 是 rotten但你在推理時把 0 當(dāng)成了 rotten。解決辦法是打印train_ds.class_to_idx確認(rèn)映射關(guān)系。坑五推理時忘了切 eval 模式?,F(xiàn)象是同一張圖推理兩次結(jié)果不一樣。原因是模型還在 train 模式dropout 和 batch norm 還在隨機(jī)行為。解決辦法是推理前必須調(diào)model.eval()并用torch.no_grad()包住推理過程。提示這五個坑里數(shù)據(jù)泄漏和標(biāo)簽搞反是最難排查的因為它們不會報錯只會讓指標(biāo)難看。養(yǎng)成先打印class_to_idx和檢查 split 來源的習(xí)慣能省下大量調(diào)試時間。5. 把模型推到能用的程度閾值調(diào)優(yōu)與輕量化部署模型訓(xùn)完不是終點能實際用起來才算。這里講兩個進(jìn)階技巧分類閾值調(diào)優(yōu)和模型輕量化。閾值調(diào)優(yōu)解決的是漏檢和誤檢的平衡問題。默認(rèn)情況下softmax 概率大于 0.5 就判為 rotten但你可以把這個閾值降到 0.3讓更多「疑似腐爛」的樣本被攔下來。代價是誤檢增加——一些新鮮草莓會被判成腐爛。在分揀場景里這個代價是值得的因為漏檢一顆爛果的損失遠(yuǎn)大于多扔幾顆好果。調(diào)閾值的方法是在驗證集上畫 ROC 曲線找到召回率滿足要求比如 0.95時對應(yīng)的閾值。from sklearn.metrics import roc_curve # 收集所有驗證集樣本的 rotten 概率 probs [] with torch.no_grad(): for imgs, labels in val_loader: imgs imgs.to(device) outputs model(imgs) p torch.softmax(outputs, dim1)[:, 1].cpu().numpy() probs.extend(p) fpr, tpr, thresholds roc_curve(all_labels, probs) # 找到 tpr 0.95 時對應(yīng)的最小閾值 idx np.where(tpr 0.95)[0][0] best_threshold thresholds[idx] print(f推薦閾值: {best_threshold:.4f}, 此時召回率: {tpr[idx]:.4f}, f誤檢率: {fpr[idx]:.4f})roc_curve返回不同閾值下的假正率和真正率tpr 0.95表示我們要求腐爛草莓的召回率至少 95%在這個前提下選最小的閾值即誤檢最少的那個。這個閾值可以直接寫進(jìn)推理代碼替換掉默認(rèn)的 0.5。模型輕量化解決的是部署到邊緣設(shè)備的問題。ResNet18 的權(quán)重文件大約 45MB在服務(wù)器上跑沒問題但如果要部署到分揀線旁邊的 Jetson 或樹莓派上可能需要更小的模型。兩條路一是換 MobileNetV3 或 EfficientNet-B0參數(shù)量只有 ResNet18 的三分之一到一半精度損失通常在 1-2 個百分點二是對訓(xùn)好的 ResNet18 做動態(tài)量化把權(quán)重從 float32 轉(zhuǎn)成 int8模型體積縮小到約 11MB推理速度提升 2-3 倍精度損失一般在 1% 以內(nèi)。# 動態(tài)量化訓(xùn)練后量化不需要重新訓(xùn)練 quantized_model torch.quantization.quantize_dynamic( model, {nn.Linear}, dtypetorch.qint8 ) torch.save(quantized_model.state_dict(), strawberry_quantized.pth)quantize_dynamic只量化nn.Linear層對 CNN 的卷積層不動這是最保守也最安全的做法。量化后的模型在 CPU 上推理速度提升明顯適合沒有 GPU 的邊緣設(shè)備。注意量化后的模型不能再在 GPU 上跑只能 CPU 推理。最后說一個我自己的習(xí)慣每次訓(xùn)完模型我都會拿十幾張「邊界樣本」單獨測一遍——半爛不爛的、光照很暗的、草莓只占畫面一小角的。這些樣本才是真正考驗?zāi)P偷牡胤綔y試集上的準(zhǔn)確率再高邊界樣本翻車了上線就得挨罵。希望幫到你。本文還有配套的精品資源點擊獲取