現(xiàn)DDPM擴(kuò)散模型:從原理到源碼實(shí)戰(zhàn)詳解)
簡(jiǎn)介基于PyTorch實(shí)現(xiàn)的DDPM去噪擴(kuò)散概率模型圖像生成完整工程面向正在學(xué)習(xí)擴(kuò)散模型或需要參考生成式AI代碼的開發(fā)者解決了從零搭建訓(xùn)練與采樣流程的難題。壓縮包共11個(gè)文件包括6個(gè)Python腳本分別覆蓋數(shù)據(jù)集處理、UNet模型定義、前向擴(kuò)散模擬、訓(xùn)練、采樣及可視化、3張運(yùn)行效果圖、1個(gè)依賴清單和1個(gè)說明文檔整體大小4.47MB目錄結(jié)構(gòu)清晰便于檢索。目前已有158人學(xué)習(xí)。代碼按照標(biāo)準(zhǔn)DDPM流程組織既有前向加噪過程演示也有完整訓(xùn)練與反向采樣模塊可幫助讀者深入理解噪聲調(diào)度、UNet結(jié)構(gòu)及生成推理機(jī)制。無需額外復(fù)雜配置安裝依賴后即可獨(dú)立運(yùn)行適合用于復(fù)現(xiàn)實(shí)驗(yàn)、二次開發(fā)或作為畢業(yè)論文的參考實(shí)現(xiàn)。 正好手頭在寫一套基于PyTorch的DDPM圖像生成模型源碼最近也總有人在后臺(tái)問擴(kuò)散模型該怎么入門、源碼從哪讀起干脆把這個(gè)項(xiàng)目從頭到尾拆一遍講講原理、環(huán)境、結(jié)構(gòu)、訓(xùn)練和踩過的坑。這篇內(nèi)容適合剛接觸擴(kuò)散模型的學(xué)生也適合想從GAN切到擴(kuò)散模型的工程師我會(huì)盡量把數(shù)學(xué)部分講得通俗把實(shí)操步驟寫清楚確保你拿到這套源碼能直接跑通、能改、能調(diào)。1. DDPM項(xiàng)目整體概覽與核心思路1.1 DDPM是什么解決了什么問題DDPM全稱是Denoising Diffusion Probabilistic Models中文常稱為去噪擴(kuò)散概率模型。它屬于生成模型的一種核心思路非常直白先把一張真實(shí)圖片逐步加噪加到變成純高斯噪聲然后訓(xùn)練一個(gè)神經(jīng)網(wǎng)絡(luò)學(xué)習(xí)逆向過程從純?cè)肼暲镆徊讲交謴?fù)出原始圖片。你把它想象成往一杯清水里滴墨墨水越來越渾直到整杯水完全變黑然后教一個(gè)模型學(xué)會(huì)用吸管把墨滴往回吸最終恢復(fù)出干凈的水。這套思路最早由Ho等人于2020年提出一經(jīng)發(fā)布就在圖像生成質(zhì)量上直接對(duì)標(biāo)GAN且訓(xùn)練穩(wěn)定性遠(yuǎn)優(yōu)于GAN。這套源碼解決的問題很明確不同背景的開發(fā)者手頭沒有多年數(shù)學(xué)功底、沒有大規(guī)模算力也希望能從一個(gè)可直接運(yùn)行的PyTorch實(shí)現(xiàn)出發(fā)理解擴(kuò)散模型內(nèi)部到底發(fā)生了什么并能把模型遷移到自己的數(shù)據(jù)集上訓(xùn)練出不錯(cuò)的生成效果。項(xiàng)目不追求刷SOTA而是把DDPM最核心的部分拆清楚、寫干凈。1.2 為什么選擇PyTorch實(shí)現(xiàn)我面試過人也帶過做算法的小伙伴幾乎統(tǒng)一感受是PyTorch在動(dòng)態(tài)圖模式下調(diào)試體驗(yàn)極佳對(duì)初學(xué)者非常友好。你可以把網(wǎng)絡(luò)前向傳播的中間張量直接打印出來看維度對(duì)不對(duì)也可以隨時(shí)用斷點(diǎn)停下來檢查某一步的輸入輸出形狀這在排查擴(kuò)散模型這種多步迭代過程時(shí)尤其重要。相比之下靜態(tài)圖框架在調(diào)試“逐步加噪、逐步采樣”這類動(dòng)態(tài)流程時(shí)會(huì)讓你有種隔靴搔癢的感覺。此外PyTorch的生態(tài)對(duì)生成模型極其完備。HuggingFace Diffusers、torchvision等庫都提供了大量預(yù)訓(xùn)練權(quán)重和經(jīng)典實(shí)現(xiàn)方便我們對(duì)照驗(yàn)證。這套源碼本身就是純PyTorch實(shí)現(xiàn)不依賴額外的重型封裝庫只用了torch、torchvision、numpy和PIL基礎(chǔ)組件這樣任何人clone下來安裝好依賴就能跑不用被復(fù)雜的工程框架勸退。1.3 項(xiàng)目適用場(chǎng)景與學(xué)習(xí)價(jià)值如果你是想發(fā)論文的研究生這套源碼可以作為baseline在此基礎(chǔ)上改損失函數(shù)、改噪聲調(diào)度、改網(wǎng)絡(luò)結(jié)構(gòu)做對(duì)比實(shí)驗(yàn)會(huì)很方便如果你是工程師想在業(yè)務(wù)里做圖像生成、數(shù)據(jù)增強(qiáng)、風(fēng)格遷移這套代碼同樣能幫你在最短時(shí)間內(nèi)跑通DDPM流程后續(xù)可以直接替換成DDIM或者Latent Diffusion做加速如果你是本科生或者自學(xué)者那這個(gè)項(xiàng)目的價(jià)值就更大了因?yàn)樗恰耙恍幸恍心茏x懂”的代碼不是工業(yè)級(jí)黑盒配合這篇文章里的解析完全可以理解生成模型的核心技術(shù)點(diǎn)。2. PyTorch環(huán)境搭建與依賴準(zhǔn)備2.1 基礎(chǔ)環(huán)境配置建議跑DDPM這套代碼說難不難但環(huán)境沒配對(duì)后面全是淚。我自己最開始在Windows上裝PyTorch癡迷于追求最新版CUDA結(jié)果跟顯卡驅(qū)動(dòng)版本不匹配跑卷積直接報(bào)錯(cuò)CUDA error: no kernel image is available。后來我總結(jié)了一套相對(duì)穩(wěn)妥的搭配方案。首先確認(rèn)顯卡驅(qū)動(dòng)版本在命令行輸入nvidia-smi查看頂部CUDA Version比如顯示12.1那么安裝CUDA 12.1及以下的PyTorch都沒有問題。Python版本我建議選擇3.9或3.10兼容性最好太新的3.12、3.13反而容易出現(xiàn)某些依賴包沒編譯好。如果你有Anaconda我推薦用以下方式創(chuàng)建獨(dú)立環(huán)境conda create -n ddpm python3.9 conda activate ddpm2.2 PyTorch安裝與關(guān)鍵依賴版本激活環(huán)境后最關(guān)鍵的一步就是安裝PyTorch。這里我不建議直接用pip install torch因?yàn)槟J(rèn)識(shí)別的CUDA版本很可能不是你機(jī)器的版本。要到PyTorch官網(wǎng)用生成的命令安裝。以CUDA 12.1為例pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121這套源碼用到的基礎(chǔ)依賴很少如果只是訓(xùn)練MNIST和CIFAR-10這種小數(shù)據(jù)集都不需要額外裝太多東西。我在實(shí)際測(cè)試中使用的版本組合如下依賴庫推薦版本說明Python3.9兼容性極佳PyTorch2.0.11.13也可運(yùn)行torchvision0.15.1用于加載數(shù)據(jù)集numpy1.24.3數(shù)學(xué)運(yùn)算Pillow10.0.0圖像處理后處理matplotlib3.7.1可視化訓(xùn)練曲線2.3 數(shù)據(jù)集準(zhǔn)備與預(yù)處理細(xì)節(jié)這套源碼數(shù)據(jù)集加載支持兩種方式一種是通過torchvision直接下載MNIST、FashionMNIST、CIFAR-10另一種是從本地文件夾讀取自定義圖片數(shù)據(jù)集。如果你用自定義數(shù)據(jù)集圖片建議統(tǒng)一resize到64×64或128×128DDPM對(duì)分辨率比較敏感因?yàn)閁-Net下采樣次數(shù)是固定的如果輸入尺寸不是2的冪次倍數(shù)最后維度對(duì)不上會(huì)直接報(bào)錯(cuò)。我在實(shí)踐中的預(yù)處理寫法是先轉(zhuǎn)成RGB三通道再縮放到目標(biāo)分辨率最后歸一化到[-1, 1]區(qū)間。這里特別提醒DDPM訓(xùn)練時(shí)加噪是在[-1, 1]數(shù)據(jù)空間上操作的如果讓網(wǎng)絡(luò)在[0, 1]區(qū)間里學(xué)訓(xùn)練很容易不平穩(wěn)采樣出來的圖也會(huì)偏灰整體對(duì)比度發(fā)悶。3. 源碼核心模塊拆解3.1 擴(kuò)散過程的數(shù)學(xué)實(shí)現(xiàn)源碼里擴(kuò)散過程的核心在noise_scheduler.py文件。它實(shí)現(xiàn)了一個(gè)線性噪聲調(diào)度器linear beta schedule這段代碼是你理解DDPM的第一道門。前向過程在數(shù)學(xué)上用公式表示對(duì)輸入圖片x0在任意時(shí)間步t直接得到加噪后的圖片xt sqrt(α_bar_t) * x0 sqrt(1 - α_bar_t) * ε其中ε是標(biāo)準(zhǔn)高斯噪聲。源碼里的實(shí)現(xiàn)分兩個(gè)階段先預(yù)定義beta從0.0001線性增加到0.02然后通過累乘操作計(jì)算出alphas_cumprod這個(gè)變量是后續(xù)計(jì)算任意時(shí)刻加噪圖片和損失函數(shù)的關(guān)鍵。這里有個(gè)很容易踩的坑計(jì)算中間變量時(shí)要用32位浮點(diǎn)數(shù)如果用64位部分GPU算子不支持訓(xùn)練時(shí)會(huì)拖慢速度如果用16位誤差又會(huì)累積采樣效果會(huì)變差。3.2 U-Net模型結(jié)構(gòu)解析源碼里的模型文件是unet.py采用標(biāo)準(zhǔn)的U-Net結(jié)構(gòu)包含編碼器、解碼器和跳躍連接。編碼器部分是由多個(gè)下采樣塊組成每一層通過卷積提取特征然后逐步降低空間分辨率、增加通道數(shù)從64通道一路升到256通道解碼器部分通過轉(zhuǎn)置卷積逐步恢復(fù)分辨率通道數(shù)逐層降低最終的輸出通道數(shù)和輸入保持一致RGB圖像就是3通道。為什么要用U-Net而不是簡(jiǎn)單的卷積網(wǎng)絡(luò)因?yàn)閿U(kuò)散模型的輸入和輸出是同一尺寸的圖片屬于稠密預(yù)測(cè)任務(wù)需要同時(shí)保留全局語義信息和局部細(xì)節(jié)信息。跳躍連接的作用就是把編碼器各層的特征直接拼接到解碼器對(duì)應(yīng)層這樣模型在生成時(shí)既能參考低層級(jí)紋理細(xì)節(jié)又能參考高層級(jí)語義信息。如果去掉跳躍連接生成出來的圖基本是糊的結(jié)構(gòu)完全崩掉。3.3 訓(xùn)練循環(huán)與損失計(jì)算訓(xùn)練入口在train.py核心邏輯非常簡(jiǎn)潔。每次隨機(jī)采樣一批真實(shí)圖片隨機(jī)從0到T-1采樣時(shí)間步t利用之前提到的alphas_cumprod參數(shù)直接得到xt然后讓模型預(yù)測(cè)加入的噪聲ε最后計(jì)算預(yù)測(cè)噪聲與真實(shí)噪聲之間的均方誤差MSE Loss。這里模型不直接預(yù)測(cè)圖像本身而是預(yù)測(cè)噪聲這是DDPM的精髓所在。預(yù)測(cè)噪聲而不是預(yù)測(cè)圖像有什么好處我個(gè)人的理解是預(yù)測(cè)噪聲的優(yōu)化空間更加平滑。圖像本身是高維復(fù)雜信號(hào)直接預(yù)測(cè)圖像會(huì)讓模型像“瞎子摸象”每張圖收斂方向都不一樣而噪聲是一個(gè)相對(duì)簡(jiǎn)單的連續(xù)目標(biāo)每個(gè)像素的誤差獨(dú)立優(yōu)化起來更穩(wěn)定。源碼里的寫法遵循了這一設(shè)定確保每一步反向傳播都直接對(duì)應(yīng)去噪質(zhì)量的提升。3.4 采樣與生成過程的實(shí)現(xiàn)采樣階段的代碼在sample.py里。訓(xùn)練完成后輸入一個(gè)隨機(jī)高斯噪聲逐步執(zhí)行T次去噪迭代。每一步根據(jù)模型預(yù)測(cè)的噪聲通過公式計(jì)算前一步的均值再加上一個(gè)帶有方差控制的隨機(jī)噪聲項(xiàng)。在采樣過程中方差控制是由噪聲調(diào)度器給定的具體計(jì)算時(shí)要注意保留scheduler內(nèi)部狀態(tài)的一致性否則每隔幾步生成結(jié)果會(huì)出現(xiàn)色偏。這里分享一個(gè)個(gè)人經(jīng)驗(yàn)采樣前先用固定的隨機(jī)種子生成噪聲觀察幾次結(jié)果穩(wěn)定性如果仍然有明顯隨機(jī)性偏差問題多半出在方差參數(shù)計(jì)算上。源碼里提供了快速采樣的參數(shù)配置可以將采樣步數(shù)從1000降到200步視覺質(zhì)量損失不算太大適合快速驗(yàn)證生成效果。4. 訓(xùn)練實(shí)操與參數(shù)調(diào)優(yōu)4.1 訓(xùn)練腳本運(yùn)行完整流程環(huán)境配好以后訓(xùn)練運(yùn)行起來很簡(jiǎn)單。先克隆或解壓源碼包進(jìn)入項(xiàng)目目錄直接執(zhí)行python train.py --dataset mnist --epochs 100 --batch_size 64 --image_size 32 --device cuda我第一次跑這個(gè)命令時(shí)MNIST數(shù)據(jù)集會(huì)自動(dòng)下載到本地大約10個(gè)epoch之后就能看到比較清晰的數(shù)字輪廓。訓(xùn)練過程中源碼會(huì)定期把生成的圖片保存到samples目錄方便直觀觀察訓(xùn)練進(jìn)展。如果你想在自定義數(shù)據(jù)集上訓(xùn)練執(zhí)行命令調(diào)整data_dir路徑即可。模型訓(xùn)練完后權(quán)重會(huì)默認(rèn)保存在checkpoints/model_epoch100.pth之后采樣直接執(zhí)行python sample.py --model_path checkpoints/model_epoch100.pth --num_samples 32 --device cuda4.2 關(guān)鍵超參數(shù)選擇與調(diào)優(yōu)心得這套源碼里幾個(gè)關(guān)鍵超參數(shù)直接影響生成質(zhì)量我分別測(cè)試過把心得匯總成一張表超參數(shù)推薦值經(jīng)驗(yàn)說明T擴(kuò)散步數(shù)1000太少會(huì)學(xué)不穩(wěn)如200太多訓(xùn)練和采樣耗時(shí)翻倍beta起始/終止0.0001 / 0.02線性調(diào)度器默認(rèn)值穩(wěn)定可靠學(xué)習(xí)率2e-4 / 1e-4Adam優(yōu)化器適合2e-4過大會(huì)崩過小收斂極慢批量大小32 / 64CPU訓(xùn)練建議16-32GPU 64以上圖片尺寸32 / 64越小訓(xùn)練越快64以上更接近真實(shí)場(chǎng)景通道數(shù)64起步每翻倍下采樣通道翻倍參數(shù)量可控調(diào)參最大的坑在于學(xué)習(xí)率過高。我測(cè)試過直接把學(xué)習(xí)率提到1e-3前10個(gè)epoch的loss下降飛快但到30個(gè)epoch就開始震蕩最終生成的圖像有嚴(yán)重的棋盤格偽影怎么都消不掉。后來回退到2e-4重新訓(xùn)練效果立刻穩(wěn)定。所以遇到生成效果炸裂先別急著改模型結(jié)構(gòu)把學(xué)習(xí)率降下來試試。4.3 訓(xùn)練效果評(píng)估與可視化盲訓(xùn)練不可取。源碼訓(xùn)練時(shí)每200個(gè)iteration會(huì)打印一組當(dāng)前l(fā)oss同時(shí)會(huì)把最新生成的樣本圖保存下來。我建議訓(xùn)練過程中全程盯著生成圖質(zhì)量而不是只看loss數(shù)值。因?yàn)閘oss可能一直在降但圖像可能在模糊和輕微噪點(diǎn)之間反復(fù)橫跳這在擴(kuò)散模型里非常常見很可能是模型容量不足或訓(xùn)練步數(shù)不夠?qū)е碌?。有一個(gè)實(shí)用的評(píng)估方式是把同一組固定噪聲在訓(xùn)練的不同階段都生成一遍。比如第10、50、100個(gè)epoch用同一個(gè)隨機(jī)種子采樣這樣能非常直觀看到模型逐步“學(xué)會(huì)”生成圖像的細(xì)節(jié)增強(qiáng)對(duì)訓(xùn)練進(jìn)度的掌控感。源碼里如果沒有現(xiàn)成實(shí)現(xiàn)我建議你在sample.py中加一行隨機(jī)種子固定的邏輯幾行代碼就能搞定收益很明顯。5. 常見問題與排查技巧5.1 訓(xùn)練loss不下降或者變成NaNloss不下降或變NaN絕大多數(shù)情況下是數(shù)據(jù)預(yù)處理出了問題。首先檢查輸入圖片是否已經(jīng)歸一化到[-1, 1]如果沒有模型輸入分布和加噪分布錯(cuò)位梯度會(huì)異常其次檢查批量大小和通道數(shù)是否匹配尤其是自定義數(shù)據(jù)集時(shí)單通道灰度圖和三通道RGB圖混用會(huì)在網(wǎng)絡(luò)中間某個(gè)卷積層維度爆炸。把這兩個(gè)問題排查完95%的loss異常都能解決。如果loss一開始正常、訓(xùn)練到一半突然變NaN這時(shí)候大概率是數(shù)值精度問題。建議檢查是否手動(dòng)啟用了混合精度訓(xùn)練且未設(shè)置合理的loss縮放DDPM里加噪過程中涉及多個(gè)連乘操作如果使用fp16非常容易上溢或下溢所以最好先全用fp32訓(xùn)練跑通后再優(yōu)化加速。5.2 顯存不足時(shí)的解決方案顯存不足是訓(xùn)練擴(kuò)散模型的標(biāo)配問題。我自己的顯卡是8G顯存訓(xùn)練64×64圖片、batch_size設(shè)為64時(shí)會(huì)直接OOM解決辦法有三個(gè)層級(jí)第一降低batch_size到16或8觀察顯存變化第二減小圖像分辨率到32×32第三使用梯度累積策略模擬大的batch_size等價(jià)于每4步或8步更新一次梯度效果接近直接加大batch但對(duì)顯存占用幾乎無影響。源碼中我加了一個(gè)--grad_accum_steps參數(shù)默認(rèn)值為1當(dāng)用戶傳到4時(shí)會(huì)在反向傳播時(shí)不立即更新參數(shù)累計(jì)梯度后再執(zhí)行優(yōu)化器這種改動(dòng)對(duì)最終訓(xùn)練效果影響很小但能讓你在有限顯存下繼續(xù)訓(xùn)練。這個(gè)技巧我在多個(gè)生成模型實(shí)戰(zhàn)里都能用到值得掌握。5.3 采樣結(jié)果模糊或者出現(xiàn)結(jié)構(gòu)崩壞如果訓(xùn)練出來采樣圖像整體發(fā)灰、輪廓模糊先檢查采樣公式是否漏乘了均值系數(shù)如果只是細(xì)節(jié)崩壞看一下U-Net的注意力機(jī)制是否被誤改。很多人喜歡在U-Net里加入額外的注意力模塊來增強(qiáng)生成效果但如果維度處理不當(dāng)反而會(huì)破壞原本穩(wěn)定的語義。模型接收Batch×Channel×Height×Width輸入一旦Height和Width被某種池化改變跳躍連接拼接時(shí)就會(huì)發(fā)生維度不匹配雖然在代碼里沒報(bào)錯(cuò)但特征分布已經(jīng)亂了。實(shí)際訓(xùn)練中我用128×128人臉數(shù)據(jù)集做驗(yàn)證時(shí)出現(xiàn)過一個(gè)非常奇怪的紗窗效應(yīng)整張圖看起來有人臉輪廓但皮膚區(qū)域全是規(guī)則的細(xì)碎網(wǎng)格。原因是圖片resize時(shí)用了簡(jiǎn)單的最近鄰插值高頻紋理信息在縮放時(shí)丟失擴(kuò)散模型學(xué)到的“地面真值”本身就是破碎的這屬于數(shù)據(jù)問題換雙線性插值后立即緩解。5.4 采樣速度太慢的工程化處理方案原始DDPM要迭代1000次才能生成一張圖片在普通顯卡上可能需要幾十秒這在業(yè)務(wù)場(chǎng)景中很難接受。我自己踩過最快的方案是走DDIM采樣器在noise_scheduler里增加一個(gè)skip參數(shù)讓采樣時(shí)每隔10步做一次去噪處理。這樣總采樣步數(shù)從1000降到100生成一張圖只需兩三秒質(zhì)量幾乎不掉檔。源碼中我已經(jīng)預(yù)留了接口調(diào)整--sampling_steps為更小值即可。如果你還想進(jìn)一步壓縮可以使用Latent Diffusion的思路——先用VAE把圖像編碼到低維潛空間在潛空間里做擴(kuò)散然后再解碼回像素空間。但這一步已經(jīng)超出當(dāng)前源碼范圍需要引入額外的自編碼器模型建議先跑通當(dāng)前版本穩(wěn)定出圖后再做這種架構(gòu)升級(jí)。實(shí)操總結(jié)與補(bǔ)充心得跑完整個(gè)DDPM項(xiàng)目我最真實(shí)的體會(huì)是模型本身并不復(fù)雜真正需要花時(shí)間的是理解噪聲調(diào)度器、訓(xùn)練目標(biāo)函數(shù)和采樣循環(huán)這三者之間的配合關(guān)系。一個(gè)常見的誤區(qū)是想一次性把所有最新技術(shù)全堆進(jìn)去比如把Unet改成注意力機(jī)制版本把DDPM換成DDIM把損失函數(shù)改成LPIPS感知損失結(jié)果一跑就崩根本不知道問題出在哪一步。正確做法是先老老實(shí)實(shí)跑通原始DDPM看清每個(gè)環(huán)節(jié)的輸入輸出和張量形狀然后再逐步附加改進(jìn)這樣排查bug有據(jù)可循實(shí)驗(yàn)對(duì)比也說得清。最后再分享一個(gè)小技巧訓(xùn)練過程中如果發(fā)現(xiàn)loss曲線下降平緩、生成圖像卻長時(shí)間沒有明顯進(jìn)步可以試著調(diào)整噪聲調(diào)度器的噪聲強(qiáng)度上下界。把beta_end從0.02稍微減小到0.015會(huì)讓模型更專注學(xué)習(xí)高頻細(xì)節(jié)反之如果想生成更多樣化的圖片就把beta_end調(diào)大到0.03。這種微調(diào)不會(huì)帶來劇烈的訓(xùn)練崩壞往往會(huì)給結(jié)果帶來意想不到的改善。本文還有配套的精品資源點(diǎn)擊獲取