據(jù)集實(shí)戰(zhàn):從數(shù)據(jù)加載到CNN訓(xùn)練與混淆矩陣分析)
簡(jiǎn)介本資源為常規(guī)茶葉葉片病害圖像分類數(shù)據(jù)集面向從事農(nóng)業(yè)圖像識(shí)別、深度學(xué)習(xí)分類任務(wù)的學(xué)生與算法工程師可用于訓(xùn)練和評(píng)估CNN分類模型。數(shù)據(jù)集已標(biāo)注共劃分5個(gè)類別包括褐枯病、灰枯萎病、紅點(diǎn)病等具體類別信息可查看包內(nèi)json文件。資源已按訓(xùn)練集、驗(yàn)證集、測(cè)試集劃分各類別圖片分別存放便于直接加載訓(xùn)練同時(shí)提供show腳本可快速可視化樣本分布與圖像內(nèi)容。壓縮包共2000個(gè)文件以1998張jpg圖像為主另含1個(gè)py腳本和1個(gè)json標(biāo)注文件整體約21.68MB結(jié)構(gòu)清晰、開箱即用。目前已有104人學(xué)習(xí)下載適合作為茶葉病害分類項(xiàng)目的基準(zhǔn)數(shù)據(jù)也可結(jié)合CNN網(wǎng)絡(luò)改進(jìn)思路進(jìn)行模型對(duì)比實(shí)驗(yàn)幫助讀者快速完成數(shù)據(jù)加載、類別核對(duì)與訓(xùn)練驗(yàn)證流程。1. 茶葉葉片病害分類數(shù)據(jù)集4000 張已標(biāo)注圖像能直接跑通什么拿到一個(gè)圖像分類數(shù)據(jù)集第一反應(yīng)不該是「有多少?gòu)垺苟恰笜?biāo)注結(jié)構(gòu)長(zhǎng)什么樣、能不能直接喂給訓(xùn)練腳本」。這份常規(guī)茶葉葉片病害圖像分類數(shù)據(jù)集約 4000 張已標(biāo)注圖像分 5 個(gè)類別——褐枯病、灰枯萎病、紅點(diǎn)病等具體類別名以隨包的 json 文件為準(zhǔn)。它已經(jīng)把訓(xùn)練集、驗(yàn)證集、測(cè)試集按同一類別分目錄存放還帶了一個(gè)可視化腳本省掉了自己寫劃分邏輯的功夫。適合誰用一是想快速驗(yàn)證 CNN 分類網(wǎng)絡(luò)改進(jìn)效果的人二是做農(nóng)業(yè)圖像識(shí)別、需要一份干凈多分類數(shù)據(jù)做 baseline 的從業(yè)者三是教學(xué)場(chǎng)景里要演示「數(shù)據(jù)加載到訓(xùn)練」完整鏈路的。不適合誰想直接拿來做目標(biāo)檢測(cè)的這份是分類標(biāo)注不是邊界框標(biāo)注別硬套 YOLO 那套流程。下面從目錄結(jié)構(gòu)、加載、訓(xùn)練、避坑到進(jìn)階一層層拆開講。2. 數(shù)據(jù)集目錄結(jié)構(gòu)與標(biāo)注格式先看清 json 再動(dòng)手2.1 目錄組織與類別映射這類分類數(shù)據(jù)集常見的組織方式是train/val/test三個(gè)根目錄每個(gè)根目錄下再按類別名建子文件夾圖片直接放在對(duì)應(yīng)類別文件夾里。這種結(jié)構(gòu)的好處是torchvision.datasets.ImageFolder能直接讀不需要額外寫索引文件。但這份資源額外帶了 json 文件說明類別名和編號(hào)的映射關(guān)系可能存在 json 里而不是單純靠文件夾名。先做一件事把 json 讀出來確認(rèn)類別數(shù)和類別名。很多翻車現(xiàn)場(chǎng)就是「以為有 5 類結(jié)果 json 里寫了 6 類最后一類只有 3 張圖」訓(xùn)練時(shí) loss 直接 NaN。import json import os # 假設(shè) json 文件在數(shù)據(jù)集根目錄名為 classes.json json_path ./tea_dataset/classes.json with open(json_path, r, encodingutf-8) as f: class_info json.load(f) print(類別總數(shù):, len(class_info)) for idx, name in class_info.items(): print(f 編號(hào) {idx} - 類別 {name})邏輯說明這段代碼只做一件事——把類別映射打印出來。參數(shù)上encodingutf-8必須加中文類別名在 Windows 下默認(rèn)編碼容易亂碼。如果 json 結(jié)構(gòu)不是{編號(hào): 名稱}而是{classes: [...]}把取值那行改成class_info[classes]即可。跑完這一步你心里就有底了到底幾類、每類叫什么。2.2 統(tǒng)計(jì)每類樣本數(shù)識(shí)別長(zhǎng)尾類別不均衡是農(nóng)業(yè)病害數(shù)據(jù)集的常態(tài)。褐枯病可能拍了 1500 張紅點(diǎn)病只有 300 張。不先統(tǒng)計(jì)就開訓(xùn)模型會(huì)偏向多數(shù)類驗(yàn)證集準(zhǔn)確率看著高實(shí)際對(duì)少數(shù)類幾乎沒識(shí)別能力。import os from collections import Counter root ./tea_dataset/train counter Counter() for cls_name in os.listdir(root): cls_dir os.path.join(root, cls_name) if os.path.isdir(cls_dir): imgs [f for f in os.listdir(cls_dir) if f.lower().endswith((.jpg, .png, .jpeg))] counter[cls_name] len(imgs) for cls_name, num in counter.most_common(): print(f{cls_name}: {num} 張) total sum(counter.values()) print(訓(xùn)練集合計(jì):, total)邏輯說明os.listdir遍歷類別文件夾endswith過濾出圖片文件避免把.DS_Store或縮略圖緩存算進(jìn)去。most_common()按數(shù)量降序排列一眼就能看出哪個(gè)類是長(zhǎng)尾。如果最多類和最少類差距超過 5 倍訓(xùn)練時(shí)就得考慮加權(quán)采樣或數(shù)據(jù)增強(qiáng)補(bǔ)償這個(gè)后面第 4 章會(huì)展開。2.3 用自帶 show 腳本做可視化抽檢資源里帶了 show 腳本直接跑之前先確認(rèn)它依賴什么。常見做法是 matplotlib 讀幾張圖拼成網(wǎng)格。如果腳本報(bào)錯(cuò)大概率是路徑寫死或者缺少Pillow。我一般會(huì)先手動(dòng)抽檢 20 張確認(rèn)圖片沒有損壞、沒有標(biāo)注錯(cuò)位。# 先看 show 腳本依賴 head -30 show.py # 常見依賴缺失時(shí)補(bǔ)裝 pip install matplotlib pillow numpy # 運(yùn)行可視化 python show.py邏輯說明head -30先看腳本頭部導(dǎo)入和路徑配置避免直接跑報(bào)一堆錯(cuò)不知道從哪查。如果腳本里路徑是絕對(duì)路徑改成相對(duì)路徑再跑??梢暬皇菫榱撕每词菫榱舜_認(rèn)「褐枯病的圖確實(shí)是褐枯病」標(biāo)注質(zhì)量決定后面所有訓(xùn)練的上限。3. 從 ImageFolder 到 DataLoader把 4000 張圖喂進(jìn) CNN3.1 構(gòu)建 Dataset 與劃分校驗(yàn)雖然數(shù)據(jù)集已經(jīng)分好 train/val/test但還是要校驗(yàn)一遍三個(gè)集合的類別是否一致。常見坑是驗(yàn)證集里少了一個(gè)類別訓(xùn)練時(shí)模型學(xué)了 5 類驗(yàn)證時(shí)只有 4 類算準(zhǔn)確率直接報(bào)錯(cuò)。from torchvision import datasets, transforms data_dir ./tea_dataset train_dir os.path.join(data_dir, train) val_dir os.path.join(data_dir, val) test_dir os.path.join(data_dir, test) # 基礎(chǔ)變換統(tǒng)一尺寸 轉(zhuǎn)張量 base_tf transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), ]) train_ds datasets.ImageFolder(train_dir, transformbase_tf) val_ds datasets.ImageFolder(val_dir, transformbase_tf) test_ds datasets.ImageFolder(test_dir, transformbase_tf) print(訓(xùn)練集類別:, train_ds.classes) print(驗(yàn)證集類別:, val_ds.classes) print(測(cè)試集類別:, test_ds.classes) print(類別是否一致:, train_ds.classes val_ds.classes test_ds.classes)邏輯說明ImageFolder會(huì)自動(dòng)按子文件夾名排序生成classes列表三個(gè)集合的classes必須完全相同。Resize((224, 224))是 ImageNet 預(yù)訓(xùn)練模型的標(biāo)配輸入尺寸如果你用別的骨干網(wǎng)絡(luò)按它的要求改。ToTensor()把像素從 0-255 歸一化到 0-1這是后續(xù) Normalize 的前提。3.2 數(shù)據(jù)增強(qiáng)與歸一化參數(shù)茶葉葉片圖像的光照、角度、背景差異大增強(qiáng)是必須的。但增強(qiáng)不能亂加比如隨機(jī)裁剪可能把病斑裁掉顏色抖動(dòng)過度會(huì)讓褐枯病和紅點(diǎn)病顏色特征混淆。train_tf transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomResizedCrop(224, scale(0.8, 1.0)), transforms.RandomHorizontalFlip(p0.5), transforms.RandomRotation(15), transforms.ColorJitter(brightness0.2, contrast0.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]), ])邏輯說明RandomResizedCrop的scale(0.8, 1.0)控制裁剪面積比例下限別低于 0.7否則病斑容易被裁沒。RandomRotation(15)限制在 15 度內(nèi)茶葉葉片方向性沒那么強(qiáng)但旋轉(zhuǎn)太大可能引入不真實(shí)樣本。Normalize的均值和標(biāo)準(zhǔn)差是 ImageNet 統(tǒng)計(jì)值用預(yù)訓(xùn)練權(quán)重時(shí)必須對(duì)齊否則特征分布偏移收斂變慢。驗(yàn)證集只用 Resize Normalize不做隨機(jī)增強(qiáng)保證評(píng)估可復(fù)現(xiàn)。3.3 DataLoader 批大小與線程設(shè)置4000 張圖不算大但批大小和線程數(shù)設(shè)錯(cuò)訓(xùn)練速度差一倍。from torch.utils.data import DataLoader train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_ds, batch_size32, shuffleFalse, num_workers4, pin_memoryTrue) test_loader DataLoader(test_ds, batch_size32, shuffleFalse, num_workers4, pin_memoryTrue)邏輯說明batch_size32是 224 尺寸下的穩(wěn)妥值顯存 8G 以上可以試 64。shuffleTrue只在訓(xùn)練集開驗(yàn)證和測(cè)試必須關(guān)否則評(píng)估結(jié)果每次不一樣。num_workers4在 Linux 下通常夠用Windows 下如果報(bào)錯(cuò)就改成 0 先跑通。pin_memoryTrue在 GPU 訓(xùn)練時(shí)能加速數(shù)據(jù)搬運(yùn)CPU 訓(xùn)練可以關(guān)掉。4. 訓(xùn)練配置與常見問題排查血淚經(jīng)驗(yàn)都在這里4.1 學(xué)習(xí)率與優(yōu)化器選擇分類任務(wù)用預(yù)訓(xùn)練骨干時(shí)學(xué)習(xí)率別設(shè)大。常見做法是主干網(wǎng)絡(luò)用 1e-4分類頭用 1e-3分參數(shù)組設(shè)置。import torch import torch.nn as nn from torchvision import models model models.resnet50(pretrainedTrue) num_classes len(train_ds.classes) model.fc nn.Linear(model.fc.in_features, num_classes) # 分參數(shù)組主干小學(xué)習(xí)率分類頭大學(xué)習(xí)率 backbone_params [p for n, p in model.named_parameters() if fc not in n] head_params [p for n, p in model.named_parameters() if fc in n] optimizer torch.optim.AdamW([ {params: backbone_params, lr: 1e-4}, {params: head_params, lr: 1e-3}, ], weight_decay1e-4) criterion nn.CrossEntropyLoss()邏輯說明pretrainedTrue加載 ImageNet 權(quán)重小數(shù)據(jù)集上這是提點(diǎn)最快的手段。AdamW比Adam多了正確的權(quán)重衰減實(shí)現(xiàn)分類任務(wù)上更穩(wěn)。weight_decay1e-4抑制過擬合數(shù)據(jù)量 4000 張不算大正則化不能省。CrossEntropyLoss是多分類標(biāo)配如果類別不均衡嚴(yán)重加weight參數(shù)傳類別權(quán)重。4.2 訓(xùn)練循環(huán)與驗(yàn)證指標(biāo)訓(xùn)練循環(huán)要記錄訓(xùn)練 loss 和驗(yàn)證準(zhǔn)確率別只看 loss 下降就以為沒問題。def train_one_epoch(model, loader, optimizer, criterion, device): model.train() total_loss, correct, total 0, 0, 0 for imgs, labels in loader: imgs, labels imgs.to(device), labels.to(device) optimizer.zero_grad() outputs model(imgs) loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() * imgs.size(0) correct (outputs.argmax(1) labels).sum().item() total imgs.size(0) return total_loss / total, correct / total torch.no_grad() def evaluate(model, loader, device): model.eval() correct, total 0, 0 for imgs, labels in loader: imgs, labels imgs.to(device), labels.to(device) outputs model(imgs) correct (outputs.argmax(1) labels).sum().item() total imgs.size(0) return correct / total邏輯說明model.train()和model.eval()必須成對(duì)出現(xiàn)影響 BatchNorm 和 Dropout 行為。torch.no_grad()在驗(yàn)證時(shí)關(guān)閉梯度省顯存也提速。loss.item() * imgs.size(0)是按樣本數(shù)加權(quán)平均避免最后一個(gè) batch 不滿時(shí) loss 計(jì)算偏差。驗(yàn)證準(zhǔn)確率才是選模型的依據(jù)訓(xùn)練 loss 低不代表泛化好。4.3 避坑與常見問題排查現(xiàn)象一訓(xùn)練 loss 一直不降準(zhǔn)確率卡在 20% 左右。原因類別數(shù)對(duì)不上或者標(biāo)簽編碼錯(cuò)位。json 里寫 5 類但文件夾有 6 個(gè)ImageFolder按文件夾名排序生成標(biāo)簽和 json 映射不一致。 解決打印train_ds.classes和 json 的類別列表逐項(xiàng)比對(duì)確保順序和數(shù)量完全一致。不一致就重命名文件夾或重寫 json?,F(xiàn)象二驗(yàn)證準(zhǔn)確率比訓(xùn)練準(zhǔn)確率高很多。原因驗(yàn)證集太小或者驗(yàn)證集和訓(xùn)練集有重疊圖片。4000 張按 7:2:1 劃分驗(yàn)證集只有 800 張波動(dòng)大正常但高太多就是數(shù)據(jù)泄漏。 解決用圖片 MD5 去重檢查 train 和 val 是否有相同文件。常見做法是劃分前先對(duì)所有圖片做哈希按哈希劃分而不是按文件名?,F(xiàn)象三訓(xùn)練幾個(gè) epoch 后 loss 突然變 NaN。原因?qū)W習(xí)率太大或者某張圖片損壞導(dǎo)致梯度爆炸。 解決先把學(xué)習(xí)率降 10 倍試。如果還 NaN在 Dataset 的__getitem__里加 try-except把讀不出來的圖片路徑打出來直接刪掉或修復(fù)。現(xiàn)象四GPU 顯存夠但利用率很低訓(xùn)練慢。原因num_workers設(shè)太小或者數(shù)據(jù)增強(qiáng)在 CPU 上成了瓶頸。 解決把num_workers加到 8 試同時(shí)用pin_memoryTrue。如果還慢檢查是不是每張圖都做了耗時(shí)的顏色變換適當(dāng)簡(jiǎn)化增強(qiáng)。現(xiàn)象五測(cè)試集準(zhǔn)確率遠(yuǎn)低于驗(yàn)證集。原因測(cè)試集分布和訓(xùn)練驗(yàn)證不一致比如測(cè)試集圖片來自不同拍攝設(shè)備或不同光照條件。 解決這不是代碼問題是數(shù)據(jù)問題。要么補(bǔ)充測(cè)試集同分布的樣本要么在訓(xùn)練時(shí)加入更強(qiáng)的光照和顏色增強(qiáng)提升模型魯棒性。5. 進(jìn)階技巧用混淆矩陣和 t-SNE 驗(yàn)證模型到底學(xué)到了什么準(zhǔn)確率只是一個(gè)數(shù)字5 分類任務(wù)里模型可能把紅點(diǎn)病全預(yù)測(cè)成褐枯病但另外三類全對(duì)準(zhǔn)確率照樣 80%。要真正判斷模型可用性得看混淆矩陣和特征分布。import numpy as np from sklearn.metrics import confusion_matrix, classification_report import matplotlib.pyplot as plt import seaborn as sns torch.no_grad() def get_all_preds(model, loader, device): model.eval() all_preds, all_labels [], [] for imgs, labels in loader: imgs imgs.to(device) outputs model(imgs) preds outputs.argmax(1).cpu().numpy() all_preds.extend(preds) all_labels.extend(labels.numpy()) return np.array(all_preds), np.array(all_labels) preds, labels get_all_preds(model, test_loader, device) cm confusion_matrix(labels, preds) plt.figure(figsize(8, 6)) sns.heatmap(cm, annotTrue, fmtd, xticklabelstest_ds.classes, yticklabelstest_ds.classes) plt.xlabel(預(yù)測(cè)類別) plt.ylabel(真實(shí)類別) plt.title(測(cè)試集混淆矩陣) plt.tight_layout() plt.savefig(confusion_matrix.png, dpi150) print(classification_report(labels, preds, target_namestest_ds.classes))邏輯說明confusion_matrix的行是真實(shí)標(biāo)簽列是預(yù)測(cè)標(biāo)簽對(duì)角線是正確分類。如果某一列特別亮說明模型偏向預(yù)測(cè)那個(gè)類。classification_report給出每個(gè)類的精確率、召回率和 F1比整體準(zhǔn)確率更有參考價(jià)值。sns.heatmap的annotTrue把數(shù)字標(biāo)在格子里fmtd保證顯示整數(shù)?;煜仃嚹芨嬖V你「哪個(gè)類被混淆了」但不知道「為什么混淆」。這時(shí)候用 t-SNE 把倒數(shù)第二層的特征降到二維看分布。from sklearn.manifold import TSNE torch.no_grad() def extract_features(model, loader, device): model.eval() # 去掉最后的全連接層取池化后的特征 feature_extractor torch.nn.Sequential( *list(model.children())[:-1] ).to(device) feats, labs [], [] for imgs, labels in loader: imgs imgs.to(device) out feature_extractor(imgs).squeeze(-1).squeeze(-1) feats.append(out.cpu().numpy()) labs.extend(labels.numpy()) return np.concatenate(feats), np.array(labs) feats, labs extract_features(model, test_loader, device) tsne TSNE(n_components2, perplexity30, random_state42) feats_2d tsne.fit_transform(feats) plt.figure(figsize(8, 6)) for i, cls_name in enumerate(test_ds.classes): mask labs i plt.scatter(feats_2d[mask, 0], feats_2d[mask, 1], labelcls_name, alpha0.6, s10) plt.legend() plt.title(測(cè)試集特征 t-SNE 分布) plt.tight_layout() plt.savefig(tsne.png, dpi150)邏輯說明model.children()[:-1]去掉 ResNet 最后的全連接層保留全局池化輸出得到每張圖的特征向量。perplexity30是 t-SNE 的常用值數(shù)據(jù)量 800 左右時(shí) 30 到 50 都合理。random_state42保證每次跑圖一致方便對(duì)比不同模型。如果 t-SNE 圖上某一類散得到處都是說明模型沒學(xué)到該類別的判別特征要么數(shù)據(jù)不夠要么增強(qiáng)過度把特征破壞了。我自己的習(xí)慣是每次訓(xùn)完一個(gè)模型混淆矩陣和 t-SNE 必須各跑一遍不看這兩個(gè)圖不敢說模型能用。有一次準(zhǔn)確率 92% 看著挺好混淆矩陣一出來發(fā)現(xiàn)紅點(diǎn)病召回率只有 0.6全被預(yù)測(cè)成褐枯病后來補(bǔ)了 200 張紅點(diǎn)病樣本才拉回來。從那以后我每次拿到分類數(shù)據(jù)集都強(qiáng)制先跑一遍類別統(tǒng)計(jì)和可視化抽檢再開始訓(xùn)。希望幫到你。本文還有配套的精品資源點(diǎn)擊獲取