註釋式擴散模型 – DDPM 實作的詳細步驟說明
TL;DR
Hugging Face 發布了一本完整且註釋過的 PyTorch notebook,實作了原始的去噪擴散概率模型(DDPM)——一個學習從純高斯噪聲中去噪圖像的神經網路——並展示了在低解析度資料集上的訓練以及高品質圖像的抽樣。此資源闡明了數學原理、網路架構、訓練迴圈與推論流程,讓擴散模型對實務工作者更易於使用。
擴散模型的功能(關鍵要點)
擴散模型學習一個逆向去噪過程,將隨機高斯噪聲轉換為真實資料;前向過程依照預先定義的時間表加入噪聲,而類 U‑Net 的神經網路在每個時間步預測加入的噪聲,從而在固定步數後完成圖像生成。
前向擴散:使用變異數時間表加入噪聲
結論:前向過程是一個閉式高斯轉移,可直接抽樣任意噪聲時間步,省去逐步加入噪聲的需求。
The forward diffusion distribution is
q(x_t | x_{t-1}) = N(x_t ; sqrt(1-β_t)·x_{t-1}, β_t·I)
where the variance schedule (β_1,…,β_T) is monotonic (e.g., linear, cosine, quadratic, or sigmoid). By repeatedly applying this transition, the marginal
q(x_t | x_0) = N(x_t ; sqrt(\bar α_t)·x_0, (1-\bar α_t)·I)
can be sampled directly using pre‑computed (\bar α_t = \prod_{s=1}^t (1-β_s)). This “nice property” lets the training algorithm draw a random timestep (t) and corrupt a real image (x_0) in one step.
逆向擴散:透過噪聲預測學習均值
結論:模型被訓練以預測在時間步 (t) 加入的精確高斯噪聲 (ε);預測的噪聲隨後用於計算逆向高斯分布的均值。
The reverse conditional is assumed Gaussian:
p_θ(x_{t-1} | x_t) = N(x_{t-1} ; μ_θ(x_t, t), σ_t^2·I)
DDPM fixes (σ_t^2) to the known forward variance and learns only (μ_θ). By re‑parameterising the mean in terms of the noise predictor (ε_θ):
μ_θ(x_t, t) = (1/√α_t)·(x_t - (β_t/√(1-\bar α_t))·ε_θ(x_t, t))
the training loss reduces to a simple mean‑squared error between the true noise (ε) and the network output (ε_θ):
L_t = || ε - ε_θ(x_t, t) ||^2
隨機在每個 batch 中抽樣 (t) 可得到變分下界的無偏估計。
網路架構:條件式 U‑Net
結論:DDPM 使用時間條件式 U‑Net,搭配正弦位置嵌入、殘差塊、注意力機制與群組正規化,以在每個空間解析度上預測噪聲。
主要組件:
- 正弦位置嵌入 編碼時間步 (t),並透過小型 MLP 注入至每個 ResNet 塊。
- 權重標準化卷積 結合群組正規化時可提升訓練穩定性。
- ResNet 塊(兩層 convolution‑norm‑SiLU,並可選擇性加入時間嵌入的 scale‑shift)提供主要特徵轉換。
- 注意力模組(可為完整多頭或線性注意力)捕捉長距離依賴。
- 群組正規化 在注意力之前應用(
PreNorm)。 - 下採樣/上採樣路徑 將空間解析度減半或加倍,同時保持通道深度,呼應經典 U‑Net 設計。
- 最終頭部 將串接的瓶頸特徵映射回圖像形狀。
完整的 PyTorch 類別 Unet 組合上述元件,接受噪聲圖像張量與時間步張量,並回傳形狀相同的噪聲張量。
訓練迴圈:隨機時間步抽樣與 Huber 損失
結論:訓練透過在每個 batch 中抽樣隨機時間步,使用閉式前向過程汙損輸入,並最小化真實噪聲與預測噪聲之間的 Huber 損失。
Pseudo‑code (simplified):
for epoch in range(num_epochs):
for batch in dataloader:
t = torch.randint(0, T, (batch_size,)).to(device) # uniform timestep
loss = p_losses(model, batch, t, loss_type='huber') # MSE/Huber on noise
loss.backward()
optimizer.step()
輔助函式 p_losses 會呼叫 q_sample 取得 (x_t),再計算與網路輸出 model(x_t, t) 的損失。定期抽樣(使用逆向過程)可視化訓練進度。
抽樣(推論):逆向擴散鏈
結論:生成從純高斯噪聲開始,並迭代套用學習到的去噪步驟;每一步使用預測的噪聲計算後驗均值,並依據已知變異數時間表加入校準的高斯噪聲。
Algorithm (Algorithm 2 in the DDPM paper):
x_T = torch.randn(shape) # start from noise
for t in reversed(range(T)):
pred_noise = model(x_t, t)
mean = (1/√α_t) * (x_t - β_t/√(1-\bar α_t) * pred_noise)
if t > 0:
x_{t-1} = mean + √posterior_variance_t * torch.randn_like(x_t)
else:
x_0 = mean
提供的 notebook 在 p_sample_loop 中實作此迴圈,並以 GIF 方式視覺化去噪軌跡。
Fashion‑MNIST 完整範例
結論:本教學在 28×28 的 Fashion‑MNIST 資料集上訓練 DDPM,於數千次訓練步驟後即可產生可辨識的服飾圖像。
- 資料載入使用 🤗
datasets函式庫,搭配即時轉換(隨機水平翻轉、縮放至 ([-1,1]))。 - 模型超參數:
dim=image_size、dim_mults=(1,2,4)、channels=1(灰階圖像)。 - 最佳化器:Adam,學習率 1e‑3。
- 訓練執行 6 個 epoch;損失快速下降至低於 0.05。
- 訓練後抽樣產生清晰的 T‑shirt 形狀,證實實作在低解析度資料上可行。
此實作在擴散文獻中的定位
結論:此註釋式實作重現了原始 DDPM(Ho 等,2020),並作為許多後續進展的基礎,如變異數學習(Nichol 等,2021)、階層式擴散(Ho 等,2021)、無分類器指導(classifier‑free guidance)以及大規模文字到圖像模型(DALL‑E 2、Imagen)。
部落格文章列出的主要後續工作:
- Improved DDPM – 同時學習均值與變異數,提升樣本品質。
- Cascaded Diffusion Models – 堆疊多個擴散模型以進行高解析度合成。
- Diffusion Models Beat GANs – 透過架構調整與分類器指導,展示更優的 FID 分數。
- Classifier‑Free Guidance – 在條件生成時免除外部分類器的需求。
- DALL‑E 2 & ImageGen – 結合擴散與 CLIP 嵌入或大型語言模型,以實現文字條件圖像合成。
主要缺點仍是需要大量的去噪步驟(通常 1000 步),儘管近期研究(例如高階求解器)已將步數降低至僅 10 步左右。
開發者的實務要點
- 此 notebook 提供即時可執行的參考實作;可更換資料集、提升
T,或將 U‑Net 替換為更大的骨幹網路以擴展規模。 - 調整變異數時間表(
linear_beta_schedule、cosine_beta_schedule等)以在樣本真實度與速度之間取得平衡。 - 使用 Huber 損失以提升魯棒性;程式碼亦支援 L1/L2。
- 定期抽樣(
save_and_sample_every)對於監測模式崩潰或訓練不穩定性至關重要。 - 若需更快推論,可考慮實作 DDIM 或隨機抽樣變體,或採用近期的「少步」求解器。
Sources
- OriginalThe Annotated Diffusion Model