U6.1 12 分鐘閱讀 2026年9月

U6.1 Progressive Distillation:學生一步,等於老師兩步

本篇重用M2.3導航每 30 秒更新一次:Euler 法與它的誤差·M2.1從一片葉子到一群葉子:Flow Map 與 Pushforward

不要先造資料,邊走邊教

上一篇 Q1 的第一個「不好」是:每一筆訓練配對都要 teacher 跑完整條 ODE。Salimans & Ho [1] 的觀察是,這件事沒有必要一次做完。

想像 teacher 是一個要走 NN 步的取樣器(本單元座標下,從 t=Tt=T 走到 t=0t=0,每步是 U1.5 的 DDIM,也就是 PF-ODE 的一步 Euler)。我們不要求學生一口氣學會整條路,只要求:學生走一步,要到 teacher 走兩步的地方。 學好之後,學生是一個 N/2N/2 步的取樣器。然後把學生當成新的 teacher,再訓一個 N/4N/4 步的學生。重複 log2N\log_2 N 輪。

每一輪的訓練資料是即時算出來的:抽一個 x0x_0、抽一個 tt(在學生的時間網格上)、加噪聲得到 xtx_t,然後讓 teacher 從 xtx_t 走兩步得到 x~\tilde x。這只要兩次 teacher 的網路呼叫,不是幾十次。上一篇那個「要先用 teacher 生成整個資料集」的成本,被攤成每一輪每一筆兩次呼叫。

課堂提問Q1

每一輪的老師都是上一輪的學生。那為什麼不乾脆讓第一個學生直接學「teacher 的 NN 步」,一次到位? 反正最後要的就是那個一步的映射。

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

「一次到位」就是 U6.0 的直接回歸——那一篇已經分析過它,所以真正要問的是:逐輪減半到底換到了什麼?

換到的第一件事是資料成本。 直接回歸要 teacher 跑完整條 ODE 才生出一筆配對;PD 每一筆只要 teacher 走兩步。但這件事上面那一節已經說了。

換到的第二件事才是重點:每一輪要學的東西都只比上一輪難一點點。 第一輪的學生要學的是「teacher 的兩小步」——那兩小步幾乎共線,所以那個映射幾乎就是一條直線,網路很好表達。學好之後它自己成為一個 N/2N/2 步的取樣器,下一輪再把兩步合成一步,又只是「比上一輪的映射彎一點」。難度被攤在 log2N\log_2 N 輪上。直接回歸沒有這個階梯:第一天就要網路表達整條 ODE 疊起來的那個映射。

代價是什麼? 每一輪都在上一輪的輸出上訓練,所以上一輪的誤差進入這一輪的訓練資料——這和 U3.3 的 reflow 是同一種代價。下一節把它寫出來。

順帶一提,這也解釋了為什麼 PD 蒸到 4 步品質還好、蒸到 1 步明顯下降 [1]:不是最後一輪比較不認真,是最後那幾輪要跨的那一段本來就最難學

寫成式子

用本單元的 VE 座標 xt=x0+tϵx_t=x_0+t\epsilon。teacher 的一步 DDIM 從 ttt<tt'<t,寫成 denoiser x^0\hat x_0 的形式是

xt=xt+(tt)xtx^0(xt,t)t=ttxt+(1tt)x^0(xt,t),x_{t'}=x_t+(t'-t)\,\frac{x_t-\hat x_0(x_t,t)}{t} =\frac{t'}{t}\,x_t+\Big(1-\frac{t'}{t}\Big)\hat x_0(x_t,t),

(中間那個分數是 ϵ^=(xtx^0)/t\hat\epsilon=(x_t-\hat x_0)/t,這一步就是上一篇 PF-ODE dxdt=E[ϵxt]\frac{dx}{dt}=\mathbb E[\epsilon\mid x_t] 的一步 Euler。)讀法:往 tt' 走,就是把 xtx_t 與去噪結果 x^0\hat x_0t/tt'/t 做線性混合。

學生的時間網格是 teacher 的每隔一格:學生從 tt 一步到 t=t2Δt''=t-2\Delta,teacher 從 ttttΔt2Δt\to t-\Delta\to t-2\Deltax~\tilde x。我們想要學生的一步落在 x~\tilde x。學生的一步也是上面那個線性混合,只是用學生自己的 denoiser x^0θ\hat x_0^{\theta};要它落在 x~\tilde x,把上式反解,得到學生該預測的 x0x_0

  x0target=x~ttxt1tt  ,LPD=E[w(t)x^0θ(xt,t)x0target2].\boxed{\;x_0^{\text{target}}=\frac{\tilde x-\frac{t''}{t}\,x_t}{1-\frac{t''}{t}}\;},\qquad \mathcal L_{\text{PD}}=\mathbb E\Big[w(t)\,\big\|\hat x_0^{\theta}(x_t,t)-x_0^{\text{target}}\big\|^2\Big].

目標是確定的。 給定 (xt,t)(x_t,t)x~\tilde x 是 teacher 兩步 DDIM 的結果,是 xtx_t 的一個函數;x0targetx_0^{\text{target}} 也是。回歸的目標是一個點,不是一個分佈,所以 U1.3 的 conditional expectation 沒有東西可以平均——與上一篇 Q1 的分析相同,來源仍是 DDIM(ODE)的確定性。這是「蒸餾」與「從頭訓練」最大的差別:從頭訓練時目標是 x0x_0,一個 xtx_t 對應很多個 x0x_0,網路學到的是後驗平均;蒸餾時目標是 teacher 的輸出,一對一。

