U7.4 15 分鐘閱讀 2026年9月

U7.4 另一條路:分佈匹配

本篇重用M5.1用錯的機率表下注一年會多輸多少:KL Divergence 與 Cross-Entropy·M5.3只看得到樣本,怎麼調機器:對 Generator 微分·M1.4去噪就是取 Posterior Mean(Tweedie)

到目前為止,全部都是回歸

回頭看這兩個單元的每一個損失:progressive distillation 讓學生一步等於 teacher 兩步;consistency 讓 fθ(xt)f_\theta(x_t) 等於 fθ(xt)f_{\theta^-}(x_{t'});MeanFlow 讓 uθu_\theta 等於 identity 的右邊;Shortcut 讓大步等於兩小步的平均。形狀全都一樣:學生在某個輸入上的輸出,要等於某個目標點,用 squared loss 拉近。

這個形狀有一個 U1.3 就講過的性質:squared loss 的最小值是條件期望。目標若給定輸入是唯一的,回歸學到它;目標若給定輸入有好幾個可能,回歸學到它們的平均。前幾個單元這件事一直是幫手——正是靠它,條件目標才能在期望上等於邊際目標。但在蒸餾裡,它有另一面。

想像 teacher 從噪聲 x0x_0 出發,會走到哪一個資料點。若 teacher 是 ODE,這是一個確定的函數,目標唯一。但只要下面任何一件事發生,目標就變成多值:teacher 用的是 SDE 取樣器(同一個 x0x_0 走出不同終點);學生用的是條件配對 (x0,x1)(x_0,x_1) 而網路還沒學好;或者 consistency training 的時間網格太粗,目標退化成「xtx_t 對應的 x1x_1」而不是「軌跡的終點」。這時回歸給出的是幾個可能終點的平均——落在兩個 mode 中間、不像任何一個樣本的點U2.3 Q1 說邊際速度是交叉直線方向的平均,於是軌跡彎;同一件事在蒸餾裡,是輸出變成終點的平均,於是模糊。交叉問題在蒸餾裡換了一身裝。

有沒有辦法完全不要求逐點對應?

只要求分佈一樣

Distribution Matching Distillation(DMD,Yin et al. [1])的出發點是:我們其實不在乎學生把哪個 zz 送到哪個 xx,只在乎學生輸出的分佈 pθp_\theta 要等於 teacher 的分佈 prealp_{\text{real}}(teacher 多步取樣得到的分佈,實務上當作資料分佈)。學生是一個一步生成器 x=Gθ(z)x=G_\theta(z)zN(0,I)z\sim\mathcal N(0,I)。目標寫成 KL:

minθ KL(pθpreal)=Expθ[logpθ(x)logpreal(x)].\min_\theta\ \mathrm{KL}\big(p_\theta\,\|\,p_{\text{real}}\big)=\mathbb E_{x\sim p_\theta}\big[\log p_\theta(x)-\log p_{\text{real}}(x)\big].

直接算不出來——兩個密度都沒有。但它的梯度可以:

  θKL(pθpreal)=Ez[(sθ(x)sreal(x))x=Gθ(z)Gθ(z)θ],s=xlogp  \boxed{\;\nabla_\theta\,\mathrm{KL}(p_\theta\|p_{\text{real}})=\mathbb E_{z}\Big[\big(s_\theta(x)-s_{\text{real}}(x)\big)\Big|_{x=G_\theta(z)}\cdot\frac{\partial G_\theta(z)}{\partial\theta}\Big],\qquad s=\nabla_x\log p\;}

密度消失了,只剩兩個 score。這和 U1.0 Q1 的觀察是同一件事:取樣(這裡是「把樣本往正確的分佈推」)需要的是 score,不是密度。讀法:在學生的每個樣本 xx 上,梯度下降把 xxsrealsθs_{\text{real}}-s_\theta 的方向推——往 teacher 密度高的方向、離開學生自己密度高的方向。前者把樣本拉向真實 mode,後者防止所有樣本擠到同一個地方。

展開細節一維四行推導:為什麼 KL 的梯度是 score 差

x=Gθ(z)x=G_\theta(z),KL =Ez[logpθ(Gθ(z))logpreal(Gθ(z))]=\mathbb E_z\big[\log p_\theta(G_\theta(z))-\log p_{\text{real}}(G_\theta(z))\big]。對 θ\theta 微分,θ\theta 出現在兩個地方:logpθ\log p_\theta 的下標,以及 GθG_\theta 的輸出。

  1. GθG_\theta 那條路(chain rule):Ez[(xlogpθ(x)xlogpreal(x))x=Gθ(z)θGθ(z)]\mathbb E_z\big[(\partial_x\log p_\theta(x)-\partial_x\log p_{\text{real}}(x))\big|_{x=G_\theta(z)}\,\partial_\theta G_\theta(z)\big],這就是 score 差乘上生成器的 Jacobian。
  2. 下標那條路:Ez[θlogpθ(x)x=Gθ(z)]=Expθ[θlogpθ(x)]\mathbb E_z\big[\partial_\theta\log p_\theta(x)\big|_{x=G_\theta(z)}\big]=\mathbb E_{x\sim p_\theta}\big[\partial_\theta\log p_\theta(x)\big]
  3. Expθ[θlogpθ(x)]=θpθ(x)dx=θpθdx=θ1=0\mathbb E_{x\sim p_\theta}[\partial_\theta\log p_\theta(x)]=\int\partial_\theta p_\theta(x)\,dx=\partial_\theta\int p_\theta\,dx=\partial_\theta 1=0(score function 的期望為零)。
  4. 所以只剩第 1 條路:θKL=Ez[(sθsreal)θGθ]\nabla_\theta\mathrm{KL}=\mathbb E_z[(s_\theta-s_{\text{real}})\,\partial_\theta G_\theta]

符號提醒θKL\nabla_\theta\mathrm{KL} 的方向是 sθsreals_\theta-s_{\text{real}}梯度下降是往 srealsθs_{\text{real}}-s_\theta 走。檢查一維例子:pθ=N(θ,1)p_\theta=\mathcal N(\theta,1)preal=N(0,1)p_{\text{real}}=\mathcal N(0,1),KL =θ2/2=\theta^2/2,梯度 θ\theta;而 sθ(x)sreal(x)=(xθ)+x=θs_\theta(x)-s_{\text{real}}(x)=-(x-\theta)+x=\thetaθG=1\partial_\theta G=1,吻合。

sreals_{\text{real}} 是 teacher,用 U1.3 從 denoiser 換算;sθs_\theta 是學生分佈的 score,另外訓一個 U1.2 式的 denoiser,訓練資料是學生當下的輸出。

為什麼要在所有 noise level 上做

上面的式子寫在乾淨樣本 xx 上,實際做不到:pθp_\thetaprealp_{\text{real}} 在高維空間都很薄,兩者幾乎不重疊的地方 score 沒有意義,訓練初期學生的樣本全在 teacher 密度為零的區域,sreals_{\text{real}} 給不出有用的方向。解法和 U1.1 選擇加噪的理由一模一樣:把兩個分佈都加噪。對每個 noise level tt,令 xt=(1t)x0+txx_t=(1-t)\,x_0'+t\,x(FM 慣例;x0x_0' 是新抽的噪聲端樣本,和本單元其他地方的 x0x_0 一樣來自 p0p_0,加一撇是因為它與生成那一支用的 zz 無關),學生的擴散分佈 pθ,tp_{\theta,t} 與 teacher 的 preal,tp_{\text{real},t} 都變得光滑、處處有密度。DMD 的目標是這些 KL 的加權和:

LDMD=Et[w(t)KL(pθ,tpreal,t)],θLDMD=Et,z,ϵ[w(t)(sfake(xt,t)sreal(xt,t))xtθ].\mathcal L_{\text{DMD}}=\mathbb E_t\Big[w(t)\,\mathrm{KL}\big(p_{\theta,t}\,\|\,p_{\text{real},t}\big)\Big],\qquad \nabla_\theta\mathcal L_{\text{DMD}}=\mathbb E_{t,z,\epsilon}\Big[w(t)\big(s_{\text{fake}}(x_t,t)-s_{\text{real}}(x_t,t)\big)\frac{\partial x_t}{\partial\theta}\Big].

兩個 score 都是denoisersreal(xt,t)s_{\text{real}}(x_t,t) 由 teacher 的 denoiser 經 Tweedie 換算——這是 teacher 本來就會的事,任何 tt 都給得出來;sfake(xt,t)s_{\text{fake}}(x_t,t) 由一個「fake denoiser」給,它用 U1.2 那五行程式碼訓練,只是資料換成學生的輸出 Gθ(z)G_\theta(z)。因為學生一直在變,fake denoiser 要同步更新——交替訓練:一步更新生成器、一步(或幾步)更新 fake denoiser。這個「一個生成器、一個追著它的 score 網路」的結構最早出現在 ProlificDreamer 的 variational score distillation(VSD,Wang et al. [3]),那裡用來做 text-to-3D;DMD 把它搬到蒸餾上。

w(t)w(t)U2.4 那第四個旋鈕、訓練端的加權(U3.5 有它和 SNR 的換算)——DMD 用一個依樣本正規化的權重讓各 noise level 的梯度量級相當。DMD 原版還加了一個小的回歸項(用少量 teacher 的 (z,x)(z,x) 配對),DMD2 [2] 把它拿掉、讓 fake denoiser 更新得比生成器頻繁(two time-scale update)、加一個 GAN 判別器補品質,並把一步推廣到少步。主線只需要記住:DMD 的訓練訊號是兩個 score 的差,其中一個 score 要即時估計。

課堂提問Q1

回歸式蒸餾(consistency、MeanFlow、Shortcut……)與分佈匹配式蒸餾(DMD),各自的失敗模式是什麼? 請分別說出它壞掉時樣本長什麼樣、以及病根在哪一個數學性質上。

先想一想,再展開看整理後的答案

課堂上通常先說得出回歸的那一半(「會模糊」),分佈匹配的那一半需要提示「這個結構像不像 GAN」才會浮現。

回歸式:目標多值時平均到中間。 病根是 squared loss 的最小值是條件期望。teacher 映射若在某個輸入上多值——SDE teacher、條件配對、粗時間網格——學生輸出幾個終點的平均,樣本落在 mode 之間、細節被抹掉。這正是 U2.3 Q1 的交叉問題:那裡平均掉的是速度方向,這裡平均掉的是終點。就算網路容量無限、訓練無限久,只要目標本身多值,這個平均就不會消失;它是目標的性質,不是優化的失敗。反過來說,回歸式的優點也在這裡:訓練是普通的 supervised learning,穩定、可預期,不需要第二個網路。

分佈匹配式:估不準 score,或塌到少數 mode。 DMD 完全不要求逐點對應——同一個 zz 送到哪裡都可以,只要整體分佈對——所以交叉問題不存在,樣本不會落在 mode 之間。代價有三層。第一,sfakes_{\text{fake}} 是一個正在移動的目標的估計:fake denoiser 永遠在追學生,追不上時梯度方向就是錯的,訓練會震盪甚至發散——這是 GAN 裡判別器與生成器互相追逐的老問題換了一個形式。第二,KL(pθpreal)\mathrm{KL}(p_\theta\|p_{\text{real}}) 這個方向是 mode-seeking 的:它重罰「學生把質量放在 teacher 沒有的地方」,卻只輕罰「學生漏掉 teacher 的某個 mode」——學生把所有樣本放進 teacher 的一個 mode 裡,KL 可以很小。實務上表現為 mode collapse、多樣性下降,DMD2 加 GAN 項與 two time-scale update 主要就是在對付這兩件事。第三,DMD 需要 teacher(sreals_{\text{real}}),沒有從頭訓練的版本;學生也不是 flow map,多步取樣要另外設計。

一句話總結:回歸式的錯是「平均」,分佈匹配式的錯是「塌縮」或「追不上」。 前者來自 conditional expectation,後者來自要估一個移動中的 score。下一篇的表把這兩種誤差來源與曲率 bias、累積並列。

延伸閱讀:同一族的其他成員

分佈匹配不只 KL 一種寫法。Adversarial Diffusion Distillation(ADD,Sauer et al. [4])直接用 GAN 的判別器衡量學生與 teacher 的分佈差,再加一個 score distillation 項——把「分佈層面的訊號」交給判別器而不是 score 差。Inductive Moment Matching(IMM,Zhou et al. [5])用 moment matching(MMD)比較「從 tt 一步跳到 ss」與「先走到中間再跳到 ss」兩批樣本的分佈——它同時是分佈匹配式的(比較的是樣本分佈而非逐點),又用了 flow map 的半群結構,而且不需要 teacher。這兩篇不在主線裡展開,但它們說明了一件事:「回歸 vs 分佈匹配」是訓練訊號的分類,「學什麼物件」是另一個獨立的軸——IMM 學的是 flow map,訊號卻是分佈層面的。下一篇的表把這兩個軸分開列。

先消化一下

想一想

DMD 的梯度是 E[(sfakesreal)θxt]\mathbb E[(s_{\text{fake}}-s_{\text{real}})\,\partial_\theta x_t]。梯度下降把學生樣本往哪個方向推?

想一想

為什麼 DMD 要在所有 noise level tt 上匹配,而不是只在乾淨樣本上?

想一想

回歸式蒸餾「平均到中間」的失敗,即使網路容量無限、訓練無限久也不會消失。原因是:

想一想

DMD 最小化的是 KL(pθpreal)\mathrm{KL}(p_\theta\|p_{\text{real}}) 而不是 KL(prealpθ)\mathrm{KL}(p_{\text{real}}\|p_\theta)。這個方向的後果是:

參考文獻

  1. Yin, T., Gharbi, M., Zhang, R., Shechtman, E., Durand, F., Freeman, W. T., Park, T. One-step Diffusion with Distribution Matching Distillation. CVPR 2024.(KL 目標、score 差梯度、fake denoiser、回歸正則項。)
  2. Yin, T., Gharbi, M., Park, T., Zhang, R., Shechtman, E., Durand, F., Freeman, W. T. Improved Distribution Matching Distillation for Fast Image Synthesis. NeurIPS 2024.(DMD2:拿掉回歸項、two time-scale update、GAN 項、少步。)
  3. Wang, Z., Lu, C., Wang, Y., Bao, F., Li, C., Su, H., Zhu, J. ProlificDreamer: High-Fidelity and Diverse Text-to-3D Generation with Variational Score Distillation. NeurIPS 2023.(VSD:生成器加追蹤 score 網路的結構。)
  4. Sauer, A., Lorenz, D., Blattmann, A., Rombach, R. Adversarial Diffusion Distillation. 2023.(延伸閱讀。)
  5. Zhou, L., Ermon, S., Song, J. Inductive Moment Matching. 2025.(延伸閱讀。)