戰(zhàn):從環(huán)境配置到模型評估)
簡介針對希望在自己多類別數(shù)據(jù)集上快速上手Unet語義分割的PyTorch開發(fā)者這份資源提供了完整可運(yùn)行的工程范例覆蓋醫(yī)學(xué)影像、遙感圖像等場景可解決自定義數(shù)據(jù)集的類別通道設(shè)計(jì)、數(shù)據(jù)讀取與訓(xùn)練評估等關(guān)鍵問題。壓縮包共46個文件以19個Python源碼為主對應(yīng)模型構(gòu)建、訓(xùn)練與評估等核心邏輯配套24個pyc緩存、2個txt說明與1個json配置整體僅69KB代碼精簡但模塊劃分清晰包含dataloaders數(shù)據(jù)讀取、modeling網(wǎng)絡(luò)結(jié)構(gòu)、utils工具函數(shù)以及train.py、demo.py等可直接執(zhí)行的腳本。已有15245人學(xué)習(xí)瀏覽適合作為從單類別分割轉(zhuǎn)向多類別任務(wù)的參考實(shí)現(xiàn)。資源內(nèi)不僅有Unet模型定義和交叉熵?fù)p失函數(shù)還實(shí)現(xiàn)了數(shù)據(jù)增強(qiáng)、學(xué)習(xí)率調(diào)度、驗(yàn)證集指標(biāo)計(jì)算如IoU等關(guān)鍵環(huán)節(jié)并附有txt說明與配置文件便于復(fù)現(xiàn)訓(xùn)練流程并遷移到自有數(shù)據(jù)集上繼續(xù)改進(jìn)。 如果你也是那種“官方demo跑得飛起一換自己的數(shù)據(jù)集就各種報錯”的狀態(tài)那這篇文章正好是寫給你的。PyTorch U-Net做多類別語義分割本身是個非常成熟的技術(shù)組合但網(wǎng)上90%的教程都停留在“兩分類背景前景”或者“OpenCV畫幾個色塊當(dāng)mask”的階段。真到自己標(biāo)注幾百張圖、分十個八個類別、還要跑出像樣的指標(biāo)時問題一個接一個mask讀取出來是全黑的、loss不降、驗(yàn)證集mIoU奇低、顯存爆掉……這些問題我全踩過。這篇文章我想把我用PyTorch實(shí)現(xiàn)U-Net對自己多類別數(shù)據(jù)集做語義分割的完整鏈路寫清楚從環(huán)境依賴、網(wǎng)絡(luò)結(jié)構(gòu)理解、數(shù)據(jù)集制作、訓(xùn)練細(xì)節(jié)到踩坑排查和評估指標(biāo)全部基于實(shí)際跑通的經(jīng)驗(yàn)來寫。適合剛開始嘗試用深度學(xué)習(xí)做分割任務(wù)、有基礎(chǔ)Python語法知識但對訓(xùn)練細(xì)節(jié)還不熟悉的讀者。核心目標(biāo)只有一個讓你拿著自己的數(shù)據(jù)也能順利把U-Net跑通并且跑得規(guī)范。1. 環(huán)境搭建先解決PyTorch與CUDA的版本“暗坑”很多人忽略環(huán)境配置的坑覺得裝完anaconda再pip install torch就完事了。但語義分割訓(xùn)練對顯存和算力有硬性要求一個匹配錯誤的CUDA版本可能讓你明明裝了GPU版PyTorch一跑torch.cuda.is_available()卻返回False或者直接報CUDA driver version is insufficient。這類問題在論壇上幾乎每天都有新帖子。1.1 我自己驗(yàn)證過的穩(wěn)定組合根據(jù)多次排障經(jīng)驗(yàn)比較穩(wěn)妥的組合是在Anaconda里創(chuàng)建獨(dú)立Python虛擬環(huán)境用conda安裝CUDA toolkit和cuDNN再安裝對應(yīng)版本的PyTorch。我目前主力環(huán)境中使用的是Python 3.10.11 PyTorch 2.8.0 CUDA 12.1的組合整體穩(wěn)定訓(xùn)練和推理都沒遇到兼容性問題。創(chuàng)建環(huán)境和安裝PyTorch的參考命令conda create -n seg python3.10.11 -y conda activate seg conda install cuda -c nvidia pip install torch2.8.0 torchvision --index-url https://download.pytorch.org/whl/cu121裝完之后用這幾行代碼驗(yàn)證import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0))如果第二行輸出True說明GPU可用。如果輸出False先用nvidia-smi看驅(qū)動支持的CUDA版本再重新安裝匹配的PyTorch版本。這里有個很容易混淆的點(diǎn)nvidia-smi顯示的CUDA版本是驅(qū)動支持的上限不是當(dāng)前環(huán)境實(shí)際使用的CUDA版本PyTorch實(shí)際用的是它自己捆綁的CUDA runtime兩者不一定要完全一致但PyTorch要求的CUDA版本不能高于驅(qū)動支持的版本。1.2 顯存不夠時的備選方案如果你的顯卡顯存不足6GB訓(xùn)練原始尺寸的U-Net會很吃力。這時候有兩個選擇一是把訓(xùn)練圖像縮放或裁剪到256x256或512x512二是考慮用CPU訓(xùn)練但把num_workers調(diào)低、batch_size調(diào)小。CPU訓(xùn)練確實(shí)慢但處理小規(guī)模數(shù)據(jù)集幾百張圖時也不是不能接受。我在無GPU環(huán)境下用CPU跑過512x512輸入、4個類別的分割任務(wù)一個epoch大概需要8-10分鐘整體訓(xùn)練時間在一個小時左右屬于可以等待的范圍。2. U-Net結(jié)構(gòu)到底在做什么編碼器、解碼器與跳躍連接在寫Dataset和訓(xùn)練循環(huán)之前先把U-Net的網(wǎng)絡(luò)邏輯講清楚。很多人直接copy網(wǎng)上的U-Net實(shí)現(xiàn)卻不理解為什么最后一層要輸出num_classes個通道也不明白為什么跳躍連接能“保留細(xì)節(jié)”。理解這些你改代碼時才能知其然也知其所以然。2.1 編碼器與解碼器的分工U-Net的骨干部分由兩條路徑組成。編碼器左側(cè)通過重復(fù)的卷積池化/下采樣把特征圖分辨率降低通道數(shù)增加。這個過程在“提煉語義”低分辨率下模型更容易回答“這張圖里有哪些類別的物體”。解碼器右側(cè)通過上采樣逐步恢復(fù)分辨率把低分辨率的特征圖“放大”回原始尺寸。這個階段在回答“每一個像素到底屬于哪個類別”。如果只有這兩條路徑典型的問題就出來了上采樣是個信息丟失的過程。編碼器在捕捉大范圍語義時會把小目標(biāo)細(xì)節(jié)磨掉解碼器想恢復(fù)這些細(xì)節(jié)卻無能為力。U-Net的解決方案就是跳躍連接——把編碼器每個stage的輸出直接拼接到解碼器對應(yīng)stage的輸入上。2.2 跳躍連接的本質(zhì)細(xì)節(jié)特征的“快遞通道”跳躍連接把編碼器某層的高分辨率特征圖復(fù)制一份在通道維度上拼接到解碼器同尺度的特征圖上。這樣解碼器在進(jìn)行卷積時不僅能“看到”上采樣后較模糊的語義信息也能“看到”來自編碼器同分辨率分支的精細(xì)邊緣和紋理信息。有點(diǎn)像修圖時把模糊的大改圖和高清的原圖疊在一起做參考兩邊對齊后效果自然更好。理解了這一點(diǎn)你就能明白為什么U-Net在醫(yī)學(xué)影像、遙感影像這類“既要語義又要邊界”的任務(wù)里表現(xiàn)優(yōu)異而在一些簡單二分類任務(wù)上反而顯得大材小用。它的設(shè)計(jì)初衷就是針對細(xì)節(jié)敏感的分割場景。2.3 輸入輸出通道的設(shè)計(jì)邏輯實(shí)現(xiàn)U-Net時有兩個關(guān)鍵參數(shù)in_channels和num_classes。in_channels對應(yīng)輸入圖像的通道數(shù)RGB圖像就是3灰度圖就是1。num_classes對應(yīng)你要分割的類別總數(shù)包括背景類。比如你要分割“背景、道路、樹木、建筑”四類那num_classes4。最后一層輸出的就是一個(B, 4, H, W)的張量每個像素在4個通道上各有一個logit值。訓(xùn)練時用交叉熵?fù)p失計(jì)算這些logit和真實(shí)標(biāo)簽間的差異預(yù)測時取4個通道中數(shù)值最大的索引作為該像素的類別。下面這段核心代碼可以幫你理解這個映射關(guān)系# 網(wǎng)絡(luò)最后一層 self.outc nn.Conv2d(64, num_classes, kernel_size1) # 前向傳播 logits self.outc(x) # shape: (B, num_classes, H, W) # 預(yù)測階段 pred logits.argmax(dim1) # shape: (B, H, W)每個像素是類別索引3. 自制多類別數(shù)據(jù)集從標(biāo)注文件到可訓(xùn)練的Tensor很多教程會直接用現(xiàn)成的VOC或Cityscapes數(shù)據(jù)集但實(shí)際項(xiàng)目中幾乎都需要用自己的數(shù)據(jù)。U-Net訓(xùn)練所需的標(biāo)注格式并不復(fù)雜核心就是每張?jiān)瓐D對應(yīng)一張與它同尺寸的mask圖mask中每個像素的灰度值代表這個像素的類別編號。3.1 標(biāo)注工具與數(shù)據(jù)整理流程我用的是Labelme做多邊形標(biāo)注。標(biāo)注時圍繞目標(biāo)邊緣打點(diǎn)保存后生成JSON文件再通過腳本把JSON轉(zhuǎn)成8位灰度mask圖。具體過程是讀取JSON中的shapes字段用labelme.utils.shapes_to_label生成標(biāo)簽圖再轉(zhuǎn)成單通道的png保存。在整理數(shù)據(jù)時我建議按下面的目錄結(jié)構(gòu)來放data/ images/ img_001.jpg img_002.jpg masks/ img_001.png img_002.png這里有個非常關(guān)鍵、也是新手最常踩的坑mask圖不能保存為三通道的彩色PNG必須是單通道灰度圖或調(diào)色板模式。如果用彩色圖保存每個類別變成RGB三元組而不是單一灰度值訓(xùn)練時取像素值做標(biāo)簽就會完全錯亂。3.2 數(shù)據(jù)增強(qiáng)的“同步變換”陷阱語義分割的數(shù)據(jù)增強(qiáng)比分類任務(wù)要小心得多。分類任務(wù)里你對圖像做翻轉(zhuǎn)、裁剪就完事了但分割任務(wù)里必須保證原圖和mask圖做完全相同的幾何變換否則模型學(xué)到的就是錯誤對應(yīng)關(guān)系。實(shí)現(xiàn)同步變換時最省事的方法是使用albumentations庫它的設(shè)計(jì)天然支持image和mask同時變換import albumentations as A train_transform A.Compose([ A.HorizontalFlip(p0.5), A.RandomBrightnessContrast(p0.2), A.RandomCrop(512, 512), ])注意RandomBrightnessContrast這種顏色類增強(qiáng)只會作用在圖像上mask不受影響這正是我們想要的。而HorizontalFlip、RandomCrop這類空間變換需要對兩者同時生效albumentations會自動處理。3.3 Dataset類實(shí)現(xiàn)要點(diǎn)PyTorch中自定義Dataset需要繼承torch.utils.data.Dataset并實(shí)現(xiàn)__len__和__getitem__。我的實(shí)現(xiàn)里每個樣本返回三樣?xùn)|西原圖Tensor、mask長整型Tensor、以及文件名。返回文件名看起來多余但在推理階段輸出預(yù)測結(jié)果時會非常有用。from torch.utils.data import Dataset from PIL import Image import numpy as np import torch class SegDataset(Dataset): def __init__(self, image_dir, mask_dir, transformNone): self.image_paths sorted(glob.glob(f{image_dir}/*.jpg)) self.mask_dir mask_dir self.transform transform def __len__(self): return len(self.image_paths) def __getitem__(self, idx): img_path self.image_paths[idx] name os.path.basename(img_path).split(.)[0] image np.array(Image.open(img_path).convert(RGB)) mask np.array(Image.open(os.path.join(self.mask_dir, name .png)).convert(L)) if self.transform: augmented self.transform(imageimage, maskmask) image augmented[image] mask augmented[mask] image torch.from_numpy(image.transpose(2, 0, 1)).float() / 255.0 mask torch.from_numpy(mask).long() return image, mask, name用convert(L)讀取mask可以強(qiáng)制按單通道灰度讀取避免某些工具把灰度png按三通道讀出來。4. 訓(xùn)練流程中的關(guān)鍵細(xì)節(jié)損失函數(shù)、評估指標(biāo)與顯存控制網(wǎng)絡(luò)結(jié)構(gòu)和數(shù)據(jù)準(zhǔn)備好了接下來就是整個訓(xùn)練流程的搭建。這一部分直接決定模型能不能收斂、收斂后的效果好不好千萬別直接照搬分類任務(wù)的代碼。4.1 損失函數(shù)的選擇邏輯多類別語義分割最基礎(chǔ)的損失函數(shù)是nn.CrossEntropyLoss()。PyTorch的交叉熵?fù)p失接收兩個參數(shù)模型輸出的logits張量形狀為(B, num_classes, H, W)以及真實(shí)標(biāo)簽張量形狀為(B, H, W)每個位置的取值是0到num_classes-1的整數(shù)。如果類別分布很不均衡——比如背景占了80%你要分割的目標(biāo)很小——建議在交叉熵的基礎(chǔ)上加一個Dice Loss。Dice Loss在醫(yī)學(xué)影像分割里使用極為廣泛對小目標(biāo)的約束比交叉熵更強(qiáng)。實(shí)際使用中我傾向于用0.7 * CrossEntropyLoss 0.3 * DiceLoss這樣的組合。有個細(xì)節(jié)容易搞錯使用nn.CrossEntropyLoss時不需要對網(wǎng)絡(luò)輸出手動做softmax因?yàn)檫@個損失函數(shù)內(nèi)部已經(jīng)包含了softmax計(jì)算。如果你先在外面套了nn.Softmax(dim1)再傳給CrossEntropy訓(xùn)練時數(shù)值穩(wěn)定性反而會受影響。4.2 mIoU指標(biāo)與訓(xùn)練日志分類任務(wù)看accuracy就夠了但分割任務(wù)中存在嚴(yán)重的類別不均衡問題時accuracy往往具有欺騙性。比如背景占95%的場景模型把所有像素都預(yù)測為背景accuracy已經(jīng)95%了但目標(biāo)一個都沒分割出來。所以需要看mIoU——預(yù)測值和真實(shí)值兩個集合的交集與并集之比對每個類別計(jì)算后取平均。下面是自己計(jì)算的mIoU代碼在驗(yàn)證集上逐batch累加混淆矩陣再統(tǒng)一計(jì)算def compute_miou(pred_mask, true_mask, num_classes): iou_list [] for c in range(num_classes): pred_c (pred_mask c) true_c (true_mask c) intersection (pred_c true_c).sum().item() union (pred_c | true_c).sum().item() if union 0: iou_list.append(float(nan)) else: iou_list.append(intersection / union) return np.nanmean(iou_list)訓(xùn)練時每跑完一個epoch在驗(yàn)證集上算一次mIoU保存mIoU最高的模型權(quán)重。別只按loss保存loss最低不代表分割效果最好。4.3 batch_size、學(xué)習(xí)率與混合精度顯存不足時優(yōu)先調(diào)整的是batch_size而不是輸入尺寸。輸入尺寸直接影響感受野和分割精細(xì)度batch_size只要不是1稍微小一點(diǎn)影響沒那么大。我用6GB顯存訓(xùn)練U-Net時512x512輸入、batch_size設(shè)為4是可以跑起來的再加混合精度還能余出一些空間。推薦圖里跑代碼時同步開啟PyTorch的自動混合精度scaler torch.cuda.amp.GradScaler() for batch in dataloader: images, masks batch images, masks images.cuda(), masks.cuda() optimizer.zero_grad() with torch.cuda.amp.autocast(): logits model(images) loss criterion(logits, masks) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()混合精度一方面降低顯存占用另一方面在部分GPU上能明顯提速RTX 30系、40系顯卡效果明顯。不需要手動做什么額外操作PyTorch已經(jīng)把細(xì)節(jié)封裝好了。4.4 訓(xùn)練參數(shù)參考如果是第一次跑自己的數(shù)據(jù)我建議按下面這組參數(shù)起步優(yōu)化器AdamW初始學(xué)習(xí)率1e-4weight_decay1e-4學(xué)習(xí)率調(diào)度CosineAnnealingLRT_maxepochsepochs先跑50看loss和mIoU趨勢再調(diào)整batch_size根據(jù)顯存調(diào)整4到16之間輸入尺寸256x256或512x512用AdamW加weight_decay的做法是我在實(shí)際項(xiàng)目中逐步穩(wěn)定下來的組合。同樣數(shù)據(jù)集下純SGD收斂慢且后期不穩(wěn)定Adam不加weight_decay會在訓(xùn)練后期出現(xiàn)振蕩AdamW在精度和穩(wěn)定性之間平衡得最好。5. 推理階段從模型輸出到彩色分割圖的完整映射模型訓(xùn)練好之后推理階段的代碼邏輯也要清晰。我見過不少人訓(xùn)練時跑通了一到預(yù)測階段輸出怎么都變不成預(yù)期的彩色圖問題往往出在索引到顏色的映射上。這一步其實(shí)很簡單先取每個像素在所有類別通道上最大值對應(yīng)的索引再查表映射到預(yù)先定義的顏色。5.1 自定義調(diào)色板的生成與映射假設(shè)你標(biāo)注了5個類別可以自己定義顏色表import numpy as np # 每行是類別索引對應(yīng)的RGB顏色 palette np.array([ [0, 0, 0], # 背景 [255, 0, 0], # 類別1 [0, 255, 0], # 類別2 [0, 0, 255], # 類別3 [255, 255, 0], # 類別4 ], dtypenp.uint8)推理時把模型輸出轉(zhuǎn)成索引圖再通過調(diào)色板映射到RGB圖model.eval() with torch.no_grad(): logits model(image_tensor.unsqueeze(0).cuda()) pred logits.argmax(dim1).squeeze(0).cpu().numpy() # (H, W) # 索引圖 - 彩色圖 color_mask palette[pred]這里稍微解釋一下palette[pred]的工作原理pred是一個形狀為(H, W)的整數(shù)數(shù)組取值是0到4把它當(dāng)作索引去查palette數(shù)組得到的就是形狀(H, W, 3)的彩色圖。這是NumPy的高級索引特性完全不需要寫循環(huán)。5.2 大圖的滑窗預(yù)測如果你要對遙感影像或全景圖這種超大圖像做分割直接整圖喂給模型很容易顯存不足因?yàn)槟P陀泄潭ㄏ虏蓸颖稊?shù)和最大分辨率限制?;邦A(yù)測是把大圖切成若干有重疊的patch分別送入模型預(yù)測最后把所有patch拼回原圖尺寸。重疊區(qū)域通常取patch尺寸的1/8到1/4拼接時對重疊區(qū)域取平均值或直接取中心部分能有效避免拼接邊緣出現(xiàn)接縫。實(shí)際使用中我是先把大圖按512x512切成patch不足512的邊緣做padding預(yù)測完裁剪掉padding部分再拼回原圖。整個過程計(jì)算量不小但能在6GB顯存下處理幾萬乘幾萬像素的影像。6. 實(shí)際排錯記錄我踩過的幾個“經(jīng)典坑”既然你已經(jīng)準(zhǔn)備用自己的數(shù)據(jù)跑了有些坑我提前幫你踩掉。這些問題是U-Net語義分割項(xiàng)目中出現(xiàn)頻率最高的而且?guī)缀醵疾皇谴a邏輯本身的問題而是數(shù)據(jù)或環(huán)境層面的細(xì)節(jié)。6.1 mask全黑導(dǎo)致的loss不下降我第一次跑二分類時訓(xùn)練loss一直穩(wěn)定在0.69左右不下降。后來排查發(fā)現(xiàn)labelme導(dǎo)出的灰度png在保存時被誤存成了三通道的RGB png用Image.open(...).convert(L)讀取后所有像素值都變成了0或255而標(biāo)簽要求類別索引是0到1。結(jié)果模型看到label中大量為255的像素自然沒法收斂。解決方案是保存mask時強(qiáng)制使用單通道模式mask_pil Image.fromarray(mask_array.astype(np.uint8), modeL) mask_pil.save(xxx.png)dataset里讀mask時也記得加convert(L)。如果mask是標(biāo)注工具自動生成的先打印一下像素值的np.unique()結(jié)果確認(rèn)類別索引符合預(yù)期再開始訓(xùn)練。6.2 類別索引從0還是從1開始語義分割的標(biāo)簽索引和分類任務(wù)一樣從0開始。背景類通常是0其他類別依次從1開始。如果你的數(shù)據(jù)標(biāo)注工具是從1開始的訓(xùn)練前需要統(tǒng)一減1否則CrossEntropyLoss會報Target ... is out of bounds的錯誤。6.3 訓(xùn)練集和驗(yàn)證集的數(shù)據(jù)泄露很多人在劃分?jǐn)?shù)據(jù)集時不注意數(shù)據(jù)的獨(dú)立性。如果同一張?jiān)瓐D經(jīng)過不同裁剪后同時出現(xiàn)在訓(xùn)練集和驗(yàn)證集驗(yàn)證mIoU會虛高換到新數(shù)據(jù)上效果明顯變差。劃分?jǐn)?shù)據(jù)時要以“原始圖像”為單位劃分而不是以裁剪后的patch為單位。6.4 輸入歸一化方式前后不一致訓(xùn)練時把圖像像素除以255映射到0到1區(qū)間推理時如果忘記歸一化直接喂原圖模型輸出基本就是噪聲。這個看似簡單的問題實(shí)際排查起來很耗時間。建議寫一個統(tǒng)一的預(yù)處理函數(shù)訓(xùn)練和推理共用一個函數(shù)從源頭避免這類問題。7. 我自己的一些補(bǔ)充建議如果你準(zhǔn)備長期做語義分割方向有幾個工具和習(xí)慣值得提前養(yǎng)成。我在多個數(shù)據(jù)集上換著跑過之后越來越覺得工作流順暢比單次訓(xùn)練出好結(jié)果更重要。建議把數(shù)據(jù)檢查和指標(biāo)統(tǒng)計(jì)腳本單獨(dú)做成工具函數(shù)每次訓(xùn)練前先打印數(shù)據(jù)集類別分布和mask像素值分布訓(xùn)練后把驗(yàn)證集預(yù)測結(jié)果可視化成網(wǎng)格圖邊訓(xùn)練邊看效果而不是等全部訓(xùn)完再調(diào)。這些“基礎(chǔ)設(shè)施”看似花費(fèi)時間但長期來看能省下大量無效試驗(yàn)時間。對于U-Net的改進(jìn)方向當(dāng)你基礎(chǔ)版跑通后可以考慮把backbone替換成ResNet34或ResNet50預(yù)訓(xùn)練權(quán)重編碼器部分直接加載ImageNet預(yù)訓(xùn)練模型通常能帶來mIoU的明顯提升。不過這里有個小提醒替換backbone后輸入圖像的歸一化方式也要按ImageNet的mean和std來做否則預(yù)訓(xùn)練權(quán)重反而成了負(fù)擔(dān)。我自己在實(shí)際項(xiàng)目中還會嘗試把U-Net的普通卷積換成深度可分離卷積參數(shù)量顯著下降推理速度變快精度損失很小。這類結(jié)構(gòu)改進(jìn)在保持原有數(shù)據(jù)管線不變的前提下就能做適合在基礎(chǔ)版穩(wěn)定之后逐步迭代。本文還有配套的精品資源點(diǎn)擊獲取