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