系統(tǒng):從原理到源碼復(fù)現(xiàn)實(shí)戰(zhàn)解析)
簡介圖像修復(fù)是計算機(jī)視覺中極具實(shí)用價值的研究方向旨在通過算法自動恢復(fù)圖像中缺失、遮擋或破損區(qū)域的像素內(nèi)容。傳統(tǒng)方法基于紋理合成或插值在面對大面積缺失時往往力不從心而基于深度學(xué)習(xí)的生成模型則能借助語義理解推斷出合理且自然的內(nèi)容。以生成對抗網(wǎng)絡(luò)GAN為核心、U-Net為生成器主干配合感知損失與對抗損失的聯(lián)合優(yōu)化已成為當(dāng)前主流技術(shù)范式。這類系統(tǒng)可用于老照片修復(fù)、物體移除、影視后期及醫(yī)學(xué)影像處理等真實(shí)場景。在實(shí)際工程落地中基于PyTorch搭建的修復(fù)項(xiàng)目需要重點(diǎn)關(guān)注數(shù)據(jù)Mask生成、生成器與判別器結(jié)構(gòu)、損失函數(shù)配比以及訓(xùn)練推理流程的穩(wěn)定性。本文圍繞一套完整的基于PyTorch的圖像修復(fù)源碼系統(tǒng)梳理其設(shè)計思路、環(huán)境配置、核心模塊實(shí)現(xiàn)與訓(xùn)練推理細(xì)節(jié)并針對常見問題給出排查方法適合需要學(xué)習(xí)或二次開發(fā)相關(guān)系統(tǒng)的開發(fā)者參考。 前兩天清理硬盤翻出一個標(biāo)注為“基于PyTorch的圖像修復(fù)系統(tǒng)”的源碼壓縮包順手解壓跑了一輪。圖像修復(fù)Image Inpainting是計算機(jī)視覺里一個很實(shí)用的方向輸入一張被遮擋、劃痕或物體缺失的圖片模型會自動把缺失區(qū)域補(bǔ)出來比傳統(tǒng)克隆印章和插值算法自然得多。這類源碼適合正在學(xué)PyTorch的開發(fā)者研究也適合有圖像修復(fù)需求的產(chǎn)品工程師做二次開發(fā)。這篇文章就把整個系統(tǒng)的設(shè)計思路、環(huán)境搭建、核心代碼和復(fù)現(xiàn)過程中的坑完整梳理一遍。這類項(xiàng)目在網(wǎng)上有不少公開實(shí)現(xiàn)但大多只給了模型結(jié)構(gòu)訓(xùn)練流程寫得模糊數(shù)據(jù)預(yù)處理也不完整。我拿到這個壓縮包后重點(diǎn)檢查了三個東西數(shù)據(jù)加載和Mask生成邏輯、生成器和判別器的實(shí)現(xiàn)細(xì)節(jié)、訓(xùn)練與推理的入口配置。只要這三個地方能跑通整個系統(tǒng)基本就能復(fù)現(xiàn)。下面從頭開始拆。1. 項(xiàng)目整體設(shè)計與核心思路1.1 圖像修復(fù)任務(wù)與適用場景圖像修復(fù)的核心任務(wù)是給定一張帶有缺失區(qū)域的圖像以及標(biāo)明缺失位置的Mask模型需要生成與周圍上下文一致的像素內(nèi)容。用數(shù)學(xué)語言描述輸入是原始圖像 (I) 和掩碼 (M)其中 (M1) 的位置代表缺失區(qū)域修復(fù)目標(biāo)是估計 (I_{out})使得 (I_{out}) 在 Mask 區(qū)域內(nèi)與真實(shí)內(nèi)容 (I_{gt}) 盡可能接近同時在視覺上沒有明顯接縫。傳統(tǒng)方法比如Telea算法和基于Patch匹配的紋理合成對細(xì)小劃痕處理還可以遇到大塊缺失區(qū)域就無能為力要么模糊要么紋理重復(fù)。深度學(xué)習(xí)模型能夠依據(jù)高層語義信息推斷出合理內(nèi)容比如補(bǔ)全一棟被遮擋的樓、一條被電線穿過的天空甚至生成原來并不確定的細(xì)節(jié)。實(shí)際落地場景主要有幾類老照片修復(fù)去除照片上的折痕、污漬、霉斑同時恢復(fù)背景紋理。圖像編輯把不需要的元素路人、水印、雜物從畫面中移除再用修復(fù)算法填充背景。影視后期對拍攝時無法規(guī)避的穿幫物體做內(nèi)容補(bǔ)全。醫(yī)學(xué)影像去除掃描圖像中的金屬偽影或運(yùn)動偽影。這個源碼的通用性還不錯只要你準(zhǔn)備了合適的Mask分布和數(shù)據(jù)集微調(diào)一下就能適配上述場景。1.2 源碼模塊結(jié)構(gòu)與架構(gòu)選型解壓壓縮包后典型工程結(jié)構(gòu)大致如下. ├── checkpoints/ # 模型權(quán)重保存位置 ├── configs/ # 訓(xùn)練和推理參數(shù)配置文件 ├── data/ # 數(shù)據(jù)集加載、Mask生成、數(shù)據(jù)增強(qiáng) ├── models/ # 生成器、判別器、損失函數(shù)定義 ├── utils/ # 圖像處理工具、評估指標(biāo) ├── train.py # 訓(xùn)練入口 ├── infer.py # 推理入口 ├── requirements.txt # 依賴列表 └── README.md # 使用說明拿到源碼千萬別急著訓(xùn)練先看README和requirements確認(rèn)作者用了哪個PyTorch版本、數(shù)據(jù)是什么格式、訓(xùn)練輸入尺寸是多大。很多人復(fù)現(xiàn)失敗問題不在模型而在版本和配置不一致。比如PyTorch 1.x和2.x在某些算子行為上差異不小模型代碼里如果用了torch.nn.functional.interpolate的align_corners參數(shù)版本不同結(jié)果可能完全不同。架構(gòu)選型上這個項(xiàng)目遵循了主流修復(fù)框架的設(shè)計生成器用U-Net變體配合部分卷積或門控卷積判別器用PatchGAN損失函數(shù)由像素重建損失、感知損失和對抗損失組成。這樣組合的原因是重建損失讓網(wǎng)絡(luò)學(xué)到穩(wěn)定的基礎(chǔ)結(jié)構(gòu)感知損失從特征層面保證語義一致對抗損失負(fù)責(zé)讓修復(fù)區(qū)域紋理更真實(shí)。2. 環(huán)境準(zhǔn)備與依賴安裝2.1 PyTorch環(huán)境搭建的完整步驟先強(qiáng)調(diào)一件事圖像修復(fù)訓(xùn)練必須GPU純CPU跑會慢得讓人懷疑人生。配置環(huán)境我建議用Anaconda管理虛擬環(huán)境不要直接裝在base環(huán)境里否則后面項(xiàng)目一多依賴沖突能煩死你。創(chuàng)建獨(dú)立環(huán)境并激活conda create -n inpaint python3.8 -y conda activate inpaint接下來安裝GPU版PyTorch。先看自己機(jī)器的CUDA支持情況nvidia-smi輸出右上角能看到驅(qū)動支持的最高CUDA版本比如12.1那么安裝PyTorch時選擇cuda 12.1或比它低的版本都可以。注意nvidia-smi顯示的CUDA版本是驅(qū)動支持的版本不是當(dāng)前環(huán)境已經(jīng)裝好的版本。實(shí)際安裝PyTorch時conda會自動把配套的CUDA runtime和cuDNN一起裝進(jìn)虛擬環(huán)境所以不需要單獨(dú)裝CUDA Toolkit除非你要編譯自定義算子。安裝命令示例conda install pytorch torchvision pytorch-cuda11.8 -c pytorch -c nvidia如果下載速度不穩(wěn)定可以用清華源加速但要注意conda和pip的源不要混著換出問題。我這里不展開說鏡像配置重點(diǎn)是你的pytorch-cuda版本要和顯卡驅(qū)動兼容。裝完之后驗(yàn)證python -c import torch; print(torch.__version__, torch.cuda.is_available())輸出True說明CUDA可用否則要重新檢查安裝步驟。2.2 依賴庫安裝與預(yù)訓(xùn)練權(quán)重準(zhǔn)備requirements.txt里面一般會有這些庫torch torchvision opencv-python numpy pillow tqdm tensorboard scikit-image批量安裝pip install -r requirements.txt如果需求里有pytorch-msssim、lpips這類評估指標(biāo)庫建議一并裝上后面測試效果會用到。安裝時如果遇到opencv-python編譯慢可以直接用阿里云或豆瓣的鏡像源安裝。數(shù)據(jù)集方面如果只是想快速跑通流程不需要一上來就下載Places2這么大的數(shù)據(jù)集??梢韵饶肅elebA或COCO的子集試跑甚至用自己拍攝的幾十張照片也能驗(yàn)證。但需要注意修復(fù)模型對數(shù)據(jù)量有要求數(shù)據(jù)太少容易過擬合表現(xiàn)就是訓(xùn)練loss很低換一張新圖效果崩。公開數(shù)據(jù)集我常用Places2和CelebA-HQ前者適合場景補(bǔ)全后者適合人臉修復(fù)。預(yù)訓(xùn)練權(quán)重一般放在checkpoints目錄。加載權(quán)重時報size mismatch是最常見的問題原因通常是模型結(jié)構(gòu)和state_dict鍵名對不上。遇到這種情況先用torch.load把權(quán)重load進(jìn)來打印model_state_dict的keys和當(dāng)前模型的keys做對比缺哪個補(bǔ)哪個多了的刪掉再加載就順暢了。2.3 訓(xùn)練與推理的關(guān)鍵參數(shù)解析打開配置文件常見參數(shù)如下表所示參數(shù)推薦值說明image_size256 / 512輸入圖像尺寸越大越消耗顯存batch_size4 ~ 8顯存不足時優(yōu)先調(diào)低learning_rate1e-4 ~ 2e-4Adam優(yōu)化器常用范圍lambda_rec10重建損失權(quán)重lambda_perceptual0.1感知損失權(quán)重lambda_adv1對抗損失權(quán)重iterations100000按迭代數(shù)訓(xùn)練更常見save_interval5000每隔多少步保存一次模型這里重點(diǎn)說下?lián)p失權(quán)重的影響。重建損失權(quán)重太高模型傾向于輸出平滑結(jié)果細(xì)節(jié)會糊對抗損失權(quán)重太高訓(xùn)練不穩(wěn)定容易出現(xiàn)色彩失真。我用過的組合里先固定lambda_rec10, lambda_perceptual0.1再逐步從0.1調(diào)到1效果會更可控。訓(xùn)練過程中需要同時關(guān)注生成器和判別器的loss比例判別器loss長期接近0說明生成器完全打不過判別器梯度幾乎沒有需要降低判別器學(xué)習(xí)率。3. 核心模塊源碼解析與實(shí)操3.1 數(shù)據(jù)加載與Mask生成邏輯數(shù)據(jù)加載是第一個容易踩坑的地方。Pytorch的Dataset類返回的樣本格式必須和模型輸入對齊。修復(fù)任務(wù)中輸入不是單一圖像而是“損壞圖像 Mask”的組合。很多實(shí)現(xiàn)會在__getitem__里做以下操作讀取完整圖像并做隨機(jī)裁剪或resize。生成隨機(jī)Mask。根據(jù)Mask對圖像做損壞處理比如把Mask區(qū)域像素置為0。將Mask歸一化為0/1并與圖像在通道維度拼接得到4通道輸入。返回(input_tensor, mask_tensor, gt_tensor, mask_for_loss)。Mask生成是重點(diǎn)。如果Mask全是隨機(jī)矩形模型只會補(bǔ)矩形區(qū)域遇到真實(shí)劃痕就失效。我在項(xiàng)目里看到比較實(shí)用的做法是生成三種類型隨機(jī)矩形塊模擬物體遮擋。隨機(jī)線條和曲線模擬劃痕。不規(guī)則多邊形模擬污漬。實(shí)現(xiàn)時可以用OpenCV畫線、畫多邊形再配合膨脹腐蝕讓Mask邊緣更自然。簡單示例import cv2 import numpy as np def random_mask(height, width): mask np.zeros((height, width), dtypenp.uint8) # 隨機(jī)矩形 x, y, w, h np.random.randint(0, width//3, 4) mask[y:yh, x:xw] 255 # 隨機(jī)線條 pts np.random.randint(0, height, (2, 2)) cv2.line(mask, tuple(pts[0]), tuple(pts[1]), 255, 10) mask cv2.dilate(mask, np.ones((5, 5), np.uint8)) return mask / 255.0這只是示例正式項(xiàng)目里會加入更多形態(tài)變化。Mask生成時一定要保證訓(xùn)練和推理的Mask分布一致。如果訓(xùn)練時Mask區(qū)域都是小面積推理時給一個大面積Mask模型就會表現(xiàn)得很差。3.2 生成器與判別器的網(wǎng)絡(luò)實(shí)現(xiàn)細(xì)節(jié)生成器最基礎(chǔ)的做法是把U-Net輸入改成4通道中間換幾個殘差塊輸出3通道。但標(biāo)準(zhǔn)卷積在處理Mask區(qū)域時有天然缺陷卷積核對所有像素一視同仁缺失區(qū)域的零像素會污染特征導(dǎo)致修復(fù)結(jié)果有灰斑和邊界模糊。所以成熟的修復(fù)項(xiàng)目會用部分卷積或門控卷積。部分卷積的核心思路是在卷積操作時只對有效像素做計算讓Mask區(qū)域不參與特征更新。輸出是這樣算的out W * (X * M) / sum(M) b mask_out 1 if sum(M) 0 else 0每一步卷積后Mask也要跟著更新這樣網(wǎng)絡(luò)能自動判斷哪些位置已經(jīng)被修復(fù)哪些還需要繼續(xù)生成。用PyTorch自定義部分卷積層要注意實(shí)現(xiàn)細(xì)節(jié)比如分母sum(M)不能為0需要加一個epsilon。判別器用PatchGAN。簡單來說它不是輸出一個全局真/假標(biāo)量而是輸出一個特征圖比如16x16的矩陣每個值代表輸入圖像局部區(qū)域是真還是假。這么做可以讓判別器關(guān)注局部紋理一致性避免出現(xiàn)“整體像局部崩”的情況。PatchGAN實(shí)現(xiàn)并不復(fù)雜import torch.nn as nn class PatchDiscriminator(nn.Module): def __init__(self, in_channels3): super().__init__() self.layers nn.Sequential( nn.Conv2d(in_channels, 64, 4, 2, 1), nn.LeakyReLU(0.2), nn.Conv2d(64, 128, 4, 2, 1), nn.BatchNorm2d(128), nn.LeakyReLU(0.2), nn.Conv2d(128, 256, 4, 2, 1), nn.BatchNorm2d(256), nn.LeakyReLU(0.2), nn.Conv2d(256, 1, 4, 1, 1) ) def forward(self, x): return self.layers(x)實(shí)際項(xiàng)目里輸入可能是4通道含Mask需要微調(diào)第一層輸入尺寸。3.3 損失函數(shù)組合與訓(xùn)練循環(huán)模型的優(yōu)化目標(biāo)不是一個loss而是多個loss的加權(quán)和。典型組合如下loss_rec l1_loss(pred, gt, mask) # 只計算Mask內(nèi)區(qū)域 loss_perceptual perceptual_loss(pred, gt) # VGG特征距離 loss_adv gan_loss(discriminator(pred), real_label) total_loss loss_rec * w_rec loss_perceptual * w_perceptual loss_adv * w_adv重建損失用L1比L2好L2會過度懲罰大誤差導(dǎo)致輸出偏向平均顏色邊緣會變模糊。L1對離群值更魯棒能保留更多細(xì)節(jié)。感知損失一般用ImageNet預(yù)訓(xùn)練的VGG16取relu1_1, relu2_1, relu3_1, relu4_1幾層特征計算生成圖與真實(shí)圖特征之間的L1距離。注意使用預(yù)訓(xùn)練VGG時輸入圖像要做相同的歸一化否則特征值分布不同感知損失的意義會打折扣。訓(xùn)練循環(huán)里生成器和判別器交替更新。一個常見的策略是每個iteration里先更新生成器再更新判別器或者每更新一次生成器更新兩次判別器。這個項(xiàng)目源碼里如果沒有控制判別器更新頻率訓(xùn)練出現(xiàn)震蕩可以自己加一個if step % 2 0的更新判別器邏輯。訓(xùn)練日志建議記錄step, total_loss, rec_loss, perc_loss, adv_loss, d_loss, psnr, ssim偽代碼可以這樣寫for step in range(total_iterations): real_img, real_gt, mask next(data_loader) masked_img real_img * (1 - mask) # 訓(xùn)練生成器 fake_img generator(torch.cat([masked_img, mask], dim1)) rec_loss l1_loss(fake_img, real_gt, mask) perc_loss vgg_loss(fake_img, real_gt) adv_loss gan_loss(discriminator(fake_img), real_label) g_loss rec_loss * w_rec perc_loss * w_perc adv_loss * w_adv g_optimizer.zero_grad() g_loss.backward() g_optimizer.step() # 訓(xùn)練判別器 real_pred discriminator(real_gt) fake_pred discriminator(fake_img.detach()) d_loss (gan_loss(real_pred, real_label) gan_loss(fake_pred, fake_label)) * 0.5 d_optimizer.zero_grad() d_loss.backward() d_optimizer.step()這里面有一個很容易忽略的細(xì)節(jié)計算重建損失時建議把mask也作為權(quán)重傳入只計算Mask區(qū)域內(nèi)的像素差。如果把全圖都算進(jìn)去背景區(qū)域占了絕對主導(dǎo)模型很容易學(xué)到“把背景復(fù)制一下就行”對被遮擋區(qū)域毫無生成能力。4. 實(shí)戰(zhàn)從零訓(xùn)練與圖像修復(fù)推理4.1 用自己的數(shù)據(jù)集跑通訓(xùn)練建議第一次跑用一個小數(shù)據(jù)集比如從公開數(shù)據(jù)集里挑500張圖片設(shè)置訓(xùn)練步數(shù)1000步目標(biāo)不是效果好而是驗(yàn)證整個鏈路是通的數(shù)據(jù)加載正確、模型前向正常、損失能下降、權(quán)重能保存和加載。這一步通過后再全量訓(xùn)練能省下大量排查時間。使用自己的數(shù)據(jù)集時先創(chuàng)建一個data_list.txt每一行是圖片路徑/path/to/train/000001.jpg /path/to/train/000002.jpg ...然后修改config把data_root指向這個txt設(shè)置好image_size、batch_size等參數(shù)。執(zhí)行訓(xùn)練python train.py --config configs/train_config.yaml訓(xùn)練啟動后觀察前幾個迭代的日志。正常情況下loss應(yīng)該在穩(wěn)步下降PSNR和SSIM緩慢上升。如果loss曲線劇烈震蕩我的經(jīng)驗(yàn)是先把學(xué)習(xí)率調(diào)低一個數(shù)量級比如從1e-4調(diào)到1e-5再看是否穩(wěn)定。如果生成器loss和判別器loss走勢極不平衡參考上一節(jié)提到的調(diào)整訓(xùn)練頻率。訓(xùn)練時長方面單張RTX 3090跑256×256輸入、batch_size810萬步大約需要2天。如果想快速驗(yàn)證可以先用--max_iters 2000跑一小段。4.2 推理流程與效果對比推理流程比訓(xùn)練簡單但精度問題同樣不可忽視。命令類似python infer.py --image samples/test.jpg --mask samples/mask.png --checkpoint output/latest.pth --output result.png推理腳本內(nèi)部一般做這四步讀取圖像和Mask統(tǒng)一resize到訓(xùn)練尺寸。圖像和Mask拼接經(jīng)過生成器前向計算。生成結(jié)果與原始圖像融合保留原始圖中未損壞區(qū)域。保存輸出圖像必要時做后處理。這里的關(guān)鍵是融合。不能把生成器的整個輸出直接作為結(jié)果因?yàn)樗诜荕ask區(qū)域也會產(chǎn)生偏移導(dǎo)致原圖背景被改。正確的融合方式是result original_image * (1 - mask) generated_image * mask如果發(fā)現(xiàn)修復(fù)區(qū)域邊緣有接縫可以對Mask做一次高斯模糊或者把Mask膨脹幾個像素讓融合過渡更自然。顏色偏色問題時檢查推理時是否做了和訓(xùn)練一樣的歸一化。很多項(xiàng)目訓(xùn)練時把像素映射到[-1,1]推理時忘了減均值除方差出來的圖就會明顯偏色。效果對比時我習(xí)慣把原圖、Mask、修復(fù)結(jié)果拼在一張畫布里方便肉眼評估。重點(diǎn)關(guān)注三塊邊緣過渡是否自然、紋理是否重復(fù)、顏色是否一致。4.3 評估指標(biāo)PSNR/SSIM如何看在有真實(shí)完整圖像的測試集上可以計算PSNR和SSIM。PSNR越高越好通常修復(fù)模型能到25~35dBSSIM越接近1越好。另一個更接近主觀感受的指標(biāo)是LPIPS值越低越好。但客觀指標(biāo)和人的感受經(jīng)常不一致。GAN生成的結(jié)果紋理細(xì)膩但像素級偏移大PSNR可能反而不如模糊的結(jié)果。因此評估時必須同時看主觀效果圖不能只看分?jǐn)?shù)。我自己的習(xí)慣是先用PSNR/SSIM篩掉明顯不行的模型再用肉眼對比候選模型的修復(fù)圖最后根據(jù)應(yīng)用場景做決定。如果業(yè)務(wù)場景對真實(shí)性要求高LPIPS和人工評測優(yōu)先級要提前。5. 常見問題與排查技巧實(shí)錄5.1 訓(xùn)練不收斂或損失異常訓(xùn)練時遇到loss變成NaN最先檢查三件事輸入圖像有沒有全黑或全白、Mask是否全0、學(xué)習(xí)率是否過大。圖像修復(fù)模型輸入是全零區(qū)域時前向計算容易出現(xiàn)梯度異常??梢韵劝褜W(xué)習(xí)率降到1e-5跑一個iteration看看如果還NaN就要檢查數(shù)據(jù)歸一化和模型權(quán)重初始化。另一個常見情況是loss快速下降但不代表效果好。生成器可能學(xué)會了一種投機(jī)取巧的方式非Mask區(qū)域直接復(fù)制Mask區(qū)域輸出平均色。這樣重建loss很低PSNR也可能不差但視覺上一塌糊涂。解決辦法是提高感知損失和對抗損失的權(quán)重并定期人工查看生成圖。我遇到的訓(xùn)練崩潰還有一個隱藏原因用FP16混合精度訓(xùn)練時loss數(shù)值不穩(wěn)定。如果源碼默認(rèn)開了amp可以關(guān)閉再試或者調(diào)整grad_scaler的初始化比例。修復(fù)模型對精度比較敏感穩(wěn)定優(yōu)先于速度。5.2 修復(fù)結(jié)果模糊或偽影嚴(yán)重模糊是因?yàn)槟P蜎]有高頻細(xì)節(jié)常見原因包括重建損失權(quán)重過高、網(wǎng)絡(luò)容量不足、輸入分辨率過低。我的做法是先降低lambda_rec同時確保lambda_perceptual不是0否則模型會大量丟失紋理。如果還模糊考慮換更大的生成器或增加中間通道數(shù)。偽影一般出現(xiàn)在GAN訓(xùn)練不穩(wěn)定時。表現(xiàn)為修復(fù)區(qū)域有斑點(diǎn)、彩色條紋或過銳利的邊緣??梢钥紤]降低對抗損失權(quán)重。使用LSGAN或Hinge Loss代替標(biāo)準(zhǔn)BCE。判別器增加譜歸一化。訓(xùn)練前期凍結(jié)判別器只訓(xùn)練生成器若干千步。偽影還可能與推理分辨率有關(guān)。如果訓(xùn)練是256×256推理時給512×512的圖模型沒見過這么大尺寸容易產(chǎn)生結(jié)構(gòu)畸變。此時優(yōu)先保證推理尺寸和訓(xùn)練一致再考慮是否使用支持任意分辨率的結(jié)構(gòu)。5.3 環(huán)境兼容性與顯存不足問題環(huán)境兼容性問題中最常見的是PyTorch和CUDA版本不匹配。安裝PyTorch時如果選了比驅(qū)動支持更高的CUDA版本會直接報CUDA driver version is insufficient。解決辦法是用nvidia-smi確認(rèn)支持版本然后重新安裝匹配的PyTorch。顯存不足的報錯形式一般是RuntimeError: CUDA out of memory. Tried to allocate X MiB我推薦的排查順序降低batch_size到2甚至1。降低圖像尺寸比如從512降到256。設(shè)置torch.backends.cudnn.benchmarkTrue有時候能省顯存。使用梯度累積模擬大batch。開啟torch.utils.checkpoint梯度檢查點(diǎn)用計算換顯存。使用混合精度訓(xùn)練FP16。注意DataLoader的num_workers調(diào)大并不會減少顯存占用反而可能因?yàn)槎噙M(jìn)程緩存導(dǎo)致整體內(nèi)存上漲。顯存不足時可以先關(guān)掉驗(yàn)證集評估因?yàn)轵?yàn)證過程也占用顯存。5.4 常見問題速查表問題可能原因解決辦法loss為NaN學(xué)習(xí)率過大或輸入包含無效值降低學(xué)習(xí)率檢查數(shù)據(jù)歸一化修復(fù)區(qū)域模糊重建損失權(quán)重過高或網(wǎng)絡(luò)容量不足降低lambda_rec增加網(wǎng)絡(luò)寬度邊界有接縫推理融合未處理Mask對Mask做膨脹或高斯模糊后融合顏色偏色推理歸一化與訓(xùn)練不一致統(tǒng)一預(yù)處理邏輯CUDA out of memorybatch_size/分辨率過大調(diào)小尺寸使用FP16梯度累積預(yù)訓(xùn)練權(quán)重加載失敗模型結(jié)構(gòu)和權(quán)重鍵名不匹配打印keys逐一對比手動修改state_dict大型Mask修復(fù)效果差訓(xùn)練時Mask面積分布不匹配增加不規(guī)則大Mask的數(shù)據(jù)增強(qiáng)6. 項(xiàng)目擴(kuò)展方向與個人心得6.1 如何擴(kuò)展到自己的業(yè)務(wù)場景單純跑通一個開源項(xiàng)目不算本事能把它用起來才是目標(biāo)。假設(shè)你要做老照片修復(fù)需要自己準(zhǔn)備一批帶劃痕的圖片或者模擬劃痕生成訓(xùn)練數(shù)據(jù)要移除水印就得專門生成文字區(qū)域的Mask。千萬不要通用模型一把梭效果大概率不如預(yù)期。擴(kuò)展方面可以考慮模型導(dǎo)出。PyTorch模型可以用ONNX導(dǎo)出再用TensorRT做推理加速方便部署到服務(wù)端torch.onnx.export( generator, (dummy_img, dummy_mask), inpaint.onnx, input_names[image, mask], output_names[output], dynamic_axes{image: {0: batch}, mask: {0: batch}} )ONNX導(dǎo)出時要注意自定義層比如部分卷積中的Mask更新可能不支持導(dǎo)出需要把部分卷積改寫成兼容算子或者用torch.onnx.export時設(shè)置custom_opsets。這塊我踩過坑建議先導(dǎo)出一個最簡模型驗(yàn)證通道。另外把修復(fù)系統(tǒng)包裝成API接口也很實(shí)用。用FastAPI加載模型接收圖像和Mask返回修復(fù)結(jié)果這樣就能接入小程序、網(wǎng)頁或后端服務(wù)。要注意的是并發(fā)場景下顯存可能不夠可以考慮單進(jìn)程單GPU隊(duì)列方式避免多個請求同時占用顯存導(dǎo)致OOM。6.2 實(shí)操后的幾點(diǎn)建議最后說幾點(diǎn)我反復(fù)踩坑之后總結(jié)出來的經(jīng)驗(yàn)。第一個是“先復(fù)現(xiàn)再創(chuàng)新”。拿到源碼后不要著急改結(jié)構(gòu)先用作者給的預(yù)訓(xùn)練權(quán)重和測試命令跑通確認(rèn)整個鏈路能出成果再開始替換模型、改損失函數(shù)否則出現(xiàn)問題根本分不清是原代碼的鍋還是你自己改的鍋。第二個是“每次只改一個變量”。圖像修復(fù)模型效果受多個因素影響如果一次改了三四個參數(shù)效果變差了你完全不知道是哪個參數(shù)導(dǎo)致的。我習(xí)慣用實(shí)驗(yàn)管理工具記錄每次訓(xùn)練的配置、log和輸出圖比如簡單地把每個實(shí)驗(yàn)放在單獨(dú)目錄命名帶上參數(shù)摘要。第三個是保存模型時把配置一起保存。只存state_dict有個問題過幾個月你自己都忘了當(dāng)時用了什么image_size、什么損失權(quán)重、什么歸一化方式。最好是存一個字典包含model_state_dict、config和源碼版本方便日后加載和新環(huán)境復(fù)現(xiàn)。在實(shí)際修復(fù)效果上別迷信單模型很多時候“檢測Mask 修復(fù) 后處理”的pipeline比單獨(dú)調(diào)模型更有效。先用分割模型自動生成Mask修復(fù)后再用銳化或顏色校正提升觀感。這個思路能大幅提高項(xiàng)目的落地價值也是我把這個源碼項(xiàng)目吃透之后最想分享的一點(diǎn)。本文還有配套的精品資源點(diǎn)擊獲取