U6.4 20 分鐘閱讀 2026年9月

U6.4 讓 CT 訓得起來:iCT 與連續時間的 sCM

本篇重用M0.3漂流的溫度計:全導數、微分穿過積分與 JVP·M5.2追一隻自己也在跑的狗:移動目標、EMA 與 Stop-Gradient·M2.3導航每 30 秒更新一次:Euler 法與它的誤差

訓練不穩,是哪三件事在作用

上一篇結尾留下的圖像是:CT 的目標是 CD 目標的條件版本,代價是 bias 與 variance 的取捨。把這個單元到目前為止碰過的所有誤差整理成三類,接下來每一項技術都會被放進其中一格:

  • boundaryf(x,ε)=xf(x,\varepsilon)=x 有沒有精確成立、參數化在兩端是否有界。
  • bias:目標的期望偏離真實終點——來自 Euler 對彎曲軌跡的 O(Δt2)O(\Delta t^2) 誤差、fθf_{\theta^-} 非線性的 Jensen gap,以及任何讓「目標」與「真值」系統性不同的東西。
  • variance:目標本身的隨機散佈——條件目標退向自己的 ϵ\epsilon、極端樣本、以及 Δt\Delta t 小時訊噪比變差。

這一篇不引入新的損失,新的物件只有一個:沿軌跡的全導數。前半是離散時間的修補(iCT [1]),後半是把 Δt\Delta t 推到零(連續時間 CM,sCM [2])。

課堂提問Q1

下一節的第一個改動是把 target 網路的 EMA 拿掉。這和一般的直覺相反:目標會動的訓練(bootstrapping)不是都要靠 EMA 才不發散嗎?

拿掉之後,fθf_\theta 為什麼不會直接崩掉?

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

關鍵是這裡的「目標會動」和 RL 那種不一樣

RL 裡 target network 要 EMA,是因為目標值 r+γmaxaQ(s,a)r+\gamma\max_a Q(s',a) 本身就是網路輸出,沒有任何外部的錨;目標追著自己跑,所以要把它拖慢。

CT 的目標是 fθ(xtn,tn)f_{\theta^-}(x_{t_n},t_n),看起來也是自己。但它被兩件事釘住了:

  1. 邊界條件寫在架構裡。 U6.2 那一節說過,fθ(x,ε)=xf_\theta(x,\varepsilon)=x 對任何一組參數都成立。所以最靠近乾淨端的那一格是正確答案,不管網路好不好。訊號是從那裡往外傳的,不是憑空自舉。
  2. tn<tn+1t_n<t_{n+1},目標永遠比輸入更靠近乾淨端。 每一對都是「已經對的那一側教還沒對的那一側」,方向是單向的,沒有兩邊互相追的迴路。

所以 CT 不需要靠 EMA 去穩住一個沒有錨的自舉。而 EMA 反過來會壞事:上一篇那條「CT 與 CD 損失只差 o(Δt)o(\Delta t)」的定理,條件是 θ=θ\theta^-=\thetaθ\theta^- 一旦是滯後的 EMA,極限就換成另一個目標——fθf_\theta 收斂到的不再是 consistency function。這是一個 bias,而且是那種「怎麼訓都到不了」的 bias,不是訓久一點就好。

換成 θ=stopgrad(θ)\theta^-=\text{stopgrad}(\theta) 之後,目標是當前參數(只是不回傳梯度),定理的條件回來了。那為什麼還要 stop-gradient?因為要梯度的話,損失就變成在最小化「fθf_\theta 在相鄰兩點的差」,而那正是上一題那個退化解喜歡的東西。

一句話:EMA 在這裡壓的不是不穩定,是收斂的位置。

iCT:五個改動,各對一格

Song & Dhariwal [1] 把 2023 年的 CT 從「勉強能訓、品質明顯輸 CD」改到「不用 teacher 也能追上或超過 CD」。改動有五個,每一個都能對回上面的三格。

一、拿掉 target 的 EMA(bias)。 原版用 θ=EMA(θ)\theta^-=\text{EMA}(\theta) 當目標網路。直覺上 EMA 是在壓 variance——讓目標平滑一點。但 iCT 證明它其實引入 bias:上一篇說 CT 損失與 CD 損失在 Δt0\Delta t\to0 時只差 o(Δt)o(\Delta t),那條定理的條件是 θ=θ\theta^-=\theta;一旦 θ\theta^- 是一個滯後的 EMA,極限就變成另一個目標,fθf_\theta 不再收斂到 consistency function。改成 θ=stopgrad(θ)\theta^-=\text{stopgrad}(\theta)——目標是當前參數、只是不回傳梯度——後品質大幅提升。這一條值得記住,因為它反直覺:一個看起來是穩定器的東西,在這裡是 bias 的來源。(U6.2 說的「目標側要 stop-gradient」仍然成立,只是不再另加 EMA。)

二、Pseudo-Huber 取代 2\ell_2(variance)。

d(x,y)=xy2+c2c,c=0.00054Dd(x,y)=\sqrt{\|x-y\|^2+c^2}-c,\qquad c=0.00054\sqrt{D}

DD 是資料維度)。它在 xyc\|x-y\|\ll c 時像 2\ell_2、在 xyc\|x-y\|\gg c 時像 1\ell_1,梯度的模有上界。CT 的目標帶著 O(Δt)O(\Delta t) 的隨機退向,偶爾會有離群的樣本把 2\ell_2 的梯度拉得很大;pseudo-Huber 把這些離群值的影響壓下來。原版用的 LPIPS 則被拿掉——它讓 fθf_\theta 學到評估指標的特徵、對 FID 有系統性的不公平影響,是另一種 bias。

