註釋式擴散模型 – 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_sizedim_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_schedulecosine_beta_schedule 等)以在樣本真實度與速度之間取得平衡。
  • 使用 Huber 損失以提升魯棒性;程式碼亦支援 L1/L2。
  • 定期抽樣(save_and_sample_every)對於監測模式崩潰或訓練不穩定性至關重要。
  • 若需更快推論,可考慮實作 DDIM 或隨機抽樣變體,或採用近期的「少步」求解器。

Sources