學生看過所有的 tt 訓練時 tt 在整個網格上抽,學生學到的是「從任何一格跳兩格」,不是只有 t=Tt=T。所以每一輪的學生都是一個合法的多步取樣器;上一篇的第二個「不好」(只看過 t=0t=0)在這裡被解掉一部分——但注意,學生只學會固定步長的跳法,不是任意步長。

學生從 teacher 初始化。 第一輪的學生就是 teacher 的權重複製一份;每一輪學的只是「把兩小步合成一大步」這個差量,訓練很快收斂。

展開細節為什麼把 target 寫在 x₀ 空間,以及 v-parametrization 從哪裡來

可以直接在 xtx_{t''} 空間算 loss(\|學生一步x~2-\tilde x\|^2),但學生的一步是 ttxt+(1tt)x^0θ\frac{t''}{t}x_t+(1-\frac{t''}{t})\hat x_0^\theta,兩者只差一個 (1t/t)(1-t''/t) 的縮放,所以等價於在 x0x_0 空間算 loss 再乘上 (1t/t)2(1-t''/t)^2 的權重。Salimans & Ho 選擇明確寫成 x0x_0 空間、再自己選權重 w(t)w(t)——他們建議的權重是 max(SNR(t),1)\max(\text{SNR}(t),1)(「truncated SNR」),這是 U3.5 那個「選 tt 的權重=選 log\logSNR 上的權重函數」在蒸餾裡的樣子。

另一個細節:當步數減到很少時,tt 很大的那一格 ϵ^=(xtx^0)/t\hat\epsilon=(x_t-\hat x_0)/t 幾乎不含資訊、而 x^0\hat x_0 又要從幾乎純噪聲裡直接猜整張圖,兩種 parametrization 各有一端不穩。他們因此提出 vv-prediction(U1.2 已經介紹過的那個座標),它在兩端都有界。這是 vv-parametrization 最早的出處。

誤差怎麼累積

每一輪,學生逼近的是 teacher 的兩步,不是真實的 PF-ODE。所以每一輪有兩種誤差進入學生:teacher 自己已經帶著的誤差,加上學生這一輪的逼近誤差。log2N\log_2N 輪之後,最終那個一步模型身上疊了每一輪的逼近誤差。

這和 U3.3 說的 reflow 代價是同一件事:「第 kk 輪的配對來自第 k1k-1 輪模型的 ODE 解,模型誤差與數值誤差一起進入下一輪的訓練資料」。兩者都是一連串蒸餾,每一輪的老師是上一輪的學生。差別是每一輪在學什麼——reflow 每輪學一個更直的速度場,PD 每輪學一個步長翻倍的跳躍。

還有一個更細的觀察。前幾輪(NN 很大、Δ\Delta 很小)幾乎沒有東西要學:teacher 的兩小步本來就接近一大步,因為那一段軌跡幾乎是直線。真正困難的是最後幾輪:從 4 步到 2 步、從 2 步到 1 步,這時一步要跨過整條軌跡最難學的中段。U3.2 的曲率在這裡以另一種形式回來——不是 Euler 的離散誤差,而是「這一大步的 flow map 有多難用網路表達」。實務上 PD 蒸餾到 4 步左右品質仍好,1 步明顯下降 [1]。

解掉了什麼、留下什麼

回到上一篇的三個「不好」:

  • 資料太貴:解掉了。不用事先生成資料集,每筆訓練樣本兩次 teacher 呼叫。
  • 只看過 t=0t=0:解掉一半。學生看過所有網格點,能多步取樣;但只會固定步長,而且最終一步模型還是得一次跨過整條軌跡。
  • 上限是 teacher,而且多輪累積:沒解,甚至更明顯了——每一輪都是一次蒸餾。

下一篇換一個角度:不要把「一步」定義成「從網格上某一格跳到下一格」,而是定義成「從任何一點跳到終點」。這會讓多輪消失、時間網格的角色改變,並且——這是最意外的——為「不要 teacher」打開一條路。

先消化一下

想一想

Progressive distillation 每一輪的訓練樣本 (xt, x0target)(x_t,\ x_0^{\text{target}}) 是怎麼來的?

想一想

PD 的回歸目標是確定的(給定 xtx_t 只有一個答案),這件事的來源是:

想一想

在 PD 的多輪流程裡,哪幾輪最困難、誤差最大?

想一想

PD 與多輪 reflow 的共同結構是:

參考文獻

  1. Salimans, T., Ho, J. Progressive Distillation for Fast Sampling of Diffusion Models. ICLR 2022.(逐輪減半、x0x_0 空間的目標、truncated SNR 權重、vv-parametrization。)
  2. Luhman, E., Luhman, T. Knowledge Distillation in Iterative Generative Models for Improved Sampling Speed. 2021.(上一篇的直接回歸法,PD 的對照。)
  3. Liu, X., Gong, C., Liu, Q. Flow Straight and Fast. ICLR 2023.(多輪 reflow 的誤差累積。)