三、時間權重 λ(tn)=1tn+1tn\lambda(t_n)=\frac{1}{t_{n+1}-t_n}(variance 的均衡)。 相鄰兩點輸出的差距是 O(Δt)O(\Delta t);用 EDM 的時間網格 [4](tt 小處密、tt 大處疏)時,不同 nnΔt\Delta t 差好幾個數量級,loss 的尺度也跟著差好幾個數量級。乘上 1/Δt1/\Delta t 把每一對的貢獻拉到同一尺度,等效於tt 小的區域加重——那裡 xtx_t 接近資料、細節在那裡決定。

四、Lognormal 時間取樣(把預算放在有訊號的地方)。 不再均勻抽 nn,而是讓 logtn\log t_n 近似服從 N(Pmean,Pstd2)\mathcal N(P_{\text{mean}},P_{\text{std}}^2)Pmean=1.1P_{\text{mean}}=-1.1Pstd=2.0P_{\text{std}}=2.0。這就是 U2.4 那第四個旋鈕——訓練端的加權(U3.5 把它換算成 logSNR\log\text{SNR} 上的權重):兩端的一致性條件太容易(tt 極小時 ff\approx 恆等、tt 極大時 xtx_t 幾乎沒資訊),中段才是 ff 真正在變的地方,訊號在那裡。

五、NN 的指數 curriculum(bias 隨時間換 variance)。 原版的 NN 線性地從 2 增加到 150;iCT 改成

