:從切片數(shù)據(jù)集到PyTorch訓練推理)
簡介面向眼底血管分割任務的Unet完整實踐項目包含已切片好的眼底圖像數(shù)據(jù)集、可直接運行的訓練與推理代碼以及訓練得到的結果文件適合醫(yī)學圖像分割入門和需要快速跑通完整流程的開發(fā)者。壓縮包共216個文件、約153.92MB以182個png眼底圖像為主數(shù)據(jù)另有8個py訓練/預測腳本、14個pyc編譯文件、txt標簽/日志、xml配置及pth權重文件目錄結構清晰已有269人學習下載。項目僅訓練10個epochs就在眼底血管二分割上取得全局像素準確率0.95、miou 0.67加大輪次后性能仍有提升空間。代碼支持數(shù)據(jù)隨機多尺度縮放與cos學習率衰減并能根據(jù)mask灰度值自動配置Unet輸出通道訓練日志及損失/IOU曲線保存于run_results便于逐類分析指標。推理時僅需將圖像放入inference目錄并運行predict腳本小白也能快速上手。1. 說清眼底血管分割這件事眼底血管分割是醫(yī)學圖像分析中的一個經(jīng)典任務輸入一張彩色眼底照片輸出每個像素屬于血管還是背景的概率圖。血管的粗細差異很大末梢血管只有幾個像素寬對比度又低因此落地方案通常不像自然圖像分割那樣直接扔給大模型而是先裁剪成patch再讓UNet去學。標題里的“切片好的數(shù)據(jù)集”指的就是這類已經(jīng)按固定窗口切好的訓練樣本配合完整代碼和訓練結果文件能少走很多彎路。它面向要跑通UNet眼底血管分割的工程師和研究生覆蓋從數(shù)據(jù)集加載、模型定義、訓練指標到推理后處理的完整鏈路。如果你有DRIVE、CHASE或本地眼底圖按文中的文件組織就能直接訓練。2. UNet結構解析與眼底血管分割的適配點2.1 編碼器下采樣血管分割需要多大的感受野UNet從FCN發(fā)展而來主體分編碼器和解碼器。編碼器通過4次下采樣把輸入從HxW降到H/16xW/16特征通道數(shù)從64增加到512。下采樣對眼底血管分割有三個實際作用一是讓卷積核看到更大范圍的視網(wǎng)膜背景從而區(qū)分血管和出血點二是減少后續(xù)計算量三是迫使模型學到不同尺度的血管響應。血管在眼底圖像上既有跨越半個視野的動脈主干也有只有兩三個像素寬的毛細血管單靠單一尺度卷積無法同時覆蓋兩種目標。階段操作序列輸出分辨率關注特征C13x3 conv, BN, ReLU x2H/1 x W/1 x 64血管邊緣、紋理P1D2maxpool conv blockH/2 x W/2 x 128局部血管走向P2D3maxpool conv blockH/4 x W/4 x 256分叉與交叉P3D4maxpool conv blockH/8 x W/8 x 512大血管區(qū)域bridgeconv blockH/16 x W/16 x 512全局上下文注意這里沒有繼續(xù)池化到H/32。眼底血管分割不是前景分類過大的下采樣會讓最細的血管在特征圖上直接消失保留1/16分辨率作為編碼器最深處是常見折中。如果顯存比較緊張可以把最深處設為H/8解碼器也相應減少一層但感受野變小后大出血塊和血管容易混在一起。改結構時要同步修改后面模型的forward內skip數(shù)量不能只改一處。2.2 解碼器與跳躍連接薄血管恢復的關鍵解碼器每次先對底層特征做2倍上采樣然后與編碼器同分辨率特征拼接再進行兩個3x3卷積。跳躍連接的貢獻不只是補細節(jié)它把編碼器前期的空間坐標信息直接傳給解碼器。血管末梢只有2-3個像素時僅靠深層語義無法定位精確邊界所以concat比sum更常用。在UNet的眼底血管分割實現(xiàn)里不要輕易刪除早期跳躍連接如果為了減少參數(shù)至少保留C1或C2層否則邊緣預測會明顯變粗。上采樣可以選用轉置卷積或雙線性插值。轉置卷積帶可學習參數(shù)能恢復更多紋理但也更容易產生棋盤偽影我通常用雙線性插值加上后面的卷積層血管邊緣反而更平滑。PyTorch里使用F.interpolate還是ConvTranspose2d會直接影響最后輸出的像素級精度。若發(fā)現(xiàn)預測圖出現(xiàn)一格一格的紋理優(yōu)先把解碼器的轉置卷積換成雙線性上采樣。2.3 一個能跑的UNet定義訓練時可改的3個參數(shù)下面是經(jīng)典UNet的PyTorch實現(xiàn)。代碼省去了注意力、空洞卷積等改造先保證能作為基線跑通后面要換backbone時只需要調整Encoder部分。# unet.py import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.conv nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), ) def forward(self, x): return self.conv(x) class UNet(nn.Module): def __init__(self, in_channels3, out_channels1, features(64, 128, 256, 512)): super().__init__() self.downs nn.ModuleList() self.ups nn.ModuleList() self.pool nn.MaxPool2d(2) # 編碼器逐級提取特征 for f in features: self.downs.append(DoubleConv(in_channels, f)) in_channels f # 最深層 bridge self.bridge DoubleConv(features[-1], features[-1] * 2) # 解碼器轉置卷積上采樣 跳躍連接 for f in reversed(features): self.ups.append(nn.ConvTranspose2d(f * 2, f, kernel_size2, stride2)) self.ups.append(DoubleConv(f * 2, f)) self.out_conv nn.Conv2d(features[0], out_channels, 1) def forward(self, x): skip [] for down in self.downs: x down(x) skip.append(x) x self.pool(x) x self.bridge(x) skip skip[::-1] for i in range(0, len(self.ups), 2): x self.ups[i](x) x torch.cat([x, skip[i // 2]], dim1) x self.ups[i 1](x) return self.out_conv(x)這段網(wǎng)絡不包含任何注意力機制但足夠作為眼底血管分割的基線。三個能直接改的參數(shù)分別是in_channelsout_channelsfeatures。in_channels對應輸入圖像通道數(shù)彩色眼底圖是3灰度眼底圖改為1out_channels是分割輸出的通道數(shù)二分類血管分割用1而不是2features代表編碼器各層通道數(shù)默認是(64,128,256,512)。如果顯存不高可以改成(32,64,128,256)模型參數(shù)會明顯下降但通常需要更多epoch才能追平效果。這里有個容易忽略的地方解碼器里bridge的輸出通道是features[-1]2而ups列表里第一個轉置卷積的輸入通道也必須是f2。修改features時只要保持f*2一致即可否則會在forward階段出現(xiàn)channels mismatch。這個錯誤在PyTorch中要等到實際跑數(shù)據(jù)時才報出來建議寫完模型后用隨機輸入跑一次前向傳播再開始訓練。3. 切片數(shù)據(jù)集的整理與加載從原圖到訓練樣本3.1 切片好的數(shù)據(jù)集長什么樣目錄約定與樣本對應關系標題里的“切片好的數(shù)據(jù)集”通常不是一張完整眼底圖而是一批固定size的patch。常見目錄結構如下images放彩色patchmasks放黑白掩膜兩個目錄的文件名一一對應。data/ ├── images/ │ ├── 01_00000.png │ ├── 01_00001.png │ └── ... ├── masks/ │ ├── 01_00000.png │ ├── 01_00001.png │ └── ... ├── train.txt └── val.txt文件名前綴由原圖編號和patch坐標組成例如01_00000表示第1張原圖的左上角。train.txt和val.txt每行一個文件名前綴模型讀取時按這個列表加載切片。訓練集和驗證集必須分到不同的原圖不能只分patch。如果同一張原圖的不同patch同時出現(xiàn)在訓練和驗證中模型會通過感受野外的周邊信息記住圖像導致驗證結果虛高。切片好的數(shù)據(jù)集和直接輸入原圖的區(qū)別在于模型看到的是局部視野。視網(wǎng)膜血管橫向跨度很大一個256x256的patch有時只包含一根血管主干的一部分。遇到這種情況不要在patch里硬分前景和背景而是依賴推理時滑動窗口的重疊預測。切片參數(shù)沒有絕對標準下面是我常用的參考值。參數(shù)小顯存中等顯存說明patch_size128256需要能被模型下采樣次數(shù)整除建議是16的倍數(shù)stridepatch_size - 32patch_size - 64重疊越多訓練樣本越多但訓練更慢過濾閾值1%1%全背景切片保留少量即可存儲格式pngpng不要用jpeg血管邊緣會失真設置patch_size時要避開一個坑UNet每層做2倍下采樣一共4次所以patch的長寬最好是16的倍數(shù)。如果尺寸不是16的倍數(shù)上采樣后的特征圖尺寸會和跳躍連接層差1個像素雖然有些情況仍能跑但整體分割結果會在邊緣處錯位。3.2 切片代碼固定步長和重疊切片如果沒有拿到別人切好的數(shù)據(jù)自己從原始眼底圖制作也很簡單。下例以256x256的patch、192的步長把大圖和掩膜同步切片。# make_patches.py import os import random from PIL import Image import numpy as np def extract_patches(image_path, mask_path, save_dir_img, save_dir_mask, patch_size256, stride192, keep_ratio0.1): image np.array(Image.open(image_path).convert(RGB)) mask np.array(Image.open(mask_path).convert(L)) h, w mask.shape name os.path.splitext(os.path.basename(image_path))[0] idx 0 for y in range(0, h - patch_size 1, stride): for x in range(0, w - patch_size 1, stride): img_patch image[y:ypatch_size, x:xpatch_size] mask_patch mask[y:ypatch_size, x:xpatch_size] foreground_ratio (mask_patch 0).mean() * 100 # 過濾完全無血管的切片按比例保留背景樣本 if foreground_ratio 1 and random.random() keep_ratio: continue Image.fromarray(img_patch).save( os.path.join(save_dir_img, f{name}_{idx:05d}.png)) Image.fromarray(mask_patch).save( os.path.join(save_dir_mask, f{name}_{idx:05d}.png)) idx 1代碼以stride192在256的patch上移動patch與patch之間有64像素重疊。重疊的作用是讓血管在多個patch中都完整出現(xiàn)推理拼接時也不會在patch邊界留下折痕。keep_ratio控制全背景切片的保留比例如果全部丟掉模型容易把視盤周圍的暗區(qū)誤判為血管所以保留10%左右的背景樣本會更穩(wěn)。如果你的數(shù)據(jù)集切片后樣本數(shù)已經(jīng)很大可以去掉這個過濾條件直接讓模型學習背景分布。保存時統(tǒng)一用PNG。PNG是無損格式掩膜的邊界不會有JPEG壓縮帶來的灰邊。若原始掩膜本身就是JPEG建議在切片前用形態(tài)學閉運算把斷裂的細小血管連接一次否則模型會把斷裂當作學習目標。3.3 Dataset加載與增強血管分割里哪些增強可以開切片做好后寫一個torch.utils.data.Dataset把patch讀進來。常見做法是加載img和mask后一起做隨機翻轉和旋轉。對于眼底血管分割顏色抖動要謹慎血管顏色是重要特征過度改變顏色會讓模型不魯棒。推薦開啟的增強包括水平翻轉、垂直翻轉、90度旋轉和輕度隨機仿射彈性形變對醫(yī)學圖像很有效但注意位移不要超過3個像素。# dataset.py import os import random import numpy as np import torch from torch.utils.data import Dataset from PIL import Image class VesselDataset(Dataset): def __init__(self, img_dir, mask_dir, file_list, trainTrue): self.img_dir img_dir self.mask_dir mask_dir self.samples [line.strip() for line in open(file_list)] self.train train def __len__(self): return len(self.samples) def __getitem__(self, idx): name self.samples[idx] img np.array(Image.open( os.path.join(self.img_dir, name .png)).convert(RGB)).astype(np.float32) mask np.array(Image.open( os.path.join(self.mask_dir, name .png)).convert(L)) mask (mask 127).astype(np.float32) if self.train: # 翻轉時圖像和掩膜必須同步 if random.random() 0.5: img img[:, ::-1, :] mask mask[:, ::-1] if random.random() 0.5: img img[::-1, :, :] mask mask[::-1, :] k random.choice([0, 1, 2, 3]) if k: img np.rot90(img, k, axes(0, 1)) mask np.rot90(mask, k, axes(0, 1)) img img / 255.0 img torch.from_numpy(img.transpose(2, 0, 1)) mask torch.from_numpy(mask).unsqueeze(0) return img, mask加載時把mask二值化大于127視為血管其余視為背景。增強順序是先翻轉再旋轉順序不能調換否則坐標對應關系會亂。這個類沒有做resize要求所有切片已經(jīng)統(tǒng)一成patch_size。如果數(shù)據(jù)集來自多個來源尺寸不一致需要在__getitem__里補上Resize但mask必須用最近鄰插值不能用線性插值否則血管邊緣會出現(xiàn)中間灰值。驗證時trainFalse只做標準化不做增強。4. 訓練UNet損失函數(shù)、評價指標與訓練結果文件的輸出4.1 損失函數(shù)為什么用BCE加Dice組合眼底血管分割是像素二分類。BCE容易優(yōu)化但正負樣本不平衡血管像素通常只占10%左右訓練初期模型會傾向把所有像素預測為背景。Dice Loss直接優(yōu)化前景/背景重疊對小目標更敏感但單獨使用梯度不平滑。常見做法是讓兩者相加total_loss bce_loss dice_loss。# losses.py import torch import torch.nn.functional as F def mixed_loss(pred, target): pred torch.sigmoid(pred) bce F.binary_cross_entropy(pred, target, reductionmean) smooth 1e-5 pred_flat pred.reshape(pred.size(0), -1) target_flat target.reshape(target.size(0), -1) intersection (pred_flat * target_flat).sum(dim1) dice 1 - (2 * intersection smooth) / ( pred_flat.sum(dim1) target_flat.sum(dim1) smooth) return bce dice.mean()smooth的作用是防止分母為0一般取1e-5。Dice部分沒有做one-hot因為血管分割只有一個前景類。這里返回的是bce dice.mean()如果只想對前景加權可以在BCE里給pos_weight傳一個小數(shù)比如1.2但通常組合損失已經(jīng)夠用。一個關鍵點是pred要先經(jīng)過sigmoid不能用logits直接算Dice。訓練時如果發(fā)現(xiàn)dice loss出現(xiàn)nan先檢查標簽是不是0/255而不是0/1。Dataset里已經(jīng)做了mask 127所以target_flat和pred_flat數(shù)值尺度一致。另一個來源是batch里某張mask全為零smooth會阻止除零但如果smooth忘了加loss就會變成無窮大。4.2 評價指標計算F1、敏感度和特異度訓練時只看loss不夠還需要每幾個epoch在驗證集上評估。眼底血管分割常用指標有accuracy、F1、sensitivity和specificity。sensitivity反映血管漏檢率specificity反映背景誤判率兩個指標都在固定閾值下計算。如果驗證集F1很高但sensitivity低說明模型只分割了大血管漏掉了末梢小血管這時要降低預測閾值或增加薄血管切片權重。# metrics.py def compute_metrics(pred_sigmoid, label, threshold0.5): pred (pred_sigmoid threshold).float() label label.float() tp (pred * label).sum().item() fp (pred * (1 - label)).sum().item() fn ((1 - pred) * label).sum().item() tn ((1 - pred) * (1 - label)).sum().item() sensitivity tp / (tp fn 1e-6) specificity tn / (tn fp 1e-6) precision tp / (tp fp 1e-6) f1 2 * precision * sensitivity / (precision sensitivity 1e-6) return { f1: f1, sensitivity: sensitivity, specificity: specificity, accuracy: (tp tn) / (tp tn fp fn) }計算指標時先把預測值轉成0/1再和標簽做逐像素比較。這里沒有用torchmetrics是為了減少訓練腳本的額外依賴。小批量驗證時可以在batch維度上累加tp/fp/fn/tn最后再算指標。如果發(fā)現(xiàn)驗證指標劇烈抖動先確認驗證集是不是只有幾十張patch樣本太少時sensitivity會受單張圖影響。建議至少保留200張patch做驗證集。4.3 訓練循環(huán)與checkpoint訓練結果文件怎么保存與恢復訓練循環(huán)可以分成train_one_epoch和evaluate兩個函數(shù)。訓練時model.train()驗證時model.eval()并包在torch.no_grad()里。保存模型時不要只存權重建議把epoch、optimizer和best_f1都放進一個字典。# train.py def train_one_epoch(model, loader, optimizer, criterion, device): model.train() total_loss 0.0 for img, mask in loader: img img.to(device) mask mask.to(device) pred model(img) loss criterion(pred, mask) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() * img.size(0) return total_loss / len(loader.dataset) # 保存與恢復邏輯 best_f1 0.0 for epoch in range(start_epoch, epochs): train_loss train_one_epoch(model, train_loader, optimizer, criterion, device) metrics evaluate(model, val_loader, criterion, device) if metrics[f1] best_f1: best_f1 metrics[f1] torch.save({ epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), best_f1: best_f1, }, best_model.pth) torch.save(model.state_dict(), last_model.pth)每個epoch結束后用驗證集F1判斷是否保存best_model.pth這是最常見的做法。訓練結果文件里包含optimizer狀態(tài)是為了之后從斷點恢復如果只做推理用torch.load加載后只取model_state_dict即可。建議把每個epoch的loss寫入csv曲線會顯示模型是真正收斂還是在震蕩。如果訓練時用了DataParallel保存的model_state_dict會帶module前綴。推理加載時把權重名字里的module.去掉再load_state_dict否則會報告unexpected key。這個問題通常在多卡訓練完、單卡推理時出現(xiàn)訓練結果文件越大越容易忽略。4.4 訓練環(huán)境與參數(shù)速查下面這套參數(shù)是從實際項目里提煉出來的起點。GPU顯存不同需要優(yōu)先改patch_size和features而不是只改batch_size。配置項小顯存6G中顯存12G說明patch_size128256顯存不夠優(yōu)先減小patchbatch_size168數(shù)值受數(shù)據(jù)加載速度影響features(32,64,128,256)(64,128,256,512)模型寬度減半顯存約降至1/4學習率1e-31e-3Adam分割任務常用1e-4到1e-3調度器CosineAnnealingLRCosineAnnealingLR避免后期loss震蕩epoch200150數(shù)據(jù)量小時需要更多輪次訓練環(huán)境只要支持CUDA即可PyTorch 2.x和1.x在這組代碼上沒有本質差別。顯存不夠時先把batch_size調成1再不行就降patch_size。patch從256降到128同batch下特征圖計算量約降為原來的1/4因為長寬各減半顯存下降會更明顯所以調整優(yōu)先級很高。5. 用訓練好的模型做推理驗證與誤分割處理5.1 重疊滑動窗口拼接模型訓練完成后驗證整圖不能把原圖直接輸入網(wǎng)絡因為顯存和patch訓練分布都不允許。常見做法是復用訓練時的patch_size以patch_size的一半作為步長滑動推理重疊區(qū)域取多次預測的平均值。拼接時維護一張prob_map和一張weight_map每個位置累加預測值和計數(shù)最后prob_map除以weight_map就得到整圖概率圖。這樣做的效果是patch邊界不會出現(xiàn)一字折痕小血管在重疊區(qū)域也會被預測兩到三次結果更連續(xù)。5.2 形態(tài)學后處理去碎屑、填孔洞概率圖轉二值圖后最常出現(xiàn)的兩個問題是視盤周圍被識別為血管以及細血管內部出現(xiàn)斷裂。用OpenCV做兩步處理先按連通域面積刪除小碎屑再做一次閉運算填補孔洞。import cv2 import numpy as np def postprocess(prob, threshold0.5, min_area10): binary (prob threshold).astype(np.uint8) * 255 n, labels, stats, _ cv2.connectedComponentsWithStats(binary, connectivity8) result np.zeros_like(binary) for i in range(1, n): if stats[i, cv2.CC_STAT_AREA] min_area: result[labels i] 255 return cv2.morphologyEx( result, cv2.MORPH_CLOSE, cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (5, 5)))min_area在普通眼底圖分辨率下取10能去掉零星噪點分辨率更高時建議按連通域外接圓半徑過濾。閉運算內核不要超過7否則會把相鄰血管粘連。后處理只作用于二值圖不要對概率圖做開閉運算否則概率值會整體偏移。5.3 閾值掃描與TTA的選擇最終概率圖不一定用0.5做閾值。在驗證集上掃描0.3到0.7找到F1最高的閾值保存到模型目錄推理時讀取。這一步比開TTA更簡單提升也更直接敏感度不足就調低閾值背景噪聲太多就調高閾值。如果離線分析時間充足再疊加一次TTA即對輸入patch做水平翻轉、垂直翻轉把三次預測翻轉回原方向取平均。TTA不改變模型參數(shù)只降低單次預測方差對末梢血管連續(xù)性有可見改善。部署到在線服務時優(yōu)先保留重疊推理不開TTA。把最優(yōu)閾值寫進模型配置文件后續(xù)復現(xiàn)時不需要重新掃。本文還有配套的精品資源點擊獲取