:UNet原理與訓(xùn)練調(diào)參全攻略)
簡介面向Python圖像處理開發(fā)者的U-Net圖像分割實踐資源包覆蓋從數(shù)據(jù)準備、模型構(gòu)建、損失函數(shù)選擇到訓(xùn)練與預(yù)測的完整流程。資源定位在Python開發(fā)與圖片處理方向適合具備一定深度學(xué)習(xí)基礎(chǔ)、想動手實現(xiàn)語義分割任務(wù)的讀者。壓縮包內(nèi)共21個文件除示例圖片、說明文檔、Python腳本和對比動圖外還包含依賴清單與許可文件整體約5.6MB其中一個Python腳本實現(xiàn)圖像塊預(yù)測結(jié)果的平滑融合可將大尺寸遙感影像切塊推理后無縫拼接避免邊界痕跡。配合衛(wèi)星圖像分割示例、訓(xùn)練前后效果對比圖與動圖資源能直觀展示U-Net對稱編碼解碼結(jié)構(gòu)、跳躍連接在保留細節(jié)與定位邊界上的優(yōu)勢對醫(yī)療影像分析、自動駕駛感知和遙感地物分類等場景均有可復(fù)現(xiàn)的參考價值。目前已有10157人學(xué)習(xí)下載適合邊看邊練、對照代碼理解圖像分割原理的開發(fā)者。1. Python圖像分割為什么繞不開UNet一個能直接落地的基礎(chǔ)網(wǎng)絡(luò)很多人在拿到分割任務(wù)時第一個想到的往往是某個剛刷榜的新模型但一旦到了真實數(shù)據(jù)集上跑不穩(wěn)、訓(xùn)不動、修不完最后回頭換回UNet反而效果更好。這個現(xiàn)象在醫(yī)學(xué)影像、遙感、工業(yè)質(zhì)檢和廣告牌分割這類場景里反復(fù)出現(xiàn)原因很簡單UNet的U形結(jié)構(gòu)和跳躍連接在標注樣本有限時依然能穩(wěn)定收斂而Python生態(tài)里從數(shù)據(jù)加載到訓(xùn)練再到部署的每一環(huán)都有成熟庫可用。這篇文章不講花哨的刷分技巧而是圍繞“Python-使用UNet進行圖像分割”這條主線把網(wǎng)絡(luò)原理、數(shù)據(jù)準備、訓(xùn)練調(diào)參、常見翻車點和驗證技巧一次講透適合剛?cè)胧址指钊蝿?wù)或者想把UNet真正用起來的開發(fā)者。2. 先看懂UNet在做什么U形結(jié)構(gòu)、跳躍連接與特征圖尺寸2.1 編碼器把圖像“壓縮”成高維語義解碼器把它“還原”成像素級分類UNet的名字來自它的結(jié)構(gòu)像字母U。左邊是編碼器不斷做卷積和池化特征圖的寬高逐漸減半、通道數(shù)逐漸加倍網(wǎng)絡(luò)在這個過程中把“哪里有目標、目標是什么”的語義學(xué)到手。右邊是解碼器把編碼器輸出的低分辨率特征圖逐步上采樣還原回原圖尺寸同時把每個像素的分類結(jié)果輸出成與輸入相同長寬的mask。落地時最常用的骨干是ResNet18或ResNet34因為它們的預(yù)訓(xùn)練權(quán)重容易拿到顯存占用也比VGG16友好。如果做的是二維圖像分割輸入通常是“通道數(shù)×高×寬”比如RGB圖像就是3×512×512。編碼器下采樣4次之后特征圖變成16×16左右這時空間信息丟了很多但語義信息最密集。2.2 跳躍連接為什么是UNet的命根子如果沒有跳躍連接解碼器只能靠編碼器最后一層的高維特征去還原細節(jié)就像只憑一句話復(fù)述一張照片邊緣和小物體會全部糊掉。UNet的跳躍連接把編碼器每一層下采樣前的特征圖直接拼接到解碼器對應(yīng)層上讓細節(jié)信息繞過深層直接參與重建。拼接用的是torch.cat維度是通道維。編碼器某一層輸出是256×64×64解碼器同層的張量也是256×64×64拼接后變成512×64×64再接一次卷積降回256。這個操作的代價是顯存翻倍所以很多改進版UNet把拼接改成逐元素相加效果略降但省顯存。第一次寫UNet時建議直接拼因為相加的改進需要搭配殘差結(jié)構(gòu)才不丟精度。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)這段代碼是UNet里最基礎(chǔ)的卷積塊兩個3×3卷積加BatchNorm和ReLU。padding1保證特征圖尺寸不變inplaceTrue省一點顯存。BatchNorm在batch size較小時可能不穩(wěn)定如果顯存只允許放2張圖可以考慮換GroupNorm。2.3 三個最容易改錯的結(jié)構(gòu)參數(shù)第一是輸入尺寸。UNet本身不限制輸入尺寸但下采樣次數(shù)決定了最小特征圖尺寸。原版下采樣4次輸入512時最小特征圖是32×32夠用如果輸入只有128最小特征圖變成8×8語義信息丟失嚴重這時應(yīng)該減少下采樣次數(shù)。第二是通道數(shù)基數(shù)常見從32或64起步數(shù)據(jù)集小就選32數(shù)據(jù)集大選64通道數(shù)翻倍規(guī)則保持2的冪次。第三是上采樣方式轉(zhuǎn)置卷積和雙線性插值各有各的坑。轉(zhuǎn)置卷積有可學(xué)習(xí)參數(shù)但容易產(chǎn)生棋盤偽影雙線性插值沒有參數(shù)圖像更平滑。醫(yī)學(xué)分割常用轉(zhuǎn)置卷積工業(yè)分割場景我一般用雙線性插值加一次卷積來恢復(fù)通道省參數(shù)也穩(wěn)。3. 用UNet跑通第一版圖像分割數(shù)據(jù)準備到訓(xùn)練的最小路徑3.1 數(shù)據(jù)集怎么擺目錄結(jié)構(gòu)一次到位很多人寫到訓(xùn)練代碼才發(fā)現(xiàn)數(shù)據(jù)加載和mask對齊是最大的坑。常見做法是把原圖和標簽放在同一個根目錄下按前綴名配對img_001.jpg對應(yīng)mask_001.png。標簽圖必須是單通道PNG像素值從0開始連續(xù)編號0是背景1是第一類2是第二類。如果標簽是調(diào)色板PNG或三通道RGB先轉(zhuǎn)成單通道再進網(wǎng)絡(luò)。data/ ├── images/ │ ├── img_001.jpg │ ├── img_002.jpg └── masks/ ├── mask_001.png └── mask_002.png目錄擺好后寫一個Dataset類讀取這兩個文件夾。因為分割任務(wù)通常不需要shuffle文件名直接按文件名排序配對即可。3.2 數(shù)據(jù)加載與增強的落地寫法import cv2 import numpy as np from torch.utils.data import Dataset class SegDataset(Dataset): def __init__(self, img_dir, mask_dir, size(512, 512)): self.img_paths sorted(glob.glob(img_dir /*.jpg)) self.mask_paths sorted(glob.glob(mask_dir /*.png)) self.size size def __len__(self): return len(self.img_paths) def __getitem__(self, idx): img cv2.imread(self.img_paths[idx]) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img cv2.resize(img, self.size) mask cv2.imread(self.mask_paths[idx], cv2.IMREAD_GRAYSCALE) mask cv2.resize(mask, self.size, interpolationcv2.INTER_NEAREST) img img.astype(np.float32) / 255.0 img torch.from_numpy(img).permute(2, 0, 1) mask torch.from_numpy(mask).long() return img, maskmask的resize插值必須用INTER_NEAREST否則類別邊界會混入不存在的中間值。原圖resize可以用線性插值但標簽必須最近鄰。mask.long()是因為PyTorch的交叉熵損失要求target是LongTensor類別值必須在0到num_classes-1之間。輸入圖像這里直接用255歸一化配合torchvision的Normalize用ImageNet均值時要確保順序是先歸一化再標準化。3.3 損失函數(shù)與評估指標怎么選分割任務(wù)最常見的是交叉熵損失類別不平衡時用帶權(quán)重的交叉熵按每類的像素占比算中位數(shù)頻率作為權(quán)重。如果目標是細長結(jié)構(gòu)或小目標Dice Loss效果更好它直接優(yōu)化區(qū)域重疊度梯度對類別不平衡不敏感。實際項目中常用組合損失0.5 * BCE DiceBCE保持像素級梯度流Dice拉高區(qū)域一致性。def dice_loss(pred, target, smooth1.0): pred torch.softmax(pred, dim1) # 取第1類之后的所有前景類或根據(jù)具體類別調(diào)整 pred_fg pred[:, 1:] target_fg target[:, None, :, :].float() intersection (pred_fg * target_fg).sum() union pred_fg.sum() target_fg.sum() return 1 - (2 * intersection smooth) / (union smooth)這里smooth加在分子和分母上防止除零。target是LongTensor需要變成float才能參與乘法。這個函數(shù)適合二分類分割多分類時需要對每個類別算Dice再取平均。訓(xùn)練期間記錄每輪loss之外還建議保存每輪的mIoU只看loss容易漏掉過擬合點。4. UNet模型改進從普通UNet到ResUNet與注意力機制4.1 殘差連接解決深層網(wǎng)絡(luò)退化原版UNet的DoubleConv在層數(shù)加深以后梯度回傳容易衰減尤其在編碼器最深層。ResUNet的思路是在每個卷積塊外加一條恒等映射讓梯度可以直接從解碼器傳到編碼器淺層。改法很直接把DoubleConv的forward改成return self.conv(x) x前提是輸入輸出通道數(shù)一致。class ResDoubleConv(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) ) self.shortcut nn.Sequential() if in_ch ! out_ch: self.shortcut nn.Conv2d(in_ch, out_ch, 1) def forward(self, x): return nn.ReLU(inplaceTrue)(self.conv(x) self.shortcut(x))通道數(shù)不一致時用1×1卷積做shortcut保持一致后相加再ReLU。這個改動幾乎不增加參數(shù)量但收斂速度明顯變快特別是batch size較小時BatchNorm不穩(wěn)定殘差路徑能緩解梯度抖動帶來的loss波動。4.2 注意力門控讓網(wǎng)絡(luò)只關(guān)注目標區(qū)域許多場景下背景像素占比超過90%普通UNet會把大量計算浪費在背景上小目標區(qū)域?qū)W不到。Attention UNet的做法是在跳躍連接拼接前給編碼器特征圖乘一個注意力權(quán)重該權(quán)重由解碼器的高層特征生成相當于告訴網(wǎng)絡(luò)“這一塊才值得看”。常用實現(xiàn)是Attention Gate核心公式為W sigmoid(phi(g) psi(x))其中g(shù)是解碼器門控信號x是編碼器特征phi和psi各是一個1×1卷積。生成的權(quán)重圖與原特征圖逐元素相乘再進入拼接操作。比起直接拼接原始特征網(wǎng)絡(luò)對前景區(qū)域的響應(yīng)更集中小目標分割的mIoU通常能提升2到4個點。4.3 用深度可分離卷積做輕量化改進如果模型要部署在CPU或嵌入式設(shè)備上可以把標準3×3卷積替換成深度可分離卷積先按通道做3×3卷積再用1×1卷積混合通道。參數(shù)量大約是原來的九分之一速度在CPU上能快一倍以上。代價是精度略降一般配合殘差連接彌補。class SeparableConv2d(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.depthwise nn.Conv2d(in_ch, in_ch, 3, padding1, groupsin_ch) self.pointwise nn.Conv2d(in_ch, out_ch, 1) def forward(self, x): return self.pointwise(self.depthwise(x))groupsin_ch是深度卷積的關(guān)鍵每個通道獨立做卷積不跨通道混合。pointwise卷積再把所有通道信息融合。替換時注意BatchNorm要放在每個卷積之后不能兩個卷積共用一個BN。5. UNet使用中的常見問題與避坑排查5.1 顯存溢出換小輸入還是減通道數(shù)現(xiàn)象訓(xùn)練到第一個epoch直接報CUDA out of memory。原因通常是輸入尺寸太大或batch size設(shè)得過高UNet因為跳躍連接會保存編碼器各層特征圖顯存占用是同尺寸分類網(wǎng)絡(luò)的4到6倍。解決先減batch size到2還溢出就降輸入尺寸到384或256或者把編碼器起始通道從64改成32。不要一上來就開混合精度AMP能省約三分之一顯存但BatchNorm在fp16下容易出現(xiàn)數(shù)值不穩(wěn)反而更難排查。5.2 loss不下降或降得很慢先查標簽現(xiàn)象訓(xùn)練20個epoch后loss幾乎不動驗證集mIoU在0.1以下。自己看mask圖往往發(fā)現(xiàn)邊緣沒問題。常見原因是標簽類別不是從0開始連續(xù)編號比如原圖標注是1和3類別0缺失Softmax輸出的第0類永遠學(xué)不到東西。解決寫個腳本統(tǒng)計np.unique(mask)確認類別集合是[0, 1, 2, ...]。還有一類問題是mask通道數(shù)不對三通道BGR的標簽圖直接當單通道讀讀出來的值是三個通道的混合結(jié)果類別數(shù)瞬間膨脹。必須用cv2.IMREAD_GRAYSCALE讀。5.3 邊緣粗糙和空洞上采樣方式與后處理現(xiàn)象預(yù)測結(jié)果整體形狀對但邊緣像鋸齒內(nèi)部有小洞。原因是轉(zhuǎn)置卷積產(chǎn)生了棋盤偽影或者交叉熵損失每個像素獨立決策缺少區(qū)域約束。解決把上采樣換成雙線性插值加卷積或在loss里加Dice項。后處理可以用形態(tài)學(xué)閉運算補洞但注意閉運算會連帶填充真實空洞小目標多就不要用。更穩(wěn)妥的方式是CRF作為后處理但對大批量推理速度影響太大一般只用于離線評測。5.4 過擬合分割任務(wù)的泛化陷阱現(xiàn)象訓(xùn)練loss越降越低驗證集mIoU反而下降從第30個epoch開始差異明顯。分割數(shù)據(jù)集往往只有幾百張圖UNet參數(shù)多過擬合來得很早。解決順序先加數(shù)據(jù)增強隨機水平翻轉(zhuǎn)、隨機旋轉(zhuǎn)、隨機亮度對比度調(diào)整這幾項對大多數(shù)場景有效其次把Dropout加在解碼器最后一層前最后才是減小模型通道數(shù)。不要一開始就換預(yù)訓(xùn)練權(quán)重輕量數(shù)據(jù)增強的收益通常比換權(quán)重更大。5.5 編程環(huán)境的坑python安裝與cv2/numpy不匹配現(xiàn)象代碼在本機能跑換個環(huán)境后cv2.imread讀出的圖是None或者numpy和opencv版本沖突。原因多數(shù)是python版本與opencv-python的wheel不匹配比如python 3.8配新版opencv容易出現(xiàn)二進制不兼容。解決固定依賴版本用pip install opencv-python4.5.5.64 numpy1.23.5 torch1.13.1這類組合四個主庫版本對齊后基本不會再出兼容問題。另一點是cv2讀取中文路徑會失敗Windows下使用cv2.imdecode(np.fromfile(path, dtypenp.uint8), cv2.IMREAD_COLOR)替代。6. 用預(yù)測結(jié)果做驗證mIoU計算與生成分割圖模型訓(xùn)練完不等于落地還要看驗證集上的預(yù)測效果和數(shù)值指標。很多人只用loss判斷模型好壞但loss下降不代表像素分類準確尤其是類別不平衡時loss可能被背景主導(dǎo)。正確做法是寫一個評估腳本在驗證集上完整跑一遍前向推理逐圖計算每個類別的IoU再取所有類別的平均值作為mIoU。這個指標能直觀反映模型對大目標和小目標的綜合表現(xiàn)。def compute_miou(pred, target, num_classes): iou_list [] for cls in range(num_classes): p (pred cls) t (target cls) intersection (p t).sum() union (p | t).sum() if union 0: iou_list.append(float(nan)) else: iou_list.append(intersection / union) mean_iou np.nanmean(iou_list) return mean_iou這段代碼的核心是逐類別計算交集和并集。union為0意味著該類別在當前圖中完全不存在此時記nan并跳過避免把該類別的IoU算成0導(dǎo)致整體mIoU被壓下去。在驗證集中每個類別至少要出現(xiàn)一次否則對應(yīng)類別永遠不參與評分模型就會完全放棄學(xué)習(xí)這個類別。跑完mIoU還要看一眼實際分割圖尤其是邊界區(qū)域。用mask_overlay cv2.addWeighted(img, 0.7, color_mask, 0.3, 0)把預(yù)測mask疊加到原圖上目視檢查邊緣是否貼合、是否存在小碎塊。訓(xùn)練結(jié)束前我會固定使用同一批測試圖做對比每輪迭代之后保存預(yù)測圖做成GIF看變化這樣能直觀看到模型從哪個epoch開始變好、從哪個epoch開始過擬合。最后的習(xí)慣是把最優(yōu)epoch的權(quán)重單獨備份一份不要覆蓋訓(xùn)練中期的模型。很多時候測試集上的表現(xiàn)最好點并不在最后一個epoch早期checkpoint可能是更好的部署候選。用torch.save(model.state_dict(), unet_best.pth)保存并附帶一個記錄mIoU和epoch數(shù)值的JSON文件這樣回頭復(fù)盤時知道那版模型是在什么狀態(tài)下產(chǎn)出的。這種留痕習(xí)慣幫我避免過多次“模型找不回來”的翻車也希望幫你在UNet落地路上少走一段彎路。本文還有配套的精品資源點擊獲取