N(k)=min(s02k/K, s1)+1,s0=10, s1=1280,N(k)=\min\big(s_0\,2^{\lfloor k/K'\rfloor},\ s_1\big)+1,\qquad s_0=10,\ s_1=1280,

kk 是訓練步、KK' 讓翻倍的次數剛好填滿整段訓練。這是上一篇那個取捨的排程:一開始 NN 小、bias 大但訊號強;每翻倍一次就把 bias 砍一半,而 fθf_\theta 已經接近、variance 造成的擾動可承受。終點 N=1280N=1280 比原版大得多,因為前面四項把 variance 壓下來了,才容得下這麼小的 Δt\Delta t

展開細節其他兩個小改動:Fourier 尺度與 dropout

iCT 還把時間 embedding 的 Fourier 尺度從 16 調到 0.02,並提高 dropout。前者的理由與下半篇的 sCM 相通:CT 比的是兩個相鄰時刻的輸出,fθf_\thetatt 越敏感,這個差就越被 tt 方向的高頻擾動污染;把 Fourier 尺度調小等於強迫 fθf_\thetatt 平滑。這一條歸在 variance。(具體數值以原文為準。)

Δt\Delta t 推到零

iCT 的五項全部是在離散時間網格上和 bias–variance 周旋。有一個更直接的想法:既然 bias 來自 Δt\Delta t 不為零,那就讓 Δt0\Delta t\to0,看損失變成什麼。

自我一致性說 f(xt,t)f(x_t,t) 沿軌跡不變。「不變」在微分的語言裡就是沿軌跡的全導數為零

  ddtf(xt,t)=tf(xt,t)+xf(xt,t)dxtdt=0  沿每一條 PF-ODE 軌跡.\boxed{\;\frac{d}{dt}f(x_t,t)=\partial_t f(x_t,t)+\nabla_x f(x_t,t)\cdot\frac{dx_t}{dt}=0\;}\qquad\text{沿每一條 PF-ODE 軌跡}.

這是這一篇唯一的新物件。逐項讀:第一項是「時間往前走、位置不動」時 ff 的變化;第二項是「位置隨軌跡移動」帶來的變化,dxtdt\frac{dx_t}{dt} 是 PF-ODE 的速度(本單元座標下 =E[ϵxt]=\mathbb E[\epsilon\mid x_t])。兩項要恰好相消U3.2 裡出現過同一個算子——那裡是對速度場 uutu+(u)u\partial_t u+(u\cdot\nabla)u 得到加速度;這裡是對 ff 取,得到「ff 沿軌跡的漂移」。

Song et al. [3] 已經推出:2\ell_2 版的 CM 損失在 NN\to\infty 時,梯度趨近

θ E[fθ(xt,t) ⁣ ddtfθ(xt,t)]\nabla_\theta\ \mathbb E\Big[f_\theta(x_t,t)^{\!\top}\ \frac{d}{dt}f_{\theta^-}(x_t,t)\Big]

θ\theta^- 側的全導數不回傳梯度)。把它推到零,就是要 fθf_\theta 沿軌跡的全導數為零。這個目標沒有 Δt\Delta t:Euler 的曲率誤差、Jensen gap,兩個 O(Δt2)O(\Delta t^2) 的 bias 來源在極限裡一起消失。

那要怎麼在沒有 teacher 的情形下算 dxtdt\frac{dx_t}{dt}?和上一篇一樣,用 U1.2conditional trickxt=x0+tϵx_t=x_0+t\epsilontt 微分,條件速度是 ϵ\epsilon 本身。而 ddtf=tf+xfx˙t\frac{d}{dt}f=\partial_tf+\nabla_xf\cdot\dot x_t x˙t\dot x_t 是線性的,所以用 ϵ\epsilon 代替 E[ϵxt]\mathbb E[\epsilon\mid x_t] 後取期望是精確的——連上一篇那個 Jensen gap 都沒有了,因為現在條件量不再被包在非線性函數裡面,而是線性地進入一個乘積。連續時間 CT 在這裡又多賺一項。

計算上,tf+xfv\partial_tf+\nabla_xf\cdot vff(x,t)(x,t) 沿方向 (v,1)(v,1) 的方向導數,一次 Jacobian-vector product(JVP)就得到,不用算整個 Jacobian:torch.func.jvp(f, (x, t), (v, ones))。成本約等於一次前向傳播。

展開細節全導數的 JVP:三行程式與為什麼不算 Jacobian

xf\nabla_xfD×DD\times D 的矩陣(DD 是資料維度),算不起、也不需要;我們只要它乘上一個向量 vv 的結果。前向模式自動微分正好算這個:給輸入的一個擾動方向,輸出對應的擾動。

def total_derivative(f, x, t, v):          # d/dt f(x_t, t) along dx/dt = v
    _, df_dt = torch.func.jvp(f, (x, t), (v, torch.ones_like(t)))
    return df_dt

tangent 用 (v,1)(v,1)xx 方向擾動 vvtt 方向擾動 11,輸出就是 xfv+tf\nabla_xf\cdot v+\partial_tfU7 的 MeanFlow 會用同一招,只是它的 tangent 多一個分量。

sCM:為什麼原本訓不穩,怎麼穩住

連續時間的目標在 2023 年就寫出來了,但當時訓不穩 [3]。Lu & Song [2] 花了一整篇找原因,結論是:問題不在目標,在全導數 ddtfθ\frac{d}{dt}f_{\theta^-} 這個量的數值行為。把它拆開來看每一項,是這篇最值得學的方法論。

TrigFlow 參數化(boundary)。 先把路徑換成

xt=cost x0+sint z,t[0,π2],zN(0,σd2I),x_t=\cos t\ x_0+\sin t\ z,\qquad t\in[0,\tfrac\pi2],\quad z\sim\mathcal N(0,\sigma_d^2I),

consistency model 寫成

fθ(xt,t)=cost xtsint σdFθ ⁣(xtσd,t).f_\theta(x_t,t)=\cos t\ x_t-\sin t\ \sigma_d\,F_\theta\!\Big(\frac{x_t}{\sigma_d},t\Big).

對照 U6.2cskipx+coutFc_{\text{skip}}x+c_{\text{out}}Fcskip=costc_{\text{skip}}=\cos tcout=σdsintc_{\text{out}}=-\sigma_d\sin tt=0t=0f=xtf=x_t,邊界精確;兩個係數都是有界、光滑、彼此正交的三角函數,沒有 EDM 那組 σd2t2+σd2\frac{\sigma_d^2}{t^2+\sigma_d^2} 式的比值在 tt 大時的數值問題。更重要的是:同一個 FθF_\theta 也是這條路徑的 PF-ODE 速度(dxtdt=σdFθ\frac{dx_t}{dt}=\sigma_dF_\theta 是 diffusion 模型的參數化),所以 diffusion 模型與 consistency model 共用一個網路形式,可以直接從 teacher 初始化、也可以直接在同一份程式碼上切換 sCD 與 sCT。(EDM 的 tt 與這裡的三角時間之間差一個 tan\tantEDM=σdtantt_{\text{EDM}}=\sigma_d\tan t。)

在這個參數化下把全導數算出來:

ddtfθ(xt,t)=cost(σdFθdxtdt)    sint(xt+σddFθdt).\frac{d}{dt}f_{\theta^-}(x_t,t) =-\cos t\,\Big(\sigma_dF_{\theta^-}-\frac{dx_t}{dt}\Big) \;-\;\sin t\,\Big(x_t+\sigma_d\,\frac{dF_{\theta^-}}{dt}\Big).

(sCT 時 dxtdt\frac{dx_t}{dt} 用條件速度 costzsintx0\cos t\,z-\sin t\,x_0;sCD 時用 teacher。)兩個括號各是一個「該為零的東西」:第一個是 FθF_\theta 與真實速度的差——這一項在 FF 從 teacher 初始化時本來就小;第二個含 dFdt\frac{dF}{dt},是網路對時間的敏感度,這是不穩定的來源

Tangent normalization(variance)。 全導數 ddtfθ\frac{d}{dt}f_{\theta^-} 的模在不同 tt、不同樣本之間差很多個數量級,直接乘進損失會讓少數大 tangent 的樣本主導梯度。sCM 把它正規化:

g  gg+c,c=0.1,g\ \leftarrow\ \frac{g}{\|g\|+c},\qquad c=0.1,

(或直接 clip 到 [1,1][-1,1])。正規化只改變每個樣本的權重、不改變方向,所以不動極限解——梯度為零的條件仍是全導數為零——只是把 variance 壓平。

Tangent warmup(穩住 dFdt\frac{dF}{dt} 那一項)。 上面的第二個括號 sint(xt+σddFdt)-\sin t\,(x_t+\sigma_d\frac{dF}{dt}) 在訓練初期最不穩,因為 dFdt\frac{dF}{dt} 是網路對時間的敏感度、初始時沒有任何理由是平滑的。sCM 在這一項前面乘一個係數 rr,前一萬步從 0 線性升到 1——先只用第一個括號(等於一個接近 diffusion 的目標)把 FF 穩住,再逐漸打開時間方向的一致性。

時間 embedding 與 normalization 層(穩住 dFdt\frac{dF}{dt} 的另一半)。 dFdt\frac{dF}{dt} 經過時間 embedding 與各種 adaptive normalization 層,它們對 tt 的敏感度直接放大進 tangent。sCM 把時間輸入改成簡單的 cnoise(t)=tc_{\text{noise}}(t)=t 加小尺度的 positional embedding,並把 adaptive group normalization 改成一個對輸入與輸出都做正規化的版本(「adaptive double normalization」)。這與 iCT 調小 Fourier 尺度是同一個動機:fftt 平滑

Adaptive weighting(variance 的均衡)。 對不同 tt 的損失尺度,sCM 不再手選 λ(t)\lambda(t),而是學一個權重網路 wϕ(t)w_\phi(t),形式上是最大化一個「加權損失 - 權重的 log」,等價於讓每個 tt 的損失被自己的尺度正規化。這是 iCT 第三項 λ=1/Δt\lambda=1/\Delta t 的連續時間、自動化版本。

四個方法各站在哪裡

技術出處對付的誤差一句話
cskip,coutc_{\text{skip}},c_{\text{out}} 參數化CM / EDMboundary邊界寫進架構,排除平凡解
TrigFlow 參數化sCMboundary(+數值)有界的 cos,sin\cos,\sin 係數;與 diffusion 共用網路
目標側 stop-gradientCM(方向)資訊從邊界單向往外傳
拿掉 target 的 EMAiCTbiasEMA 讓極限目標不再是 consistency function
Pseudo-HuberiCTvariance壓離群樣本的梯度
λ=1/Δt\lambda=1/\Delta t、adaptive weightingiCT / sCMvariance(均衡)不同 tt 的損失拉到同一尺度
Lognormal 時間取樣iCTvariance(預算)把樣本放在 ff 真正在變的中段
NN 指數 curriculumCM / iCTbias↔variance先強訊號、再逐步消 bias
Δt0\Delta t\to0(連續時間)CM / sCMbias消掉 Euler 曲率誤差與 Jensen gap
條件速度代 x˙t\dot x_tsCT線性進入全導數,期望精確
Tangent normalization / clippingsCMvariance只改權重不改方向
Tangent warmup、時間 embedding、double normalizationsCMvariance(dFdt\frac{dF}{dt} 的穩定)fftt 平滑

三格都填滿了。值得注意的是最終那一格「Δt0\Delta t\to0」:它把 bias 整個消掉,代價是必須算全導數、必須面對 dFdt\frac{dF}{dt} 的數值行為——sCM 的其餘技術幾乎都是為了付這個代價。Lu & Song 用這套方法把 CT 放大到十億級參數、在 ImageNet 512 上兩步取樣接近 teacher 的品質 [2]。

先消化一下

想一想

iCT 把 target 網路的 EMA 拿掉(改成 θ=stopgrad(θ)\theta^-=\text{stopgrad}(\theta))。這一項改動在對付的是:

想一想

連續時間的自我一致性條件 tf+xfx˙t=0\partial_tf+\nabla_xf\cdot\dot x_t=0,相對於離散的相鄰點損失,多消掉了哪些誤差?

想一想

sCM 的 tangent normalization 把 ddtfθ\frac{d}{dt}f_{\theta^-} 除以 +c\|\cdot\|+c。它不會改變最終學到的 ff,理由是:

想一想

TrigFlow 用 xt=costx0+sintzx_t=\cos t\,x_0+\sin t\,zfθ=costxtsintσdFθf_\theta=\cos t\,x_t-\sin t\,\sigma_dF_\theta。相對於 EDM 那組 cskip,coutc_{\text{skip}},c_{\text{out}},它的好處不包括

參考文獻

  1. Song, Y., Dhariwal, P. Improved Techniques for Training Consistency Models. ICLR 2024.(iCT:EMA target 的 bias 分析、pseudo-Huber、λ=1/Δt\lambda=1/\Delta t、lognormal 時間取樣、指數 curriculum。)
  2. Lu, C., Song, Y. Simplifying, Stabilizing and Scaling Continuous-Time Consistency Models. 2024.(sCM:TrigFlow、全導數的分解、tangent normalization / warmup、時間 embedding 與 normalization 層、adaptive weighting、JVP。)
  3. Song, Y., Dhariwal, P., Chen, M., Sutskever, I. Consistency Models. ICML 2023.(連續時間 CM 的極限形式與原始的 NN curriculum。)
  4. Karras, T., Aittala, M., Aila, T., Laine, S. Elucidating the Design Space of Diffusion-Based Generative Models. NeurIPS 2022.(EDM 時間網格與 cskip,coutc_{\text{skip}},c_{\text{out}}。)