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