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