:從數(shù)據(jù)準(zhǔn)備到模型訓(xùn)練全流程)
簡介這是一份基于PyTorch實現(xiàn)Unet多類別語義分割的實戰(zhàn)資源面向具備一定深度學(xué)習(xí)基礎(chǔ)、希望掌握圖像分割模型訓(xùn)練與調(diào)優(yōu)的開發(fā)者。資源共46個文件以Python源碼為核心包含19個py腳本與24個pyc編譯文件另有2個txt配置說明和1個json參數(shù)文件整體僅69KB結(jié)構(gòu)緊湊便于直接閱讀與修改。已有15245人學(xué)習(xí)這份資源屬于同類教程中熱度較高的實踐型內(nèi)容。作者圍繞Encoder-Decoder結(jié)構(gòu)、跳躍連接、多類別輸出通道設(shè)計等關(guān)鍵點組織代碼配套提供數(shù)據(jù)加載、自定義變換、訓(xùn)練評估、損失函數(shù)、學(xué)習(xí)率調(diào)度、指標(biāo)計算等模塊并附有數(shù)據(jù)集劃分與類別權(quán)重計算腳本基本覆蓋了從數(shù)據(jù)預(yù)處理到模型推理的完整流程。通過閱讀源碼與對應(yīng)博客讀者可掌握Unet在自定義多類別數(shù)據(jù)集上的遷移方法、訓(xùn)練流程設(shè)計及調(diào)優(yōu)技巧適合用于醫(yī)學(xué)影像、遙感圖像等場景的入門與項目參考。1. 項目概述為什么選擇Pytorch配合Unet做多類別分割先說結(jié)論如果你手里有一批自己的圖片數(shù)據(jù)想按像素把圖中不同物體分出來比如道路、建筑、植被、水體這類地物目標(biāo)那“Pytorch Unet 多類別數(shù)據(jù)集”這套組合是目前開源社區(qū)里最省心、最不容易把自己繞暈的路線。這個標(biāo)題之所以有這么多人在搜是因為它幾乎覆蓋了從零入門語義分割的所有關(guān)鍵環(huán)節(jié)框架選型、模型結(jié)構(gòu)、數(shù)據(jù)組織、訓(xùn)練調(diào)參、結(jié)果評估。我自己最早接觸語義分割時也糾結(jié)過用TensorFlow還是Pytorch后來徹底轉(zhuǎn)向Pytorch原因很樸素調(diào)試方便報錯信息看得懂?dāng)帱c能直接打在張量運算那一行上。配合Unet這種編碼器-解碼器結(jié)構(gòu)哪怕數(shù)據(jù)集只有幾百張圖也能訓(xùn)出一個效果不錯的多類別分割模型。這篇文章不會跟你聊太虛的理論而是把我實際跑通整個流程的步驟、參數(shù)、踩坑記錄都擺出來你照著操作就能在自己的多類別數(shù)據(jù)集上復(fù)現(xiàn)。文章適合三類人看剛?cè)腴T語義分割的學(xué)生、需要用自己數(shù)據(jù)做分割實驗的工程師、以及想快速驗證Unet效果的產(chǎn)品人員。2. 整體設(shè)計思路與方案選型2.1 為什么Unet依然是多類別分割的首選基線Unet之所以經(jīng)典核心在于它同時保住了“細(xì)節(jié)”和“語義”。下采樣路徑不斷縮小特征圖尺寸讓模型能看到更大的感受野上采樣路徑則把高層的語義信息逐步還原到原圖分辨率。中間那一圈跳連接把下采樣時各層的位置細(xì)節(jié)直接拼到上采樣路徑上相當(dāng)于給模型開了一條“記憶通道”小目標(biāo)邊緣不容易丟。對于多類別分割任務(wù)比如五類甚至十類地物Unet每一條跳連接都在幫助模型區(qū)分“這里是邊界還是內(nèi)部”。我自己做過對比實驗在同樣的數(shù)據(jù)集上把Unet換成PSPNet發(fā)現(xiàn)小目標(biāo)類別的交并比下降明顯。原因不復(fù)雜PSPNet把重點放在全局池化上對小目標(biāo)的敏感度反而不如Unet這種逐層傳遞細(xì)節(jié)的結(jié)構(gòu)。所以如果你的數(shù)據(jù)里存在較多小目標(biāo)或細(xì)長條目標(biāo)Unet是最穩(wěn)的起點。2.2 Pytorch生態(tài)里的三個關(guān)鍵選擇第一框架版本跟進。建議用Pytorch 2.x系列如果你需要GPU訓(xùn)練記得提前確認(rèn)CUDA、cuDNN和顯卡驅(qū)動的匹配關(guān)系。以2024年之后的環(huán)境為例Pytorch 2.0以上版本對自動混合精度訓(xùn)練的支持更完善顯存占用更友好。第二模型實現(xiàn)方式。可以直接從網(wǎng)上找Unet的Pytorch實現(xiàn)也可以自己按論文結(jié)構(gòu)手寫。實話說手寫一遍Unet比復(fù)制十遍別人的代碼都有用結(jié)構(gòu)細(xì)節(jié)會刻在你腦子里。第三預(yù)訓(xùn)練編碼器。如果你用torchvision里的ResNet作為Unet的骨干加載ImageNet預(yù)訓(xùn)練權(quán)重訓(xùn)練收斂速度會快不少尤其是當(dāng)你的數(shù)據(jù)集規(guī)模不大的時候。選型邏輯很直接小數(shù)據(jù)集靠預(yù)訓(xùn)練權(quán)重大數(shù)據(jù)集靠模型容量。如果你的數(shù)據(jù)只有兩三百張建議選擇ResNet34做編碼器如果數(shù)據(jù)到了一千張以上可以嘗試ResNet50或者直接換EfficientNet。這個不是鐵律而是我在不同規(guī)模數(shù)據(jù)上反復(fù)試出來的經(jīng)驗。3. 數(shù)據(jù)準(zhǔn)備多類別數(shù)據(jù)集的整理與預(yù)處理3.1 目錄結(jié)構(gòu)與標(biāo)注格式的統(tǒng)一很多人在模型跑不起來時才發(fā)現(xiàn)問題出在數(shù)據(jù)上而不是代碼上。多類別語義分割的數(shù)據(jù)集標(biāo)準(zhǔn)的組織方式是這樣dataset/ ├── images/ │ ├── img_001.jpg │ ├── img_002.jpg │ └── ... └── masks/ ├── img_001.png ├── img_002.png └── ...圖片格式一般用jpg或png都行但掩膜mask必須是png而且是單通道的灰度圖或調(diào)色板模式。為什么必須png因為jpg是有損壓縮會導(dǎo)致標(biāo)注類別邊緣出現(xiàn)偽色模型會學(xué)到錯誤信息。特別注意掩膜每個像素點的數(shù)值背景為0第一個類別為1第二個類別為2以此類推。如果你用Labelme這類工具標(biāo)注導(dǎo)出的掩膜是調(diào)色板模式需要用代碼轉(zhuǎn)成類別索引不然訓(xùn)練時一算損失就是一片NaN。3.2 數(shù)據(jù)增強與樣本均衡技巧多類別分割里最頭疼的問題就是類別不平衡。比如一棟建筑物可能只占畫面面積的5%背景卻占了70%。如果直接訓(xùn)練模型會傾向把所有像素都預(yù)測成背景。我的做法是對每個類別統(tǒng)計像素占比然后給損失函數(shù)里的每個類別分配權(quán)重權(quán)重和像素占比成反比。數(shù)據(jù)增強我用的是albumentations庫比torchvision的transform靈活得多。我常用的一套增強組合包括水平翻轉(zhuǎn)、垂直翻轉(zhuǎn)、隨機旋轉(zhuǎn)90度、隨機裁剪和亮度對比度調(diào)整。這里有個細(xì)節(jié)對圖像做翻轉(zhuǎn)和旋轉(zhuǎn)時掩膜必須做同樣的變換albumentations保證了這一點。增強不是越多越好如果你的類別是建筑物這類剛性目標(biāo)翻轉(zhuǎn)和旋轉(zhuǎn)沒問題如果你處理的是文本行這類有方向性的目標(biāo)旋轉(zhuǎn)90度會把標(biāo)注搞亂。3.3 自定義Dataset類的關(guān)鍵代碼寫Dataset類時最容易犯的一個錯誤是忘記把掩膜里的類別索引壓到從0開始連續(xù)分布。以下是我常用的代碼import torch from torch.utils.data import Dataset from PIL import Image import os import numpy as np class SegmentationDataset(Dataset): def __init__(self, image_dir, mask_dir, transformNone, class_mappingNone): self.image_dir image_dir self.mask_dir mask_dir self.transform transform self.class_mapping class_mapping or {} self.images sorted([f for f in os.listdir(image_dir) if f.endswith((.jpg, .png))]) def __len__(self): return len(self.images) def __getitem__(self, idx): img_name self.images[idx] img_path os.path.join(self.image_dir, img_name) mask_path os.path.join(self.mask_dir, img_name.replace(.jpg, .png)) image np.array(Image.open(img_path).convert(RGB)) mask np.array(Image.open(mask_path)) # 不要convert(RGB)保持灰度 # 如果mask是調(diào)色板模式0-255任意值需要做類別映射 if self.class_mapping: mapped_mask np.zeros_like(mask) for old_id, new_id in self.class_mapping.items(): mapped_mask[mask old_id] new_id mask mapped_mask if self.transform: transformed self.transform(imageimage, maskmask) image transformed[image] mask transformed[mask] image_tensor torch.from_numpy(image).permute(2, 0, 1).float() / 255.0 mask_tensor torch.from_numpy(mask).long() return image_tensor, mask_tensor這段代碼里有個非常重要的點掩膜轉(zhuǎn)成tensor用.long()不要用.float()。因為后面損失函數(shù)CrossEntropyLoss期望輸入是整數(shù)類別索引如果用float會報錯或者產(chǎn)生錯誤結(jié)果。我在一開始就踩過這個坑整整調(diào)了一個晚上才發(fā)現(xiàn)是類型不匹配。4. 模型實現(xiàn)Unet結(jié)構(gòu)與多類別輸出適配4.1 Unet核心結(jié)構(gòu)拆解我在這里不放完整的三百行Unet代碼了因為網(wǎng)上開源實現(xiàn)非常多我建議你找一個star數(shù)高的倉庫讀一遍結(jié)構(gòu)但有幾個核心參數(shù)必須搞清楚。Unet整體分為編碼器、瓶頸、解碼器三段。編碼器是若干個卷積塊加下采樣特征圖尺寸減半通道數(shù)翻倍瓶頸在最底層解碼器逐步上采樣通道數(shù)減半并與對應(yīng)的編碼器特征圖拼接。關(guān)鍵在于最后一層卷積的輸出通道數(shù)必須等于你的類別數(shù)。舉例如果是五類分割最后一層卷積輸出通道數(shù)就設(shè)為5。每個通道對應(yīng)一個類別的置信度分?jǐn)?shù)。訓(xùn)練時用CrossEntropyLoss它對每個像素在通道維度上做softmax后計算損失。這部分不需要你自己寫softmaxPytorch的CrossEntropyLoss內(nèi)部已經(jīng)包含了。4.2 多類別輸出的通道設(shè)置與損失函數(shù)選擇我遇到不少人在修改Unet時只改了模型最后一層的輸出通道數(shù)但忽略了編碼器預(yù)訓(xùn)練權(quán)重的加載方式。如果是自己寫的Unet從頭訓(xùn)練沒問題如果你的編碼器要加載預(yù)訓(xùn)練權(quán)重前幾層的通道數(shù)必須和ImageNet預(yù)訓(xùn)練模型一致通常就是RGB三通道輸入輸出通道按骨干網(wǎng)絡(luò)設(shè)定。損失函數(shù)方面多類別分割最常用的是CrossEntropyLoss加DiceLoss的組合。我實際測試下來純用CrossEntropyLoss小目標(biāo)類別的分割效果一般純用DiceLoss訓(xùn)練初期損失波動劇烈。兩者的加權(quán)和比較穩(wěn)具體公式是loss 0.7 * ce_loss 0.3 * dice_loss這個比例可以根據(jù)你的數(shù)據(jù)調(diào)整。如果類別特別不均衡把dice_loss的權(quán)重調(diào)高到0.5甚至0.7。DiceLoss對不平衡不敏感它能直接優(yōu)化類別區(qū)域的重合度。5. 訓(xùn)練配置與完整實操流程5.1 環(huán)境準(zhǔn)備與關(guān)鍵參數(shù)設(shè)置環(huán)境方面推薦用Anaconda創(chuàng)建獨立的環(huán)境。以Ubuntu系統(tǒng)為例常見組合是Python 3.10.11加Pytorch 2.8.0加CUDA 12.1這一套搭配在GTX 30系和40系顯卡上表現(xiàn)穩(wěn)定。Windows下的流程類似只是CUDA環(huán)境變量配置要格外小心。如果顯卡顯存只有6G建議輸入圖片尺寸用256x256批次大小設(shè)為4到8顯存12G以上的話輸入尺寸可以提高到512x512批次大小設(shè)8到16。我整理了訓(xùn)練階段幾個關(guān)鍵的超參數(shù)參考值參數(shù)推薦值備注輸入尺寸256x256 / 512x512小顯存用256大顯存用512批次大小4 / 8 / 16視顯存而定初始學(xué)習(xí)率1e-4 / 1e-3使用AdamW優(yōu)化器學(xué)習(xí)率調(diào)整CosineAnnealingLR避免后期震蕩訓(xùn)練輪數(shù)50 / 100看驗證集指標(biāo)早停優(yōu)化器AdamW比Adam穩(wěn)定性好5.2 訓(xùn)練循環(huán)中的關(guān)鍵代碼訓(xùn)練循環(huán)本身并不復(fù)雜但有一些細(xì)節(jié)會讓訓(xùn)練過程順暢很多。比如啟用自動混合精度用autocast和GradScaler顯存能省下近一半訓(xùn)練速度還能提升。另外每個epoch保存一次checkpoint別只保存在最后一個epoch——訓(xùn)練中斷是家常便飯有checkpoint才能續(xù)上繼續(xù)訓(xùn)。一個簡化但完整可跑的訓(xùn)練循環(huán)結(jié)構(gòu)import torch from torch.cuda.amp import autocast, GradScaler scaler GradScaler() optimizer torch.optim.AdamW(model.parameters(), lr1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_maxepochs) for epoch in range(epochs): model.train() total_loss 0 for images, masks in train_loader: images, masks images.to(device), masks.to(device) optimizer.zero_grad() with autocast(): outputs model(images) loss ce_criterion(outputs, masks) 0.3 * dice_loss(outputs, masks) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() total_loss loss.item() scheduler.step() avg_loss total_loss / len(train_loader) print(fEpoch {epoch}/{epochs}, Loss: {avg_loss:.4f}) if epoch % 5 0: torch.save(model.state_dict(), fcheckpoints/unet_epoch_{epoch}.pth)注意dice_loss需要自己實現(xiàn)def dice_loss(pred, target, smooth1.0): pred torch.softmax(pred, dim1) target_onehot torch.nn.functional.one_hot(target, num_classesnum_classes).permute(0, 3, 1, 2).float() intersection (pred * target_onehot).sum(dim(2, 3)) union pred.sum(dim(2, 3)) target_onehot.sum(dim(2, 3)) dice (2 * intersection smooth) / (union smooth) return 1 - dice.mean()5.3 評估指標(biāo)與結(jié)果可視化訓(xùn)練完成后不要只看loss要算每個類別的IoU和mIoU。mIoU是所有類別IoU的平均值這是語義分割最通用的評價指標(biāo)。我需要提醒你計算IoU時要對每個類別單獨統(tǒng)計不能直接拿混淆矩陣整體算。具體實現(xiàn)可以用sklearn的confusion_matrix輔助。from sklearn.metrics import confusion_matrix def compute_iou(pred_mask, true_mask, num_classes): pred_flat pred_mask.flatten() true_flat true_mask.flatten() cm confusion_matrix(true_flat, pred_flat, labelslist(range(num_classes))) intersection np.diag(cm) union cm.sum(axis0) cm.sum(axis1) - np.diag(cm) iou intersection / (union 1e-6) return iou可視化方面推薦把預(yù)測結(jié)果疊加到原圖上用半透明色塊顯示不同類別。我覺得看分割效果比死磕指標(biāo)更直觀很多邊界問題光看IoU是發(fā)現(xiàn)不了的。別只看訓(xùn)練集的預(yù)測效果一定要去驗證集上抽幾張圖看邊界質(zhì)量——是不是出現(xiàn)鋸齒狀、是不是有空洞。6. 常見問題與排查技巧實錄6.1 損失降至Nan的排查過程這是最多人碰到的坑。我遇到過一次損失降到NaN排查步驟是先檢查掩膜里是否有超出類別數(shù)的值。比如類別數(shù)是5掩膜數(shù)值范圍卻是0到255這會讓CrossEntropyLoss計算出NaN。然后檢查數(shù)據(jù)歸一化圖片輸入是否除以了255掩膜是否保持整數(shù)類型。最后檢查學(xué)習(xí)率如果初始學(xué)習(xí)率設(shè)到0.1梯度爆炸也會導(dǎo)致NaN。我習(xí)慣把初始學(xué)習(xí)率控制在1e-4到1e-3配合AdamW很少再碰到NaN。6.2 顯存不足的應(yīng)對策略顯存不足是個很現(xiàn)實的問題。我一開始用512x512輸入、批次16直接爆顯存。后來做了三件事?lián)Q成分批訓(xùn)練、縮小輸入尺寸、啟用混合精度。假設(shè)你的顯卡是8G顯存推薦直接用256x256輸入加批次8加混合精度這樣訓(xùn)練速度反而可能比大尺寸低批次更穩(wěn)定。另外在forward里加一句torch.cuda.empty_cache()也能清理一部分碎片顯存但不要在每個step都調(diào)用會拖慢速度。6.3 模型訓(xùn)練不收斂或過擬合的調(diào)整訓(xùn)練不收斂先看損失曲線是震蕩還是不下降。震蕩說明學(xué)習(xí)率太高調(diào)低一個數(shù)量級不下降說明可能模型結(jié)構(gòu)或數(shù)據(jù)出了問題。我遇到過一次模型輸出恒為背景的情況檢查發(fā)現(xiàn)掩膜數(shù)值沒有對齊所有類別都被映射成了0。過擬合則表現(xiàn)為訓(xùn)練集損失低但驗證集IoU不增這時加大數(shù)據(jù)增強強度、加Dropout、縮小模型容量都有效。我把這些高頻問題整理成了一個速查表癥狀可能原因解決方案損失NaN掩膜值超出類別范圍檢查mask像素值確保0到num_classes-1顯存不足輸入尺寸過大/批次過大降輸入尺寸、分批、開混合精度驗證IoU不漲過擬合或?qū)W習(xí)率過大加增強、降學(xué)習(xí)率、早停預(yù)測全是背景類別索引不對齊檢查數(shù)據(jù)集類映射邏輯邊界粗糙跳連接被忽略檢查模型是否真的用到了skip connection6.4 模型保存、加載與推理的完整流程訓(xùn)練完模型后保存方式我推薦只存state_dict不存整個模型因為后者在Pytorch版本升級后容易反序列化失敗。加載模型后做推理要注意輸入圖片必須做和訓(xùn)練時一樣的預(yù)處理resize到相同尺寸、歸一化到0到1、轉(zhuǎn)成張量、加batch維度。預(yù)測輸出是一個形狀為(1, num_classes, H, W)的張量用argmax(dim1)取每個像素的類別索引。最后轉(zhuǎn)換成彩色圖時準(zhǔn)備一個調(diào)色板數(shù)組把類別索引映射成RGB顏色。import torch import numpy as np from PIL import Image import torchvision.transforms as transforms def inference_single_image(model, image_path, device, num_classes): transform transforms.Compose([ transforms.Resize((256, 256)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) image Image.open(image_path).convert(RGB) input_tensor transform(image).unsqueeze(0).to(device) model.eval() with torch.no_grad(): output model(input_tensor) pred torch.argmax(output, dim1).squeeze(0).cpu().numpy() return pred推理后如果想保存成彩色分割圖可以用這樣一個簡單的映射color_map np.array([ [0, 0, 0], # 背景 [255, 0, 0], # 類別1紅色 [0, 255, 0], # 類別2綠色 [0, 0, 255], # 類別3藍色 [255, 255, 0], # 類別4黃色 ], dtypenp.uint8) segmentation_rgb color_map[pred] Image.fromarray(segmentation_rgb).save(prediction.png)7. 多類別數(shù)據(jù)標(biāo)注與類別映射的實戰(zhàn)心得7.1 標(biāo)注工具選擇與格式轉(zhuǎn)換常見坑多類別分割最關(guān)鍵的第一步其實是標(biāo)注質(zhì)量。我試過Labelme、EISeg、CVAT最終常用Labelme配合腳本轉(zhuǎn)成Unet需要的png掩膜。Labelme導(dǎo)出的是json文件每個多邊形對應(yīng)一個label需要逐張解析json并把多邊形填充成掩膜。這里有個非常隱蔽的坑json里標(biāo)注的label名稱和你最終想要的類別索引可能不是一回事一定要建立一個字典做映射。import json import numpy as np import cv2 import os def json_to_mask(json_path, height, width, label_map): with open(json_path, r, encodingutf-8) as f: data json.load(f) mask np.zeros((height, width), dtypenp.uint8) for shape in data[shapes]: label shape[label] points np.array(shape[points], dtypenp.int32) if label in label_map: cv2.fillPoly(mask, [points], label_map[label]) return mask這段函數(shù)把json里的每個多邊形填充到掩膜上label_map里存的是比如{建筑: 1, 道路: 2, 植被: 3}這樣的鍵值對。實際操作中我發(fā)現(xiàn)最花時間的不是寫轉(zhuǎn)換腳本而是清洗標(biāo)注數(shù)據(jù)。比如相鄰圖片邊緣處多邊形沒有貼合圖像邊界導(dǎo)致交界處出現(xiàn)一條無類別帶狀區(qū)域這一條區(qū)域會成為模型預(yù)測錯誤的高發(fā)區(qū)。處理方法是對掩膜做一個形態(tài)學(xué)閉運算把細(xì)小的空洞縫補上。7.2 數(shù)據(jù)集劃分比例與驗證集選擇多類別分割的數(shù)據(jù)劃分和普通分類不太一樣。除了按文件數(shù)量比例劃分還要考慮類別分布。我遇到過一種情況訓(xùn)練集里“水體”樣本很多驗證集里“水體”只出現(xiàn)在一張圖的一角結(jié)果導(dǎo)致驗證集水體IoU特別低模型并沒有過擬合純粹是驗證集抽樣偏差。比較穩(wěn)妥的做法是按圖像整體劃分保證每一類在訓(xùn)練集和驗證集都出現(xiàn)如果某個類別的圖像特別少可以考慮將數(shù)據(jù)增強用到驗證集上但這并不常規(guī)更合理的是用k-fold交叉驗證。對數(shù)據(jù)量只有幾百張的情況k-fold交叉驗證是評估模型真實水平的有效方式。最后分享一個細(xì)節(jié)在訓(xùn)練過程中如果發(fā)現(xiàn)某幾個類別的IoU一直偏低先別急著改模型結(jié)構(gòu)回到標(biāo)注數(shù)據(jù)里看看這些類別的標(biāo)注質(zhì)量是否有重疊標(biāo)注、漏標(biāo)、邊界粗糙的情況。我之前處理遙感影像分割時“陰影”類別的IoU怎么都上不去后來仔細(xì)拉大圖片比對發(fā)現(xiàn)標(biāo)注員把很多本來就模糊的陰影邊界標(biāo)歪了模型學(xué)到的邊界自然混亂。重新清洗了一批標(biāo)注后這個類別的IoU直接提升了十幾個點。數(shù)據(jù)質(zhì)量決定了模型效果的上限這話在語義分割里體現(xiàn)得特別明顯。本文還有配套的精品資源點擊獲取