U6.4 讓 CT 訓得起來:iCT 與連續時間的 sCM
本篇重用M0.3漂流的溫度計:全導數、微分穿過積分與 JVP·M5.2追一隻自己也在跑的狗:移動目標、EMA 與 Stop-Gradient·M2.3導航每 30 秒更新一次:Euler 法與它的誤差
訓練不穩,是哪三件事在作用
上一篇結尾留下的圖像是:CT 的目標是 CD 目標的條件版本,代價是 bias 與 variance 的取捨。把這個單元到目前為止碰過的所有誤差整理成三類,接下來每一項技術都會被放進其中一格:
- boundary: 有沒有精確成立、參數化在兩端是否有界。
- bias:目標的期望偏離真實終點——來自 Euler 對彎曲軌跡的 誤差、 非線性的 Jensen gap,以及任何讓「目標」與「真值」系統性不同的東西。
- variance:目標本身的隨機散佈——條件目標退向自己的 、極端樣本、以及 小時訊噪比變差。
這一篇不引入新的損失,新的物件只有一個:沿軌跡的全導數。前半是離散時間的修補(iCT [1]),後半是把 推到零(連續時間 CM,sCM [2])。
課堂提問Q1
下一節的第一個改動是把 target 網路的 EMA 拿掉。這和一般的直覺相反:目標會動的訓練(bootstrapping)不是都要靠 EMA 才不發散嗎?
拿掉之後, 為什麼不會直接崩掉?
先想一想,再展開看整理後的答案
關鍵是這裡的「目標會動」和 RL 那種不一樣。
RL 裡 target network 要 EMA,是因為目標值 本身就是網路輸出,沒有任何外部的錨;目標追著自己跑,所以要把它拖慢。
CT 的目標是 ,看起來也是自己。但它被兩件事釘住了:
- 邊界條件寫在架構裡。 U6.2 那一節說過, 對任何一組參數都成立。所以最靠近乾淨端的那一格是正確答案,不管網路好不好。訊號是從那裡往外傳的,不是憑空自舉。
- ,目標永遠比輸入更靠近乾淨端。 每一對都是「已經對的那一側教還沒對的那一側」,方向是單向的,沒有兩邊互相追的迴路。
所以 CT 不需要靠 EMA 去穩住一個沒有錨的自舉。而 EMA 反過來會壞事:上一篇那條「CT 與 CD 損失只差 」的定理,條件是 。 一旦是滯後的 EMA,極限就換成另一個目標—— 收斂到的不再是 consistency function。這是一個 bias,而且是那種「怎麼訓都到不了」的 bias,不是訓久一點就好。
換成 之後,目標是當前參數(只是不回傳梯度),定理的條件回來了。那為什麼還要 stop-gradient?因為要梯度的話,損失就變成在最小化「 在相鄰兩點的差」,而那正是上一題那個退化解喜歡的東西。
一句話:EMA 在這裡壓的不是不穩定,是收斂的位置。
iCT:五個改動,各對一格
Song & Dhariwal [1] 把 2023 年的 CT 從「勉強能訓、品質明顯輸 CD」改到「不用 teacher 也能追上或超過 CD」。改動有五個,每一個都能對回上面的三格。
一、拿掉 target 的 EMA(bias)。 原版用 當目標網路。直覺上 EMA 是在壓 variance——讓目標平滑一點。但 iCT 證明它其實引入 bias:上一篇說 CT 損失與 CD 損失在 時只差 ,那條定理的條件是 ;一旦 是一個滯後的 EMA,極限就變成另一個目標, 不再收斂到 consistency function。改成 ——目標是當前參數、只是不回傳梯度——後品質大幅提升。這一條值得記住,因為它反直覺:一個看起來是穩定器的東西,在這裡是 bias 的來源。(U6.2 說的「目標側要 stop-gradient」仍然成立,只是不再另加 EMA。)
二、Pseudo-Huber 取代 (variance)。 用
( 是資料維度)。它在 時像 、在 時像 ,梯度的模有上界。CT 的目標帶著 的隨機退向,偶爾會有離群的樣本把 的梯度拉得很大;pseudo-Huber 把這些離群值的影響壓下來。原版用的 LPIPS 則被拿掉——它讓 學到評估指標的特徵、對 FID 有系統性的不公平影響,是另一種 bias。
三、時間權重 (variance 的均衡)。 相鄰兩點輸出的差距是 ;用 EDM 的時間網格 [4]( 小處密、 大處疏)時,不同 的 差好幾個數量級,loss 的尺度也跟著差好幾個數量級。乘上 把每一對的貢獻拉到同一尺度,等效於對 小的區域加重——那裡 接近資料、細節在那裡決定。
四、Lognormal 時間取樣(把預算放在有訊號的地方)。 不再均勻抽 ,而是讓 近似服從 ,、。這就是 U2.4 那第四個旋鈕——訓練端的加權(U3.5 把它換算成 上的權重):兩端的一致性條件太容易( 極小時 恆等、 極大時 幾乎沒資訊),中段才是 真正在變的地方,訊號在那裡。
五、 的指數 curriculum(bias 隨時間換 variance)。 原版的 線性地從 2 增加到 150;iCT 改成
是訓練步、 讓翻倍的次數剛好填滿整段訓練。這是上一篇那個取捨的排程:一開始 小、bias 大但訊號強;每翻倍一次就把 bias 砍一半,而 已經接近、variance 造成的擾動可承受。終點 比原版大得多,因為前面四項把 variance 壓下來了,才容得下這麼小的 。
展開細節其他兩個小改動:Fourier 尺度與 dropout
iCT 還把時間 embedding 的 Fourier 尺度從 16 調到 0.02,並提高 dropout。前者的理由與下半篇的 sCM 相通:CT 比的是兩個相鄰時刻的輸出, 對 越敏感,這個差就越被 方向的高頻擾動污染;把 Fourier 尺度調小等於強迫 對 平滑。這一條歸在 variance。(具體數值以原文為準。)
把 推到零
iCT 的五項全部是在離散時間網格上和 bias–variance 周旋。有一個更直接的想法:既然 bias 來自 不為零,那就讓 ,看損失變成什麼。
自我一致性說 沿軌跡不變。「不變」在微分的語言裡就是沿軌跡的全導數為零:
這是這一篇唯一的新物件。逐項讀:第一項是「時間往前走、位置不動」時 的變化;第二項是「位置隨軌跡移動」帶來的變化, 是 PF-ODE 的速度(本單元座標下 )。兩項要恰好相消。U3.2 裡出現過同一個算子——那裡是對速度場 取 得到加速度;這裡是對 取,得到「 沿軌跡的漂移」。
Song et al. [3] 已經推出: 版的 CM 損失在 時,梯度趨近
( 側的全導數不回傳梯度)。把它推到零,就是要 沿軌跡的全導數為零。這個目標沒有 :Euler 的曲率誤差、Jensen gap,兩個 的 bias 來源在極限裡一起消失。
那要怎麼在沒有 teacher 的情形下算 ?和上一篇一樣,用 U1.2 的 conditional trick。 對 微分,條件速度是 本身。而 對 是線性的,所以用 代替 後取期望是精確的——連上一篇那個 Jensen gap 都沒有了,因為現在條件量不再被包在非線性函數裡面,而是線性地進入一個乘積。連續時間 CT 在這裡又多賺一項。
計算上, 是 在 沿方向 的方向導數,一次 Jacobian-vector product(JVP)就得到,不用算整個 Jacobian:torch.func.jvp(f, (x, t), (v, ones))。成本約等於一次前向傳播。
展開細節全導數的 JVP:三行程式與為什麼不算 Jacobian
是 的矩陣( 是資料維度),算不起、也不需要;我們只要它乘上一個向量 的結果。前向模式自動微分正好算這個:給輸入的一個擾動方向,輸出對應的擾動。
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_dttangent 用 : 方向擾動 、 方向擾動 ,輸出就是 。U7 的 MeanFlow 會用同一招,只是它的 tangent 多一個分量。
sCM:為什麼原本訓不穩,怎麼穩住
連續時間的目標在 2023 年就寫出來了,但當時訓不穩 [3]。Lu & Song [2] 花了一整篇找原因,結論是:問題不在目標,在全導數 這個量的數值行為。把它拆開來看每一項,是這篇最值得學的方法論。
TrigFlow 參數化(boundary)。 先把路徑換成
consistency model 寫成
對照 U6.2 的 :、。 時 ,邊界精確;兩個係數都是有界、光滑、彼此正交的三角函數,沒有 EDM 那組 式的比值在 大時的數值問題。更重要的是:同一個 也是這條路徑的 PF-ODE 速度( 是 diffusion 模型的參數化),所以 diffusion 模型與 consistency model 共用一個網路形式,可以直接從 teacher 初始化、也可以直接在同一份程式碼上切換 sCD 與 sCT。(EDM 的 與這裡的三角時間之間差一個 :。)
在這個參數化下把全導數算出來:
(sCT 時 用條件速度 ;sCD 時用 teacher。)兩個括號各是一個「該為零的東西」:第一個是 與真實速度的差——這一項在 從 teacher 初始化時本來就小;第二個含 ,是網路對時間的敏感度,這是不穩定的來源。
Tangent normalization(variance)。 全導數 的模在不同 、不同樣本之間差很多個數量級,直接乘進損失會讓少數大 tangent 的樣本主導梯度。sCM 把它正規化:
(或直接 clip 到 )。正規化只改變每個樣本的權重、不改變方向,所以不動極限解——梯度為零的條件仍是全導數為零——只是把 variance 壓平。
Tangent warmup(穩住 那一項)。 上面的第二個括號 在訓練初期最不穩,因為 是網路對時間的敏感度、初始時沒有任何理由是平滑的。sCM 在這一項前面乘一個係數 ,前一萬步從 0 線性升到 1——先只用第一個括號(等於一個接近 diffusion 的目標)把 穩住,再逐漸打開時間方向的一致性。
時間 embedding 與 normalization 層(穩住 的另一半)。 經過時間 embedding 與各種 adaptive normalization 層,它們對 的敏感度直接放大進 tangent。sCM 把時間輸入改成簡單的 加小尺度的 positional embedding,並把 adaptive group normalization 改成一個對輸入與輸出都做正規化的版本(「adaptive double normalization」)。這與 iCT 調小 Fourier 尺度是同一個動機:讓 對 平滑。
Adaptive weighting(variance 的均衡)。 對不同 的損失尺度,sCM 不再手選 ,而是學一個權重網路 ,形式上是最大化一個「加權損失 權重的 log」,等價於讓每個 的損失被自己的尺度正規化。這是 iCT 第三項 的連續時間、自動化版本。
四個方法各站在哪裡
| 技術 | 出處 | 對付的誤差 | 一句話 |
|---|---|---|---|
| 參數化 | CM / EDM | boundary | 邊界寫進架構,排除平凡解 |
| TrigFlow 參數化 | sCM | boundary(+數值) | 有界的 係數;與 diffusion 共用網路 |
| 目標側 stop-gradient | CM | (方向) | 資訊從邊界單向往外傳 |
| 拿掉 target 的 EMA | iCT | bias | EMA 讓極限目標不再是 consistency function |
| Pseudo-Huber | iCT | variance | 壓離群樣本的梯度 |
| 、adaptive weighting | iCT / sCM | variance(均衡) | 不同 的損失拉到同一尺度 |
| Lognormal 時間取樣 | iCT | variance(預算) | 把樣本放在 真正在變的中段 |
| 指數 curriculum | CM / iCT | bias↔variance | 先強訊號、再逐步消 bias |
| (連續時間) | CM / sCM | bias | 消掉 Euler 曲率誤差與 Jensen gap |
| 條件速度代 | sCT | — | 線性進入全導數,期望精確 |
| Tangent normalization / clipping | sCM | variance | 只改權重不改方向 |
| Tangent warmup、時間 embedding、double normalization | sCM | variance( 的穩定) | 讓 對 平滑 |
三格都填滿了。值得注意的是最終那一格「」:它把 bias 整個消掉,代價是必須算全導數、必須面對 的數值行為——sCM 的其餘技術幾乎都是為了付這個代價。Lu & Song 用這套方法把 CT 放大到十億級參數、在 ImageNet 512 上兩步取樣接近 teacher 的品質 [2]。
先消化一下
參考文獻
- Song, Y., Dhariwal, P. Improved Techniques for Training Consistency Models. ICLR 2024.(iCT:EMA target 的 bias 分析、pseudo-Huber、、lognormal 時間取樣、指數 curriculum。)
- Lu, C., Song, Y. Simplifying, Stabilizing and Scaling Continuous-Time Consistency Models. 2024.(sCM:TrigFlow、全導數的分解、tangent normalization / warmup、時間 embedding 與 normalization 層、adaptive weighting、JVP。)
- Song, Y., Dhariwal, P., Chen, M., Sutskever, I. Consistency Models. ICML 2023.(連續時間 CM 的極限形式與原始的 curriculum。)
- Karras, T., Aittala, M., Aila, T., Laine, S. Elucidating the Design Space of Diffusion-Based Generative Models. NeurIPS 2022.(EDM 時間網格與 。)