習(xí)的轎車(chē)背景去除:U-Net語(yǔ)義分割實(shí)戰(zhàn))
簡(jiǎn)介基于深度學(xué)習(xí)的轎車(chē)背景去除算法課程設(shè)計(jì)資源包面向計(jì)算機(jī)、數(shù)學(xué)、電子信息類(lèi)專(zhuān)業(yè)學(xué)生尤其適合作為課程設(shè)計(jì)、期末大作業(yè)或畢業(yè)設(shè)計(jì)項(xiàng)目。資源以Python源碼為核心13個(gè)py腳本覆蓋數(shù)據(jù)加載、模型定義、損失函數(shù)、訓(xùn)練與配置等完整環(huán)節(jié)2個(gè)h5預(yù)訓(xùn)練權(quán)重文件支持直接加載模型進(jìn)行推理2份docx說(shuō)明文檔與1份md說(shuō)明詳細(xì)闡述算法原理、軟件體系結(jié)構(gòu)與設(shè)計(jì)模式的應(yīng)用1個(gè)pptx可用于答辯匯報(bào)。整個(gè)壓縮包共19個(gè)文件大小約37.26MB超過(guò)六成為Python腳本并附帶開(kāi)發(fā)工作日志目錄結(jié)構(gòu)清晰、模塊拆分規(guī)范。目前已有八十六人瀏覽學(xué)習(xí)。通過(guò)這份資料讀者可以快速掌握基于深度學(xué)習(xí)的圖像分割任務(wù)流程理解軟件架構(gòu)與設(shè)計(jì)模式在實(shí)際項(xiàng)目中的落地方式并基于完整源碼進(jìn)行二次開(kāi)發(fā)或功能擴(kuò)展。1. 從課程設(shè)計(jì)到落地轎車(chē)背景去除到底在解決什么問(wèn)題期末前兩周才確定題目既要交軟件體系結(jié)構(gòu)與設(shè)計(jì)模式的課程設(shè)計(jì)又想體現(xiàn)深度學(xué)習(xí)算法能力多數(shù)人最后都卡在“模型跑通了但說(shuō)不清工程結(jié)構(gòu)”這步。這個(gè)基于深度學(xué)習(xí)的轎車(chē)背景去除項(xiàng)目正是用于解決這類(lèi)問(wèn)題的完整樣例輸入一張任意場(chǎng)景下的轎車(chē)照片輸出只有車(chē)身保留、背景被置為純色的掩碼圖本質(zhì)是逐像素的語(yǔ)義分割任務(wù)。相比人臉摳圖、通用物體分割車(chē)輛目標(biāo)輪廓清晰但包含車(chē)窗反光、地面陰影、車(chē)漆高光等干擾很適合作為入門(mén)級(jí)深度圖像分割實(shí)戰(zhàn)。代碼庫(kù)劃分為數(shù)據(jù)集加載、模型定義、損失函數(shù)、訓(xùn)練配置、推理應(yīng)用五個(gè)模塊直接映射軟件體系結(jié)構(gòu)課程里分層與解耦的考核點(diǎn)。適合正在做圖像分割入門(mén)、準(zhǔn)備課程設(shè)計(jì)答辯或期末大作業(yè)、希望把設(shè)計(jì)模式落到代碼里的學(xué)生與開(kāi)發(fā)者。2. 任務(wù)定義與數(shù)據(jù)準(zhǔn)備mask 標(biāo)注與數(shù)據(jù)增強(qiáng)管線(xiàn)的搭建2.1 背景去除為什么是逐像素語(yǔ)義分割背景去除和常見(jiàn)的物體檢測(cè)有本質(zhì)區(qū)別。物體檢測(cè)輸出的是邊界框而背景去除要對(duì)圖像中的每一個(gè)像素做二分類(lèi)判斷屬于轎車(chē)還是屬于背景。這個(gè)任務(wù)在計(jì)算機(jī)視覺(jué)里被稱(chēng)為語(yǔ)義分割它比分類(lèi)任務(wù)多保留了空間位置信息也比檢測(cè)任務(wù)更精細(xì)。轎車(chē)背景去除的難點(diǎn)在于三個(gè)區(qū)域車(chē)窗玻璃會(huì)反射周?chē)h(huán)境、車(chē)漆顏色與背景接近時(shí)會(huì)混淆、車(chē)輪與地面陰影的邊界難以劃分。課程設(shè)計(jì)的考核重點(diǎn)通常不只是“效果好不好”還包括“為什么這么設(shè)計(jì)”。語(yǔ)義分割采用編碼器-解碼器結(jié)構(gòu)編碼器逐層下采樣提取高層語(yǔ)義特征解碼器逐層上采樣恢復(fù)空間分辨率。車(chē)輛輪廓的精細(xì)程度取決于解碼器對(duì)邊緣信息的恢復(fù)能力這也是后續(xù)第 3 章選擇 U-Net 作為骨干網(wǎng)絡(luò)的原因。理解這一點(diǎn)才能在軟件設(shè)計(jì)說(shuō)明文檔里交代清楚模型選型的依據(jù)而不是簡(jiǎn)單寫(xiě)一句“使用了深度學(xué)習(xí)”。2.2 數(shù)據(jù)目錄結(jié)構(gòu)與 dataset.py 的實(shí)現(xiàn)項(xiàng)目中的數(shù)據(jù)由原始車(chē)輛圖像和對(duì)應(yīng)的 mask 標(biāo)注組成。標(biāo)準(zhǔn)目錄結(jié)構(gòu)如下第一部分是課程設(shè)計(jì)里交付時(shí)要講清楚的內(nèi)容dataset/ ├── train/ │ ├── input/ # 原始轎車(chē)圖像 │ │ ├── 0001.jpg │ │ └── ... │ └── mask/ # 二值掩碼圖白色為轎車(chē)黑色為背景 │ ├── 0001.png │ └── ... └── val/ ├── input/ └── mask/注意區(qū)分這里的 train/input 與模型訓(xùn)練環(huán)節(jié)的 train 數(shù)據(jù)集前者是磁盤(pán)上的數(shù)據(jù)組織后者是訓(xùn)練循環(huán)中的批次數(shù)據(jù)。加載數(shù)據(jù)時(shí)使用 torchvision 的 transforms 做尺寸統(tǒng)一和增強(qiáng)代碼實(shí)現(xiàn)如下這也是 dataset.py 的核心內(nèi)容class CarSegDataset(Dataset): def __init__(self, img_dir, mask_dir, img_size(512, 512), augFalse): self.img_paths sorted(glob.glob(os.path.join(img_dir, *.jpg))) self.mask_paths sorted(glob.glob(os.path.join(mask_dir, *.png))) self.img_size img_size self.aug aug # 兩個(gè)目錄下的文件應(yīng)一一對(duì)應(yīng)常見(jiàn)錯(cuò)誤是按文件名排序不一致導(dǎo)致圖文錯(cuò)位 assert len(self.img_paths) len(self.mask_paths), image count ! mask count def __len__(self): return len(self.img_paths) def __getitem__(self, idx): img Image.open(self.img_paths[idx]).convert(RGB) mask Image.open(self.mask_paths[idx]).convert(L) img img.resize(self.img_size, Image.BILINEAR) mask mask.resize(self.img_size, Image.NEAREST) # 掩碼不能做插值平滑 if self.aug: if random.random() 0.5: img img.transpose(Image.FLIP_LEFT_RIGHT) mask mask.transpose(Image.FLIP_LEFT_RIGHT) if random.random() 0.5: img img.transpose(Image.FLIP_TOP_BOTTOM) mask mask.transpose(Image.FLIP_TOP_BOTTOM) # 顏色抖動(dòng)只作用于原圖不作用于 mask img transforms.ColorJitter(brightness0.2, contrast0.2)(img) img_tensor transforms.ToTensor()(img) mask_tensor torch.as_tensor(np.array(mask), dtypetorch.float32) / 255.0 mask_tensor mask_tensor.unsqueeze(0) # 增加通道維變成 [1, H, W] return img_tensor, mask_tensor代碼邏輯上有兩個(gè)容易踩坑的參數(shù)要重點(diǎn)說(shuō)明。第一mask 縮放必須用Image.NEAREST最近鄰插值不能用BILINEAR雙線(xiàn)性插值因?yàn)?mask 是離散的二值標(biāo)簽雙線(xiàn)性插值會(huì)在邊緣產(chǎn)生 0.3、0.7 這類(lèi)中間灰度值直接污染損失函數(shù)的計(jì)算。第二mask_tensor / 255.0是為了把像素值從 0255 歸一化到 01與模型輸出的 sigmoid 概率值對(duì)齊。unsqueeze(0)是為配合 PyTorch 的通道維度約定語(yǔ)義分割的數(shù)據(jù)格式為[batch, channel, height, width]單通道 mask 需要補(bǔ)上 channel 維度。2.3 數(shù)據(jù)增強(qiáng)參數(shù)怎么定數(shù)據(jù)增強(qiáng)解決的是模型泛化問(wèn)題不是越多越好每一項(xiàng)增強(qiáng)都有代價(jià)。整理出下面的參數(shù)對(duì)照表課程設(shè)計(jì)文檔里直接描述為“訓(xùn)練階段采用輕量數(shù)據(jù)增強(qiáng)”即可注意本表不涉及訓(xùn)練超參數(shù)的內(nèi)容僅描述數(shù)據(jù)增強(qiáng)環(huán)節(jié)增強(qiáng)方式推薦參數(shù)作用代價(jià)與陷阱水平翻轉(zhuǎn)p0.5消除左右視角偏差樣本量翻倍車(chē)牌文字鏡像但不影響分割任務(wù)垂直翻轉(zhuǎn)p0.5增加多樣化天空與地面語(yǔ)義被顛倒慎用于有方向性數(shù)據(jù)集隨機(jī)裁剪0.8 比例范圍模擬局部遮擋增強(qiáng)目標(biāo)局部特征可能裁掉整個(gè)車(chē)需配合重采樣色彩抖動(dòng)brightness0.2, contrast0.2增強(qiáng)對(duì)光照變化的魯棒性只作用于原圖絕不作用于 mask歸一化mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]統(tǒng)一像素分布加速收斂必須與預(yù)訓(xùn)練權(quán)重配套不可隨意換一個(gè)常見(jiàn)誤區(qū)是認(rèn)為增強(qiáng)越強(qiáng)效果越好。對(duì)于轎車(chē)背景去除這種目標(biāo)相對(duì)居中的任務(wù)過(guò)度的隨機(jī)裁剪會(huì)導(dǎo)致訓(xùn)練樣本中經(jīng)常丟失整車(chē)結(jié)構(gòu)模型反而學(xué)不到完整的車(chē)身形態(tài)。實(shí)踐中的做法更傾向于保留水平翻轉(zhuǎn)與輕量色彩抖動(dòng)垂直翻轉(zhuǎn)根據(jù)數(shù)據(jù)分布判斷如果數(shù)據(jù)集中車(chē)頭朝向無(wú)規(guī)律則保留。3. U-Net骨干實(shí)現(xiàn)與損失函數(shù)設(shè)計(jì)3.1 為什么課程設(shè)計(jì)選 U-Net 而不是 DeepLabV3模型選型是需要給出理由的環(huán)節(jié)。DeepLabV3 使用空洞卷積在保持分辨率的同時(shí)擴(kuò)大感受野在 PASCAL VOC、Cityscapes 這類(lèi)大規(guī)模分割數(shù)據(jù)集上表現(xiàn)更好但它的結(jié)構(gòu)復(fù)雜、預(yù)訓(xùn)練權(quán)重體積大在課程設(shè)計(jì)這種單 GPU、少量數(shù)據(jù)、短周期的場(chǎng)景下并不劃算。U-Net 的優(yōu)勢(shì)在三個(gè)地方結(jié)構(gòu)對(duì)稱(chēng)、包含跳躍連接、實(shí)現(xiàn)代碼短。編碼器下采樣 4 次解碼器對(duì)應(yīng)上采樣 4 次中間通過(guò) concat 把同尺度的低層特征拼接到解碼器讓邊緣信息不會(huì)因?yàn)橹饘酉虏蓸佣鴣G失。轎車(chē)車(chē)輪與背景的交界處只需要 2 到 4 個(gè)像素的精度U-Net 的跳躍連接恰好能保住這個(gè)級(jí)別的細(xì)節(jié)。另外U-Net 幾乎不依賴(lài)特定預(yù)訓(xùn)練權(quán)重也能在幾百?gòu)垐D上收斂出可用效果屬于訓(xùn)練策略里“從零訓(xùn)練也能有基礎(chǔ)效果”的模型。換個(gè)角度從軟件體系結(jié)構(gòu)的視角看U-Net 是天然的模塊化結(jié)構(gòu)編碼器與解碼器可以拆成兩個(gè)獨(dú)立組件中間通過(guò)接口對(duì)接這個(gè)特征在寫(xiě)軟件設(shè)計(jì)說(shuō)明時(shí)就非常容易畫(huà)出組件圖。對(duì)于課程設(shè)計(jì)考核“架構(gòu)設(shè)計(jì)能力”的評(píng)分項(xiàng)這一條是額外的加分點(diǎn)。3.2 Encoder-Decoder 殘差塊與跳躍連接的代碼實(shí)現(xiàn)U-Net 的核心實(shí)現(xiàn)拆成卷積塊、編碼器、解碼器三部分下面的代碼對(duì)應(yīng) model.py 的核心邏輯class DoubleConv(nn.Module): 雙層卷積塊卷積 批歸一化 ReLUU-Net 的基本組成單元 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_ch3, out_ch1, base_ch64): super().__init__() # base_ch 表示第一層卷積的輸出通道數(shù)之后每下采樣一次通道翻倍 self.enc1 DoubleConv(in_ch, base_ch) self.enc2 DoubleConv(base_ch, base_ch * 2) self.enc3 DoubleConv(base_ch * 2, base_ch * 4) self.enc4 DoubleConv(base_ch * 4, base_ch * 8) self.pool nn.MaxPool2d(2) self.bottleneck DoubleConv(base_ch * 8, base_ch * 16) self.up4 nn.ConvTranspose2d(base_ch * 16, base_ch * 8, 2, stride2) self.dec4 DoubleConv(base_ch * 16, base_ch * 8) self.up3 nn.ConvTranspose2d(base_ch * 8, base_ch * 4, 2, stride2) self.dec3 DoubleConv(base_ch * 8, base_ch * 4) self.up2 nn.ConvTranspose2d(base_ch * 4, base_ch * 2, 2, stride2) self.dec2 DoubleConv(base_ch * 4, base_ch * 2) self.up1 nn.ConvTranspose2d(base_ch * 2, base_ch, 2, stride2) self.dec1 DoubleConv(base_ch * 2, base_ch) self.out nn.Conv2d(base_ch, out_ch, 1) def forward(self, x): e1 self.enc1(x) # [B, 64, H, W] e2 self.enc2(self.pool(e1)) # [B, 128, H/2, W/2] e3 self.enc3(self.pool(e2)) # [B, 256, H/4, W/4] e4 self.enc4(self.pool(e3)) # [B, 512, H/8, W/8] b self.bottleneck(self.pool(e4)) # [B, 1024, H/16, W/16] d4 self.up4(b) d4 torch.cat([d4, e4], dim1) # 跳躍連接沿通道拼接 d4 self.dec4(d4) d3 self.up3(d4) d3 torch.cat([d3, e3], dim1) d3 self.dec3(d3) d2 self.up2(d3) d2 torch.cat([d2, e2], dim1) d2 self.dec2(d2) d1 self.up1(d2) d1 torch.cat([d1, e1], dim1) d1 self.dec1(d1) return self.out(d1)base_ch64是指第一層輸出 64 個(gè)特征圖每下采樣一次通道翻倍到最底層是base_ch * 16 1024。通道數(shù)翻倍與分辨率減半同步進(jìn)行這樣模型的計(jì)算量基本維持穩(wěn)定。跳躍連接對(duì)應(yīng)的代碼是torch.cat([d4, e4], dim1)注意 dim1 是通道拼接不是在空間上疊加這要求編碼器第 4 層輸出與解碼器上采樣后的空間尺寸必須一致否則拼接會(huì)直接報(bào)維度錯(cuò)誤。訓(xùn)練階段輸入尺寸應(yīng)該能被 16 整除這是為什么前面 dataset 里把圖像縮放成 512×512 而不是 500×500 的深層原因。3.3 BCE與Dice Loss的組合邏輯轎車(chē)背景去除是二分類(lèi)問(wèn)題最直接的損失函數(shù)是 BCEBinary Cross Entropy。單獨(dú)使用 BCE 在正負(fù)樣本極度不平衡時(shí)有明顯缺陷一張圖里背景像素經(jīng)常占 80% 以上模型只要把所有像素預(yù)測(cè)為背景就能把 BCE 壓到很低但輸出的 mask 里根本沒(méi)有車(chē)。Dice Loss 是從評(píng)估指標(biāo) Dice 系數(shù)反推出來(lái)的損失函數(shù)直接優(yōu)化“預(yù)測(cè)區(qū)域與真實(shí)區(qū)域的重疊度”對(duì)類(lèi)別不平衡不敏感。實(shí)踐中更穩(wěn)定的是兩者組合即 BCE 加 Diceclass BCEDiceLoss(nn.Module): BCE Dice 組合損失bce_weight 控制兩者占比 def __init__(self, bce_weight0.5): super().__init__() self.bce_weight bce_weight def forward(self, pred, target): pred torch.sigmoid(pred) # 把 logits 壓縮到 0~1 bce F.binary_cross_entropy(pred, target) smooth 1e-6 # 防止分母為 0 的平滑項(xiàng) intersection (pred * target).sum() union pred.sum() target.sum() dice 1 - (2.0 * intersection smooth) / (union smooth) return self.bce_weight * bce (1 - self.bce_weight) * dicebce_weight0.5表示兩者等權(quán)混合。如果訓(xùn)練時(shí)發(fā)現(xiàn) loss 下降到 0.3 左右就停滯但預(yù)測(cè)的掩碼粘連、邊緣粗糙可以調(diào)成bce_weight0.7加大逐像素約束如果發(fā)現(xiàn)訓(xùn)練前期模型輸出的區(qū)域覆蓋不全把bce_weight調(diào)低到 0.3讓 Dice 主導(dǎo)模型聚焦整體結(jié)構(gòu)。下表列出三種損失函數(shù)的適用差異便于答辯時(shí)說(shuō)明損失組合優(yōu)勢(shì)劣勢(shì)適用場(chǎng)景BCE 單獨(dú)梯度平穩(wěn)實(shí)現(xiàn)簡(jiǎn)單正負(fù)樣本不平衡時(shí)偏向背景前景占比均衡時(shí)Dice 單獨(dú)直接優(yōu)化重疊度梯度震蕩明顯小目標(biāo)不穩(wěn)定前景占比極低時(shí)BCE Dice兩者互補(bǔ)收斂平滑需要多調(diào)一個(gè)權(quán)重參數(shù)車(chē)輛分割首選方案需要特別說(shuō)明sigmoid BCE的組合在數(shù)值上不如nn.BCEWithLogitsLoss穩(wěn)定后者內(nèi)部做了數(shù)值保護(hù)。上面的代碼為了直觀展示梯度計(jì)算流程才顯式調(diào)用sigmoid在損失函數(shù)中先 sigmoid 再計(jì)算 BCE梯度會(huì)經(jīng)過(guò)兩次非線(xiàn)性變換實(shí)際項(xiàng)目中直接用nn.BCEWithLogitsLoss會(huì)更安全這個(gè)細(xì)節(jié)可以寫(xiě)進(jìn)課程設(shè)計(jì)的改進(jìn)說(shuō)明里。4. 訓(xùn)練配置與設(shè)計(jì)模式視角下的工程化重構(gòu)4.1 config.py如何統(tǒng)一管理超參數(shù)訓(xùn)練階段涉及的參數(shù)數(shù)量遠(yuǎn)比想象中多學(xué)習(xí)率、批次大小、迭代輪數(shù)、圖像尺寸、數(shù)據(jù)路徑、損失權(quán)重分散在代碼各處時(shí)調(diào)參就是一場(chǎng)災(zāi)難。軟件體系結(jié)構(gòu)課程設(shè)計(jì)里提倡的高內(nèi)聚低耦合落到訓(xùn)練代碼上就是先把所有可調(diào)參數(shù)集中到 config.py 統(tǒng)一管理class Config: 集中管理訓(xùn)練與推理參數(shù)避免魔法數(shù)字散落在各模塊 # 數(shù)據(jù)路徑 train_img_dir dataset/train/input train_mask_dir dataset/train/mask val_img_dir dataset/val/input val_mask_dir dataset/val/mask # 圖像與訓(xùn)練 img_size 512 # 必須能被 16 整除U-Net 下采樣 4 次 batch_size 8 # 顯存不足時(shí)優(yōu)先降到 4而不是調(diào)小圖片 epochs 40 learning_rate 1e-4 # Adam 下 1e-4 比默認(rèn) 1e-3 更穩(wěn) num_workers 4 # Windows 上建議設(shè)為 0否則可能報(bào)錯(cuò) # 損失與優(yōu)化器 bce_weight 0.5 weight_decay 1e-5 save_path checkpoints/best_model.pthimg_size512對(duì)應(yīng)之前提到的 16 整除要求batch_size8在單張 1080Ti 上剛好合適learning_rate1e-4是實(shí)踐中最穩(wěn)的選擇默認(rèn)的1e-3在分割任務(wù)上經(jīng)常出現(xiàn)訓(xùn)練早期 loss 震蕩甚至直接發(fā)散這一點(diǎn)會(huì)在訓(xùn)練循環(huán)里通過(guò)學(xué)習(xí)率策略進(jìn)一步控制。參數(shù)集中之后所有模塊通過(guò)Config.xxx訪(fǎng)問(wèn)參數(shù)后續(xù)做實(shí)驗(yàn)只需要改這一個(gè)文件答辯演示時(shí)也比較直觀。下表匯總了一份可直接套用的訓(xùn)練超參數(shù)規(guī)劃其中優(yōu)化器、學(xué)習(xí)率策略對(duì)收斂影響最顯著參數(shù)推薦值說(shuō)明優(yōu)化器Adam對(duì)學(xué)習(xí)率不敏感適合課程設(shè)計(jì)階段初始學(xué)習(xí)率1e-4高于 1e-3 時(shí)容易震蕩learning rate 策略ReduceLROnPlateau指標(biāo)停滯時(shí)降低為原來(lái)的 0.1批次大小8顯存不足時(shí)降低 batch_size訓(xùn)練輪數(shù)305040 輪左右 val loss 進(jìn)入平臺(tái)期權(quán)重初始化kaiming_normal配合 ReLU 使用4.2 訓(xùn)練循環(huán)的實(shí)現(xiàn)與學(xué)習(xí)率策略訓(xùn)練循環(huán)是每個(gè)課程設(shè)計(jì)必須提交的核心代碼。完整邏輯包括前向傳播、計(jì)算損失、反向傳播、梯度更新、驗(yàn)證集評(píng)估、模型保存六個(gè)步驟。下面的代碼去掉了無(wú)關(guān)的打印信息保留主干def train_one_epoch(model, dataloader, optimizer, criterion, device): model.train() total_loss 0.0 for imgs, masks in dataloader: imgs imgs.to(device) masks masks.to(device) optimizer.zero_grad() outputs model(imgs) loss criterion(outputs, masks) loss.backward() optimizer.step() total_loss loss.item() return total_loss / len(dataloader) def evaluate(model, dataloader, criterion, device): model.eval() total_loss 0.0 with torch.no_grad(): for imgs, masks in dataloader: imgs imgs.to(device) masks masks.to(device) outputs model(imgs) loss criterion(outputs, masks) total_loss loss.item() return total_loss / len(dataloader)訓(xùn)練主循環(huán)部分結(jié)合前面提到的學(xué)習(xí)率策略加進(jìn)去ReduceLROnPlateau的完整調(diào)用model UNet(in_ch3, out_ch1).to(device) optimizer torch.optim.Adam(model.parameters(), lrConfig.learning_rate, weight_decayConfig.weight_decay) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, modemin, factor0.1, patience5, verboseTrue ) best_val_loss float(inf) for epoch in range(Config.epochs): train_loss train_one_epoch(model, train_loader, optimizer, criterion, device) val_loss evaluate(model, val_loader, criterion, device) # 檢測(cè) val_loss 連續(xù)多個(gè) epoch 不下降時(shí)降低學(xué)習(xí)率 scheduler.step(val_loss) if val_loss best_val_loss: best_val_loss val_loss torch.save(model.state_dict(), Config.save_path) print(fEpoch {epoch:02d}, saved best model, val_loss: {val_loss:.4f})optimizer.zero_grad()放在每個(gè) batch 之前作用是清空上一次反向傳播累積的梯度這個(gè)步驟遺漏會(huì)導(dǎo)致梯度跨 batch 累加、loss 異常波動(dòng)。torch.save(model.state_dict(), ...)只保存權(quán)重不保存模型結(jié)構(gòu)加載時(shí)需要先用UNet()實(shí)例化模型再load_state_dict。ReduceLROnPlateau的modemin表示監(jiān)控指標(biāo)越低越好factor0.1表示每次降為原來(lái)的十分之一設(shè)成 0.5 會(huì)更平滑但會(huì)拉長(zhǎng)訓(xùn)練時(shí)間。4.3 用策略模式與工廠(chǎng)模式解耦數(shù)據(jù)與損失模塊軟件體系結(jié)構(gòu)與設(shè)計(jì)模式課程設(shè)計(jì)的核心考核點(diǎn)體現(xiàn)在這里。數(shù)據(jù)加載與損失函數(shù)是兩個(gè)最容易替換的擴(kuò)展點(diǎn)換數(shù)據(jù)集、換損失函數(shù)是調(diào)優(yōu)過(guò)程中的高頻操作。如果代碼里到處是if dataset_type carvana這樣的分支每加一種數(shù)據(jù)集就要?jiǎng)右延写a違反開(kāi)閉原則。用工廠(chǎng)模式封裝數(shù)據(jù)加載器用策略模式封裝損失函數(shù)class LossFactory: 策略模式根據(jù)名稱(chēng)返回對(duì)應(yīng)的損失函數(shù)實(shí)例 _losses { bce_dice: BCEDiceLoss, dice: DiceLoss, bce: nn.BCEWithLogitsLoss, } classmethod def create(cls, name, **kwargs): if name not in cls._losses: raise ValueError(fUnknown loss: {name}) return cls._losses[name](**kwargs) class DatasetFactory: 工廠(chǎng)模式按數(shù)據(jù)集類(lèi)型構(gòu)造對(duì)應(yīng)的 Dataset staticmethod def create(dataset_type, img_dir, mask_dir, **kwargs): if dataset_type car: return CarSegDataset(img_dir, mask_dir, **kwargs) if dataset_type general: return GeneralSegDataset(img_dir, mask_dir, **kwargs) raise ValueError(fUnsupported dataset: {dataset_type})兩個(gè)工廠(chǎng)類(lèi)的設(shè)計(jì)意圖不同LossFactory是對(duì)創(chuàng)建邏輯的集中封裝用字典注冊(cè)類(lèi)名與類(lèi)的映射新增損失函數(shù)時(shí)只需要在_losses字典里加一行其余訓(xùn)練代碼零改動(dòng)DatasetFactory做的是條件分發(fā)當(dāng)新增一種數(shù)據(jù)集時(shí)不必在每個(gè)用到Dataset的地方加判斷。實(shí)際項(xiàng)目中如果只做課程設(shè)計(jì)不需要過(guò)度設(shè)計(jì)但這兩個(gè)工廠(chǎng)類(lèi)的代碼量很少又恰好覆蓋了設(shè)計(jì)模式的兩個(gè)經(jīng)典考核點(diǎn)屬于性?xún)r(jià)比很高的工程化改造。5. 從IoU到批量摳圖評(píng)估腳本與推理后處理5.1 IoU / Dice評(píng)估與常見(jiàn)統(tǒng)計(jì)誤區(qū)訓(xùn)練完成后需要回答一個(gè)關(guān)鍵問(wèn)題這個(gè)模型到底好不好。評(píng)估指標(biāo)不應(yīng)只看 loss因?yàn)?BCE Loss 很小不代表分割結(jié)果好。語(yǔ)義分割的標(biāo)準(zhǔn)評(píng)估指標(biāo)是 IoU即預(yù)測(cè)區(qū)域與真實(shí)區(qū)域的交集除以并集。另一個(gè)常用指標(biāo)是 Dice 系數(shù)它與 IoU 之間可以互相換算Dice 2 * IoU / (1 IoU)。計(jì)算代碼很短但統(tǒng)計(jì)過(guò)程有一個(gè)常見(jiàn)誤區(qū)def compute_metrics(pred_mask, gt_mask, threshold0.5): pred_mask: 模型輸出的概率圖, gt_mask: 真實(shí)標(biāo)簽 pred_bin (pred_mask threshold).astype(int) gt_bin (gt_mask threshold).astype(int) intersection (pred_bin gt_bin).sum() union (pred_bin | gt_bin).sum() iou intersection / union dice (2 * intersection) / (pred_bin.sum() gt_bin.sum()) return iou, dice誤區(qū)在于不要把 batch 內(nèi)所有樣本的 IoU 先求平均而應(yīng)該先累加所有樣本的 intersection 和 union最后再統(tǒng)一相除兩種統(tǒng)計(jì)方式在小樣本測(cè)試集上可能相差 2 到 3 個(gè)百分點(diǎn)。誤用場(chǎng)景是當(dāng)某張圖完全沒(méi)有車(chē)時(shí)union 為 0直接計(jì)算會(huì)產(chǎn)生除零錯(cuò)誤正確做法是跳過(guò)該樣本或在分子分母同時(shí)加平滑項(xiàng)。課程設(shè)計(jì)里只需寫(xiě)清楚你用的是哪種統(tǒng)計(jì)口徑。5.2 單張推理與批量摳圖的可執(zhí)行步驟最后一步是把訓(xùn)練好的模型應(yīng)用到真實(shí)圖片上。推理腳本需要完成加載權(quán)重、預(yù)處理、前向傳播、后處理、保存結(jié)果五個(gè)步驟。后處理部分有一個(gè)容易被忽略的操作預(yù)測(cè)出的概率圖直接以 0.5 為閾值二值化后可能會(huì)出現(xiàn)一些小面積噪點(diǎn)或細(xì)小孔洞用形態(tài)學(xué)開(kāi)運(yùn)算去除噪點(diǎn)、用閉運(yùn)算填充孔洞是標(biāo)準(zhǔn)做法import cv2 import torch def inference_one_image(model, img_path, device, save_path, thresh0.5): # 1. 預(yù)處理讀圖、縮放、歸一化、轉(zhuǎn) tensor img cv2.imread(img_path) img_rgb cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img_resized cv2.resize(img_rgb, (Config.img_size, Config.img_size)) # 2. 歸一化ImageNet 均值方差 img_norm img_resized / 255.0 img_tensor torch.from_numpy(img_norm).permute(2, 0, 1).unsqueeze(0).float() img_tensor img_tensor.to(device) # 3. 前向傳播得到概率圖 model.eval() with torch.no_grad(): prob torch.sigmoid(model(img_tensor)).cpu().numpy()[0, 0] # 4. 二值化 形態(tài)學(xué)后處理 mask (prob thresh).astype(np.uint8) * 255 kernel cv2.getStructuringElement(cv2.MORPH_RECT, (5, 5)) mask cv2.morphologyEx(mask, cv2.MORPH_OPEN, kernel) # 先開(kāi)運(yùn)算去噪點(diǎn) mask cv2.morphologyEx(mask, cv2.MORPH_CLOSE, kernel) # 再閉運(yùn)算填空洞 # 5. 原圖尺寸恢復(fù)并疊加背景 mask_resized cv2.resize(mask, (img.shape[1], img.shape[0]), interpolationcv2.INTER_NEAREST) result img.copy() result[mask_resized 0] [255, 255, 255] # 背景置為白色 cv2.imwrite(save_path, result)批量推理時(shí)把inference_one_image放進(jìn)一個(gè)循環(huán)遍歷目錄下所有圖片即可無(wú)需額外寫(xiě)多進(jìn)程版本。MORPH_OPEN先腐蝕再膨脹能去除小于卷積核尺寸的白色噪點(diǎn)MORPH_CLOSE先膨脹再腐蝕能填充黑色區(qū)域里的白色空洞。對(duì)于轎車(chē)背景去除5×5 的卷積核大小適中改大會(huì)讓車(chē)輪邊緣的細(xì)小結(jié)構(gòu)被抹掉。最終保存結(jié)果時(shí)用INTER_NEAREST把 mask 恢復(fù)為原圖尺寸保持邊緣銳利不產(chǎn)生鋸齒色偏。跑通這個(gè)流程之后就完成了從課程設(shè)計(jì)考核的代碼邏輯說(shuō)明到真實(shí)場(chǎng)景應(yīng)用的完整銜接。本文還有配套的精品資源點(diǎn)擊獲取