數(shù)據(jù)集訓(xùn)練:驗(yàn)證集劃分與類(lèi)別不平衡實(shí)戰(zhàn)指南)
簡(jiǎn)介眼睛疾病分類(lèi)數(shù)據(jù)集是一份可直接用于圖像分類(lèi)任務(wù)的中小型醫(yī)學(xué)影像資源包含白內(nèi)障、青光眼、正常、視網(wǎng)膜疾病四個(gè)類(lèi)別適合臨床篩查模型練手、課程實(shí)驗(yàn)或YOLOv5分類(lèi)項(xiàng)目。數(shù)據(jù)按train和test目錄整理訓(xùn)練集481張、測(cè)試集120張均為JPEG格式配合JSON分類(lèi)字典和Python可視化腳本可快速完成數(shù)據(jù)劃分查看與模型迭代。壓縮包共604個(gè)文件除601張圖片外還有1個(gè)字典文件、1個(gè)可視化腳本和1張示例圖整包約61MB輕量易下載。目前已有413人學(xué)習(xí)下載腳本支持隨機(jī)抽取4張圖展示并保存結(jié)果無(wú)需改動(dòng)即可運(yùn)行能幫助使用者快速核對(duì)數(shù)據(jù)質(zhì)量和類(lèi)別分布。1. 拿到眼睛疾病分類(lèi)數(shù)據(jù)集先別急著訓(xùn)訓(xùn)練集和驗(yàn)證集到底在做什么接手一個(gè)醫(yī)學(xué)圖像的眼睛疾病分類(lèi)數(shù)據(jù)集時(shí)真正讓人栽跟頭的往往不是模型而是文件夾里那兩個(gè)split訓(xùn)練集和驗(yàn)證集。很多人直接把train和val合并重訓(xùn)或者反復(fù)拿驗(yàn)證集調(diào)參最后精度漂亮得可疑一上真實(shí)場(chǎng)景就露餡。要解決的問(wèn)題很具體這個(gè)分類(lèi)數(shù)據(jù)集該按什么結(jié)構(gòu)讀取、類(lèi)別分布怎么看、訓(xùn)練流程怎么寫(xiě)、驗(yàn)證集怎么用才不作弊。適合用PyTorch做醫(yī)學(xué)圖像分類(lèi)的算法工程師、做畢設(shè)的學(xué)生以及想從yolo那套自定義數(shù)據(jù)習(xí)慣切到分類(lèi)任務(wù)的人。2. 拆解眼睛疾病分類(lèi)數(shù)據(jù)集目錄結(jié)構(gòu)、標(biāo)簽格式與劃分合理性檢查2.1 先看目錄結(jié)構(gòu)和標(biāo)簽格式再?zèng)Q定用什么姿勢(shì)讀取常見(jiàn)眼睛疾病分類(lèi)數(shù)據(jù)集的組織方式通常是train目錄下按類(lèi)別建子文件夾val目錄同樣按類(lèi)別建子文件夾圖片文件散落在各自的類(lèi)別文件夾里。類(lèi)別名即標(biāo)簽文件夾名就是醫(yī)生給的診斷結(jié)論。公開(kāi)數(shù)據(jù)集里的類(lèi)別體系大體圍繞眼底鏡圖像展開(kāi)正常、糖尿病視網(wǎng)膜病變、青光眼、白內(nèi)障、黃斑變性、高血壓視網(wǎng)膜病變、近視等。這類(lèi)圖像通常由眼底相機(jī)采集也有部分是醫(yī)院病歷系統(tǒng)里導(dǎo)出的彩色照片。拿到數(shù)據(jù)后第一件事不是寫(xiě)訓(xùn)練腳本而是確認(rèn)兩類(lèi)元信息圖片擴(kuò)展名是否統(tǒng)一.jpg和.png混用非常常見(jiàn)類(lèi)別文件夾里有沒(méi)有混入非圖片文件比如隱藏的desktop.ini或macOS的.DS_Store。最有效的檢查方式是直接對(duì)每個(gè)split做一次文件統(tǒng)計(jì)把類(lèi)別和數(shù)量一次性打出來(lái)。import os from collections import Counter data_root eye_disease_dataset for split in [train, val]: split_path os.path.join(data_root, split) if not os.path.isdir(split_path): print(f{split} 目錄不存在先檢查數(shù)據(jù)集路徑) continue classes [d for d in os.listdir(split_path) if os.path.isdir(os.path.join(split_path, d))] per_class {} for cls in sorted(classes): per_class[cls] len(os.listdir(os.path.join(split_path, cls))) total sum(per_class.values()) print(f[{split}] 共 {total} 張{len(classes)} 個(gè)類(lèi)別) for cls, n in sorted(per_class.items(), keylambda x: -x[1]): print(f {cls}: {n} ({n / total * 100:.2f}%))這段腳本輸出每個(gè)類(lèi)別在訓(xùn)練集和驗(yàn)證集的數(shù)量占比。眼睛疾病數(shù)據(jù)集的通病是類(lèi)別不平衡正常眼通常是數(shù)量最多的類(lèi)而糖尿病視網(wǎng)膜病變的早期樣本可能只有正常眼的零頭。如果某個(gè)類(lèi)在驗(yàn)證集里只有個(gè)位數(shù)對(duì)應(yīng)的acc、precision都不可信后面必須換成per-class指標(biāo)。另一個(gè)判斷依據(jù)是目錄層級(jí)torchvision的ImageFolder要求類(lèi)別文件夾直接掛在split下如果數(shù)據(jù)集是train/class/subfolder這種二次封裝結(jié)構(gòu)或者用csv索引標(biāo)簽就要寫(xiě)自定義Dataset不能硬套現(xiàn)成工具。2.2 驗(yàn)證集、測(cè)試集和訓(xùn)練集標(biāo)題只給了兩個(gè)split時(shí)怎么補(bǔ)第三個(gè)很多公開(kāi)醫(yī)學(xué)圖像數(shù)據(jù)集只劃分了train和val沒(méi)有test。原因通常是數(shù)據(jù)量少官方想給使用者預(yù)留調(diào)參空間。但作為落地的人必須自己補(bǔ)出一個(gè)test split否則報(bào)出來(lái)的所有指標(biāo)都可能被驗(yàn)證集“污染”。驗(yàn)證集用來(lái)做模型選擇、超參調(diào)優(yōu)和早停測(cè)試集用來(lái)估計(jì)最終交給業(yè)務(wù)方時(shí)的真實(shí)性能。常見(jiàn)做法是從訓(xùn)練集里再切一小塊出來(lái)當(dāng)測(cè)試集。假如訓(xùn)練集有8000張按分層抽樣切出10%約800張作為test剩下的做train。切分時(shí)用sklearn的train_test_splitstratify按類(lèi)別標(biāo)簽分層保證每個(gè)類(lèi)在test里的比例和train一致。import shutil from pathlib import Path from sklearn.model_selection import train_test_split train_root Path(eye_disease_dataset/train) test_root Path(eye_disease_dataset/test) test_root.mkdir(exist_okTrue) for cls_dir in train_root.iterdir(): if not cls_dir.is_dir(): continue imgs list(cls_dir.glob(*)) _, test_imgs train_test_split(imgs, test_size0.1, random_state42) dest test_root / cls_dir.name dest.mkdir(parentsTrue, exist_okTrue) for img in test_imgs: shutil.copy(str(img), str(dest / img.name))用random_state固定隨機(jī)種子保證切分可復(fù)現(xiàn)用copy而不是move防止改主意后數(shù)據(jù)被搬走。有一個(gè)細(xì)節(jié)容易被忽略如果這張數(shù)據(jù)來(lái)自同一病人的多角度拍攝這個(gè)腳本是不安全的先按病人分組再切具體做法在2.3和避坑章展開(kāi)。測(cè)試集切出來(lái)之后只碰一次不要在它上面反復(fù)調(diào)參否則測(cè)試集就變成了第二個(gè)驗(yàn)證集失去了終極驗(yàn)證的意義。2.3 劃分合理性檢查同一個(gè)人可能出現(xiàn)在兩個(gè)集合里嗎眼睛疾病分類(lèi)數(shù)據(jù)集大多來(lái)自醫(yī)院采集同一個(gè)病人可能有兩眼甚至多張不同時(shí)間拍攝的眼底圖。如果切分按文件隨機(jī)打散同一病人的多張圖很可能同時(shí)出現(xiàn)在訓(xùn)練集和驗(yàn)證集。模型會(huì)把“這個(gè)病人的視盤(pán)形態(tài)”記下來(lái)而不是學(xué)習(xí)“這類(lèi)疾病的通用特征”驗(yàn)證集acc會(huì)虛高。檢查方式先找數(shù)據(jù)集自帶的元數(shù)據(jù)csv或DICOM頭看有沒(méi)有patient_id字段。沒(méi)有元數(shù)據(jù)時(shí)部分?jǐn)?shù)據(jù)集文件名會(huì)帶patient前綴。如果兩者都沒(méi)有可以用感知哈希做近似重復(fù)圖片檢測(cè)from PIL import Image import numpy as np def phash(path, size16): img Image.open(path).convert(L).resize((size, size)) pixels np.array(img, dtypenp.float32) avg pixels.mean() return .join(1 if p avg else 0 for p in pixels.flatten()) def hamming(a, b): return sum(c1 ! c2 for c1, c2 in zip(a, b))phash把圖像縮小成16x16的灰度指紋漢明距離小于等于4的兩張圖基本可以認(rèn)定是重復(fù)或近似重復(fù)。但這個(gè)方法只能找出“拷貝/裁剪”級(jí)別的重復(fù)同一個(gè)病人兩只眼的外觀差異明顯phash查不出來(lái)。最穩(wěn)妥還是靠patient_id切分先按病人分組再在所有病人上做train/val/test的分層切分。切完再回頭看一眼train和val的類(lèi)別分布確認(rèn)兩個(gè)集合里的病人集合沒(méi)有交集。3. 用 PyTorch 跑通眼睛疾病分類(lèi)的最小訓(xùn)練流程ResNet 路線3.1 數(shù)據(jù)讀取ImageFolder 的兩個(gè)注意點(diǎn)torchvision的ImageFolder天然適配第2章的目錄結(jié)構(gòu)不用寫(xiě)任何自定義Dataset。訓(xùn)練集和驗(yàn)證集分別掛不同的transform訓(xùn)練集做隨機(jī)增強(qiáng)驗(yàn)證集只做尺寸統(tǒng)一和標(biāo)準(zhǔn)化import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms, models train_tf transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomResizedCrop(224, scale(0.7, 1.0)), transforms.RandomHorizontalFlip(p0.5), transforms.ColorJitter(brightness0.15, contrast0.15), 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]) ]) train_ds datasets.ImageFolder(eye_disease_dataset/train, transformtrain_tf) val_ds datasets.ImageFolder(eye_disease_dataset/val, transformval_tf) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_ds, batch_size32, shuffleFalse, num_workers4, pin_memoryTrue)驗(yàn)證集不應(yīng)用RandomResizedCrop和Flip驗(yàn)證要的是確定性的結(jié)果訓(xùn)練集用RandomResizedCrop能模擬眼底相機(jī)拍攝角度和視場(chǎng)范圍的差異。Normalize沿用ImageNet的mean/std對(duì)眼底圖這種紅色調(diào)為主的圖像其實(shí)夠用。如果發(fā)現(xiàn)圖像分布差異很大可以從數(shù)據(jù)集中采樣幾千張算出自己的mean/std替換但多數(shù)場(chǎng)景沒(méi)必要。注意Windows上num_workers設(shè)為0最穩(wěn)Linux下再按CPU核數(shù)往上加。先跑通流程再優(yōu)化加載速度。3.2 訓(xùn)練腳本骨架損失函數(shù)、優(yōu)化器與驗(yàn)證時(shí)機(jī)用預(yù)訓(xùn)練ResNet50作為backbone是多數(shù)眼睛疾病分類(lèi)項(xiàng)目入門(mén)標(biāo)配。參數(shù)少、權(quán)重好找、微調(diào)穩(wěn)定。替換最后一層全連接損失函數(shù)先用最樸素的CrossEntropyLoss驗(yàn)證集每個(gè)epoch都算一次acc保存val acc最高的checkpointmodel models.resnet50(weightsmodels.ResNet50_Weights.IMAGENET1K_V1) num_classes len(train_ds.classes) model.fc nn.Linear(model.fc.in_features, num_classes) model model.cuda() criterion nn.CrossEntropyLoss() optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-4) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, modemax, factor0.5, patience3) def validate(model, loader): model.eval() correct 0 total 0 all_preds, all_labels [], [] with torch.no_grad(): for images, labels in loader: images, labels images.cuda(), labels.cuda() outputs model(images) _, predicted torch.max(outputs, 1) correct (predicted labels).sum().item() total labels.size(0) all_preds.extend(predicted.cpu().tolist()) all_labels.extend(labels.cpu().tolist()) return correct / total, all_preds, all_labels best_acc 0 no_improve 0 early_stop_patience 5 for epoch in range(30): model.train() for images, labels in train_loader: images, labels images.cuda(), labels.cuda() optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() acc, preds, labels validate(model, val_loader) scheduler.step(acc) print(fepoch {epoch1}: val_acc{acc:.4f} lr{optimizer.param_groups[0][lr]:.2e}) if acc best_acc: best_acc acc no_improve 0 torch.save(model.state_dict(), best_eye_cls.pth) else: no_improve 1 if no_improve early_stop_patience: print(f{early_stop_patience} 個(gè)epoch無(wú)提升早停) breaklr1e-4是微調(diào)全模型的安全起點(diǎn)如果只解凍最后一層fc訓(xùn)練可以用1e-3但全模型微調(diào)降到1e-4更穩(wěn)。weight_decay1e-4抑制醫(yī)學(xué)圖像上容易出現(xiàn)的過(guò)擬合。batch_size32搭配ResNet50在24G以下顯存基本舒適顯存受限改16時(shí)學(xué)習(xí)率相應(yīng)減半。早停條件用“val acc連續(xù)epoch無(wú)提升”而不是“val loss無(wú)下降”醫(yī)學(xué)圖像噪聲大loss和acc并非總是同步。一個(gè)常見(jiàn)的翻車(chē)點(diǎn)val acc已經(jīng)連續(xù)5個(gè)epoch沒(méi)漲但因沒(méi)保存最優(yōu)模型交付的是最后一個(gè)epoch的權(quán)重性能大幅回退。上面把早停和模型保存寫(xiě)在一起訓(xùn)完直接加載best_eye_cls.pth才是真正能用的模型。3.3 關(guān)鍵參數(shù)設(shè)置圖像尺寸、batch size 與學(xué)習(xí)率的搭配眼睛疾病分類(lèi)數(shù)據(jù)集里的圖像尺寸通常不統(tǒng)一。眼底相機(jī)常見(jiàn)2048x1536、1600x1200也有手機(jī)翻拍的病歷圖。Resize到256再中心裁剪到224是ImageNet時(shí)代的標(biāo)準(zhǔn)做法。想追求速度可以縮到192或160但代價(jià)是視網(wǎng)膜小血管、微動(dòng)脈瘤這類(lèi)細(xì)節(jié)可能被模糊掉建議先用224跑通再壓縮。batch和學(xué)習(xí)率的搭配遵循線性縮放原則batch32配lr1e-4batch16則lr減半到5e-5batch64可以嘗試2e-4。下表只適用于單卡小batch場(chǎng)景多卡時(shí)不這么算。batch size學(xué)習(xí)率起點(diǎn)典型場(chǎng)景165e-5顯存受限的舊卡321e-4最常見(jiàn)配置642e-412G以上顯存DataLoader里還有個(gè)容易被忽略的參數(shù)drop_last。醫(yī)學(xué)圖像數(shù)據(jù)集樣本數(shù)經(jīng)常不是batch_size的整數(shù)倍最后一個(gè)batch可能只有幾張圖BN層的統(tǒng)計(jì)會(huì)不穩(wěn)定。訓(xùn)練時(shí)建議設(shè)置drop_lastTrue驗(yàn)證時(shí)保持drop_lastFalse以便統(tǒng)計(jì)所有樣本。4. 驗(yàn)證集評(píng)估與精度調(diào)優(yōu)眼睛疾病分類(lèi)的 3 個(gè)必調(diào)參數(shù)4.1 用混淆矩陣看模型到底錯(cuò)在哪一類(lèi)總acc對(duì)醫(yī)學(xué)圖像分類(lèi)并不夠。類(lèi)別不平衡嚴(yán)重時(shí)正常眼占大頭acc會(huì)被正常類(lèi)拉高模型把所有病變都判成正常也能到60%以上。要在驗(yàn)證集上算per-class的recall和混淆矩陣from sklearn.metrics import classification_report, confusion_matrix import seaborn as sns import matplotlib.pyplot as plt classes val_ds.classes # 這兩個(gè)列表來(lái)自validate()函數(shù)返回的preds和labels report classification_report(labels, preds, target_namesclasses, digits3) print(report) cm confusion_matrix(labels, preds) cm_norm cm.astype(float) / cm.sum(axis1, keepdimsTrue) plt.figure(figsize(10, 8)) sns.heatmap(cm_norm, annotTrue, fmt.2f, cmapBlues, xticklabelsclasses, yticklabelsclasses) plt.xlabel(Predicted) plt.ylabel(True) plt.tight_layout() plt.savefig(eye_confusion_matrix.png, dpi150)混淆矩陣呈現(xiàn)的是“真實(shí)類(lèi)別vs預(yù)測(cè)類(lèi)別”。在醫(yī)學(xué)圖像場(chǎng)景里關(guān)注重點(diǎn)不是對(duì)角線多高而是哪些非對(duì)角線值得警惕。糖尿病視網(wǎng)膜病變和黃斑變性早期都表現(xiàn)為黃斑區(qū)異常模型容易把兩者搞混如果模型把青光眼判成正常這種錯(cuò)誤在臨床上屬于漏診。看矩陣時(shí)先圈出“正常眼那行”的false negative因?yàn)椴』悸┰\比誤診更危險(xiǎn)?,F(xiàn)實(shí)里還有一個(gè)常見(jiàn)現(xiàn)象模型對(duì)驗(yàn)證集里“背景亮度過(guò)高”的樣本常常成片判錯(cuò)。這類(lèi)樣本往往在混淆矩陣某一列扎堆先別急著加數(shù)據(jù)回去看那一類(lèi)圖像是不是存在設(shè)備差異。4.2 類(lèi)別不平衡損失函數(shù)替換與樣本權(quán)重用CrossEntropyLoss時(shí)class weight是最直接的平衡手段。先統(tǒng)計(jì)訓(xùn)練集的每類(lèi)樣本數(shù)再算權(quán)重注意歸一化import os from collections import Counter import torch split_path eye_disease_dataset/train class_counts Counter() for cls in sorted(os.listdir(split_path)): cls_path os.path.join(split_path, cls) if os.path.isdir(cls_path): class_counts[cls] len(os.listdir(cls_path)) counts_tensor torch.tensor([class_counts[c] for c in sorted(class_counts)]) weights 1.0 / counts_tensor.float() weights weights / weights.mean() # 歸一化讓權(quán)重均值保持在1附近 criterion nn.CrossEntropyLoss(weightweights.cuda())直接用1/count會(huì)出現(xiàn)極端類(lèi)權(quán)重過(guò)大的問(wèn)題比如某類(lèi)只有80張權(quán)重會(huì)變成正常眼的幾十倍訓(xùn)練反而震蕩。除以均值把正常類(lèi)壓回1附近正常眼權(quán)重小于1稀有類(lèi)權(quán)重大于1但不會(huì)離譜。如果加了class weight后召回率還是上不去可以換Focal Loss。它自動(dòng)降低易分樣本的loss貢獻(xiàn)讓模型把注意力放在難分的病變類(lèi)別上class FocalLoss(nn.Module): def __init__(self, gamma2.0, alphaNone): super().__init__() self.gamma gamma self.alpha alpha def forward(self, logits, targets): ce_loss nn.functional.cross_entropy(logits, targets, weightself.alpha, reductionnone) pt torch.exp(-ce_loss) focal_loss (1 - pt) ** self.gamma * ce_loss return focal_loss.mean()gamma2.0是常用起點(diǎn)。對(duì)眼睛疾病分類(lèi)我建議先從class weight入手因?yàn)樗桓囊粋€(gè)參數(shù)跑兩個(gè)epoch就能看出趨勢(shì)focal loss要調(diào)gammagamma太大模型會(huì)過(guò)度聚焦難分樣本出現(xiàn)驗(yàn)證集acc原地抖動(dòng)。見(jiàn)過(guò)有人把a(bǔ)lpha和class weight混著用結(jié)果正常眼權(quán)重被壓到0.1以下模型開(kāi)始大量誤報(bào)沒(méi)必要疊這么多。4.3 學(xué)習(xí)率策略從視頻動(dòng)作分類(lèi)實(shí)戰(zhàn)里常用的余弦退火說(shuō)起很多做視頻動(dòng)作分類(lèi)比如跑UCF101這類(lèi)基準(zhǔn)的團(tuán)隊(duì)長(zhǎng)訓(xùn)練時(shí)幾乎默認(rèn)用余弦退火。這不是眼睛疾病分類(lèi)里的新東西但確實(shí)好用。第3.2節(jié)用的ReduceLROnPlateau是驗(yàn)證集驅(qū)動(dòng)的適合訓(xùn)練中期如果數(shù)據(jù)集不大余弦退火的確定性調(diào)度往往更穩(wěn)scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max30, eta_min1e-6) for epoch in range(30): # 現(xiàn)有訓(xùn)練循環(huán) scheduler.step()T_max設(shè)為總epoch數(shù)eta_min設(shè)為初始學(xué)習(xí)率的百分之一從1e-4退到1e-6足夠。用它替代ReduceLROnPlateau時(shí)要注意余弦退火是“從當(dāng)前值一路往下”沒(méi)有回頭漲的機(jī)會(huì)所以初始lr寧可偏低。有人把初始lr設(shè)成1e-3跑余弦退火前幾個(gè)epoch loss直接炸穿。在眼睛疾病分類(lèi)這類(lèi)中小規(guī)模醫(yī)學(xué)圖像數(shù)據(jù)集上我的使用順序是先用ReduceLROnPlateau跑20個(gè)epoch看baseline如果尾部loss震蕩明顯再換余弦退火重跑一次驗(yàn)證集acc通常能再上1到2個(gè)點(diǎn)。先有baseline再調(diào)調(diào)度器比一上來(lái)就堆各種trick更省時(shí)間。5. 眼睛疾病分類(lèi)數(shù)據(jù)集落地避坑5 條真實(shí)的血淚經(jīng)驗(yàn)下面這幾條都是從實(shí)際訓(xùn)練過(guò)程中踩出來(lái)的按“現(xiàn)象、原因、解決”的順序?qū)懹龅筋?lèi)似問(wèn)題可以直接對(duì)號(hào)入座。5.1 驗(yàn)證集acc漂亮得可疑訓(xùn)練集里混進(jìn)了驗(yàn)證集圖像現(xiàn)象訓(xùn)練到一半驗(yàn)證集acc飆到99%但換一批新采集的圖acc直接掉到60%。原因數(shù)據(jù)清洗不徹底。很多公開(kāi)數(shù)據(jù)集的train和val是從原始資料里分出來(lái)的原始文件里有重復(fù)截圖、圖像拷貝同一個(gè)病例的不同版本文檔被誤放進(jìn)了兩個(gè)集合。解決用2.3節(jié)的phash全庫(kù)跑一遍去重漢明距離小于等于4的圖像對(duì)確認(rèn)后只保留一份。更穩(wěn)的是在訓(xùn)練前用文件名或元數(shù)據(jù)查一下“同一病人文件是否被分到兩個(gè)split”。5.2 灰度圖與RGB通道不一致現(xiàn)象訓(xùn)練到中途dataloader報(bào)錯(cuò)“Expected a 3-channel input”或者loss變成nan。原因眼底相機(jī)輸出一般是彩色JPEG但醫(yī)院導(dǎo)出的歷史數(shù)據(jù)里存在灰度PNG和帶透明通道的圖。ImageFolder遇到灰度圖時(shí)ToTensor會(huì)把它變成單通道和Normalize的三通道統(tǒng)計(jì)不匹配。解決在transform之前統(tǒng)一轉(zhuǎn)RGBdef load_as_rgb(path): img Image.open(path) if img.mode ! RGB: img img.convert(RGB) return img把load_as_rgb放進(jìn)自定義Dataset里?;叶葓D轉(zhuǎn)RGB是復(fù)制通道RGBA圖則丟棄Alpha。不要指望現(xiàn)成的ImageFolder幫你處理這些。5.3 驗(yàn)證集acc虛高的背后沒(méi)有按病人切分現(xiàn)象模型在驗(yàn)證集上對(duì)青光眼類(lèi)acc達(dá)到98%業(yè)務(wù)方拿新數(shù)據(jù)實(shí)測(cè)準(zhǔn)確率遠(yuǎn)低于預(yù)期。原因醫(yī)院采集中同一個(gè)病人可能提供兩只眼的圖片隨機(jī)劃分后同一病人的兩只眼一個(gè)在train一個(gè)在val模型學(xué)到的其實(shí)是病人特征。病人ID可能隱藏在文件名里比如patient_001_left.png和patient_001_right.png直接listdir根本看不出來(lái)。解決先抽出patient_id按病人劃分import pandas as pd from sklearn.model_selection import train_test_split df pd.read_csv(image_patient_labels.csv) patients df[patient_id].unique() train_patients, val_patients train_test_split(patients, test_size0.2, random_state42) train_df df[df[patient_id].isin(train_patients)] val_df df[df[patient_id].isin(val_patients)]注意分層如果病人總數(shù)少還要按主診斷做stratify否則可能出現(xiàn)某類(lèi)病人只進(jìn)val的情況。5.4 從 yolo 自定義數(shù)據(jù)集的習(xí)慣遷移過(guò)來(lái)的誤區(qū)現(xiàn)象做檢測(cè)的人第一次拿到分類(lèi)數(shù)據(jù)集時(shí)習(xí)慣性去找label文件、找標(biāo)注框txt發(fā)現(xiàn)沒(méi)有這些文件不知道該怎么訓(xùn)練。原因分類(lèi)數(shù)據(jù)集和yolov8、yolo26乃至deim這類(lèi)目標(biāo)檢測(cè)框架的數(shù)據(jù)約定不同。檢測(cè)用邊界框加txt標(biāo)簽分類(lèi)數(shù)據(jù)集的標(biāo)簽全部隱含在文件夾名里不需要再生成任何label文件。解決直接用ImageFolder按目錄讀。實(shí)在習(xí)慣用CSV就自己寫(xiě)一個(gè)映射文件import csv from pathlib import Path rows [] for split in [train, val]: root Path(feye_disease_dataset/{split}) for cls_dir in root.iterdir(): if not cls_dir.is_dir(): continue for img in cls_dir.glob(*): rows.append([str(img), cls_dir.name]) with open(eye_cls_labels.csv, w, newline) as f: writer csv.writer(f) writer.writerow([path, label]) writer.writerows(rows)CSV的作用是方便后續(xù)做病人級(jí)切分和臟樣本過(guò)濾而不是替代目錄結(jié)構(gòu)。5.5 早停判斷標(biāo)準(zhǔn)別只盯val loss現(xiàn)象訓(xùn)練時(shí)val loss一直在降但驗(yàn)證集acc紋絲不動(dòng)多跑幾個(gè)epoch后acc突然跳幾個(gè)點(diǎn)另一邊val loss止跌你以為可以停了結(jié)果再跑兩個(gè)epoch又漲一點(diǎn)。原因醫(yī)學(xué)圖像類(lèi)別特征差異大loss和acc并非同步變化。小類(lèi)別在loss里的貢獻(xiàn)占比低loss下降反映的只是大類(lèi)特征收斂小類(lèi)的acc沒(méi)有變化。解決以“驗(yàn)證集acc連續(xù)patience個(gè)epoch無(wú)提升”作為早停條件patience設(shè)5到8。如果用了class weight或focal loss同時(shí)盯per-class recall的調(diào)和平均不要只盯總acc因?yàn)榭俛cc會(huì)被正常眼主導(dǎo)。每個(gè)epoch都備份一次最優(yōu)checkpoint就算誤停也有后悔藥。6. 把驗(yàn)證集利用到極致錯(cuò)誤分析是醫(yī)學(xué)圖像分類(lèi)的最后一公里訓(xùn)練結(jié)束不等于交付。從驗(yàn)證集里篩出預(yù)測(cè)錯(cuò)誤的樣本一張一張看才是真正提升模型價(jià)值的部分。做法是保存驗(yàn)證集的softmax輸出抽最底部的錯(cuò)誤樣本probs torch.softmax(outputs, dim1) max_probs, preds torch.max(probs, dim1) filter_mask (preds ! labels) | (max_probs 0.6)把mask篩出來(lái)的圖像路徑和預(yù)測(cè)結(jié)果寫(xiě)成csv對(duì)照訓(xùn)練集里的人工復(fù)核清單再查一遍。這步常會(huì)發(fā)現(xiàn)“預(yù)測(cè)錯(cuò)誤”其實(shí)是標(biāo)注錯(cuò)誤比如早期白內(nèi)障被標(biāo)成正常眼。這類(lèi)臟樣本如果不清理會(huì)一直污染指標(biāo)。從val里剔除后重新評(píng)測(cè)模型能力才是真實(shí)的。我的習(xí)慣是每個(gè)epoch結(jié)束都保存val acc和混淆矩陣訓(xùn)練完成后用驗(yàn)證集里置信度低于0.6的樣本生成一份人工復(fù)核清單而不是把模型輸出當(dāng)裁決者。眼睛疾病分類(lèi)的落地價(jià)值不在acc多高而在于輔助醫(yī)生把漏診率降下來(lái)所以讓模型學(xué)會(huì)說(shuō)“我不確定”比強(qiáng)行輸出一個(gè)錯(cuò)誤類(lèi)別更安全。置信度閾值在驗(yàn)證集上掃描一遍再定0.6或0.7看不同閾值下被標(biāo)記為需復(fù)核的樣本數(shù)量、以及復(fù)核樣本里的誤檢率選一個(gè)業(yè)務(wù)能接受的操作點(diǎn)。我在設(shè)備色彩偏移的新數(shù)據(jù)集上翻過(guò)車(chē)教訓(xùn)是驗(yàn)證集只能證明模型在同類(lèi)數(shù)據(jù)上有效真要在新采集設(shè)備上跑還得單獨(dú)留一批數(shù)據(jù)做上線前驗(yàn)證。希望這些從數(shù)據(jù)集結(jié)構(gòu)到驗(yàn)證集用法的經(jīng)驗(yàn)?zāi)苷嬲龓偷侥?。本文還有配套的精品資源點(diǎn)擊獲取