本篇重用M2.2站內人數變化 = 進 − 出:Continuity Equation·M1.2天氣預報該報幾度:MSE 的最小值是 Conditional Expectation·M1.1從鞋子猜身高:Conditional Expectation
目標寫得出來,可是算不出來
上一篇的結論是:要學的東西是速度場。那 loss 幾乎是自動的——就是一個回歸:
LFM(θ)=Et,x∼pt[∥uθ(x,t)−ut(x)∥2].
意思是:抽一個時刻 t、抽一個那個時刻會出現的位置 x,然後要求網路在那裡輸出的速度,跟正確答案一樣。
麻煩就出現在 ut(x) 這一項——那是「在位置 x、時刻 t 的正確速度」,本篇下一節會把它精確定義出來。它不是我們手上有的東西:要知道 x 這個位置的正確速度,得先知道整條分布路徑 pt 長什麼樣,而 pt 又要有整個 pdata 才定得下來。我們有的只是五萬個資料點。
要學的目標本身,是一個我們算不出來的量。
這個處境見過了。U1.2 想要的 q(xt−1∣xt)、U1.3 想要的 ∇logpt,都是同樣寫不出來的邊際量,最後都是靠「條件的量很好算」繞過去的。
把一個整體問題切成一堆小任務
U1.0 的倉庫搬貨可以直接搬過來用。整批貨要「搬得像 B 倉庫的擺法」很難評估,可是一旦每一箱的搬運路線都先寫好,每一箱該往哪走就只是查表。
Flow matching 就是照這個順序做的,分三步。
步驟 1:先抽一組配對。 抽一個起點 x0∼p0=N(0,I)、抽一個資料點 x1∼pdata,把這一對記成 z=(x0,x1)。兩邊各自獨立抽——這是一個選擇,不是必然,U3 會把它鬆開。
步驟 2:給定這組配對,把路線寫死。 例如
xt=αtx1+σtx0,α0=0, α1=1, σ0=1, σ1=0.
邊界條件保證 t=0 時人在 x0、t=1 時人在 x1。給定 z,這是一條完全確定的曲線,速度直接微分就有:
ut(x∣z)=α˙tx1+σ˙tx0.
這是條件速度場。它好算到不需要網路——兩個純量乘兩個向量,加起來就是答案。
符號alpha 與 sigma 要滿足什麼,不用滿足什麼
αt 與 σt 是我們自己挑的兩條純量函數,唯一的硬性要求是那四個邊界值(t=0 全是噪聲、t=1 全是資料)與可微。
不需要 αt2+σt2=1。那是 DDPM 的 VP 路徑才有的關係;最常用的 αt=t, σt=1−t 就不滿足它。
步驟 3:把配對平均掉。 一組配對只描述一箱貨;我們要的是整批貨的行為。對 z 積分:
pt(x)=∫pt(x∣z)p(z)dz,ut(x):=∫ut(x∣z)pt(x)pt(x∣z)p(z)dz=E[ut(xt∣z) xt=x].
用白話說:站在 x 這個位置往外看,會有很多組配對的路線都經過這裡;把它們在這裡的速度,按「有多可能是它」加權平均起來,就是邊際速度。權重 pt(x∣z)p(z)/pt(x) 就是後驗 p(z∣xt=x),所以離 x 很遠的那些路線權重接近零,不會來干擾。
補充給定配對的位置是完全確定的,那密度是什麼意思
取 z=(x0,x1) 時,xt 完全被 z 決定,所以 pt(⋅∣z) 嚴格說是一個 delta。想要一個真正的密度,可以改成只條件在資料點上、取 z=x1:這時
pt(x∣x1)=N(αtx1, σt2I),ut(x∣x1)=α˙tx1+σtσ˙t(x−αtx1).兩種寫法給出同一個 ut——由 tower property,先對 x0 平均再對 x1 平均,跟一次對 (x0,x1) 平均是一樣的。下面的 demo 為了讓權重看得見,把 delta 給了一點寬度。
課堂提問Q1
既然邊際速度就是一個加權平均,那最直接的做法應該是先把這個平均估出來:對每一個 (x,t) 抽很多組配對、算它們的權重、加權平均,拿這個估計值當回歸目標。
這樣做會遇到什麼?
先想一想,再展開看整理後的答案
先看它到底要算什麼。權重是 wi∝pt(x∣zi),這是一個 self-normalized importance sampling:抽 M 組配對,估計值是 u^=∑iwiut(x∣zi)。
問題在於有效樣本數。wi 由 ∥x−(αtx1,i+σtx0,i)∥2 決定,而在 d 維裡這個距離的變動是隨 d 累加的,所以權重會集中在極少數幾組配對上——有效樣本數 1/∑iwi2 迅速掉到個位數,甚至掉到 1。下面的 demo 在 2D 就已經看得到這件事:一個位置真正「有權重」的只有十幾條。要在 d=3072 得到一個堪用的估計,M 得大到不可能。
這個平均我們根本不需要自己算。
回頭想想 U1.2 的結論——L2 回歸的最佳解就是條件期望。如果我們把 loss 寫成「拿條件速度當目標」:
LCFM(θ)=Et,z,x∼pt(⋅∣z)[∥uθ(x,t)−ut(x∣z)∥2],那它的最小值就是 E[ut(xt∣z)∣xt=x],而這正好是我們要的 ut(x)。目標每一次都只用一組配對算,噪聲很大,但那個噪聲的平均值是對的——回歸會把它平掉。
所以整件事的順序是:不要估平均,把平均交給 loss。
互動 demo:條件速度疊成邊際速度。 拖著 x 在圖上走,看有哪幾條配對路線會經過它、它們各自的速度指向哪裡。細箭頭是條件速度(每一條都好算),粗箭頭是加權平均後的邊際速度——注意兩者的長度比例,以及「有效參與的條數」這個數字。demo 用的是最簡單的直線路徑 αt=t, σt=1−t,也就是下一篇的主角。
為什麼可以拿條件目標代替?
上面那段話裡藏了兩個需要驗證的環節:ut 真的推得動 pt 嗎?以及 LCFM 真的能學到 ut 嗎?
定理 1(ut 是對的目標)。 若每個 ut(⋅∣z) 生成 pt(⋅∣z),則步驟 3 定義的 ut 生成 pt。
定理 2(CFM 能學到它)。 ∇θLFM=∇θLCFM。
第二個比「最小值相同」更強:梯度處處相同,意思是連訓練過程中每一步走的方向都一樣。我們不是在解一個近似問題,而是在跑同一條 gradient descent,只是每一步的目標換成了算得出來的版本。
展開細節定理 1 的證明:continuity equation 是線性的
對每一個 z,條件路徑滿足
∂tpt(x∣z)=−∇⋅(pt(x∣z)ut(x∣z)).兩邊乘 p(z) 再對 z 積分。左邊直接得到 ∂tpt(x)。右邊把散度移到積分外(∇ 只對 x 作用):
−∇⋅∫pt(x∣z)ut(x∣z)p(z)dz=−∇⋅(pt(x)ut(x)),最後一步用的就是 ut 的定義——它是被定義成「讓這個積分等於 ptut」的那個東西。□
關鍵在於 continuity equation 對 p 是線性的,所以「每一箱貨各自守恆」加起來就是「整批貨守恆」。
這個證明需要 pt(⋅∣z) 真的是一個密度,所以它走的是上面 <Remark> 裡 z=x1 的那個版本。訓練時用的 z=(x0,x1) 版本是它的推論——由 tower property,先對 x0 平均再對 x1 平均,跟一次對 (x0,x1) 平均得到同一個 ut。
展開細節定理 2 的證明:展開平方,只有交叉項要處理
把兩個 loss 都展開成三項。∥uθ∥2 那一項在兩邊完全相同,因為 LCFM 裡 x 的邊際分布正是 pt。不含 θ 的那一項對梯度沒有貢獻。剩下交叉項:
Ex∼pt[uθ(x)⊤ut(x)]=Ex∼pt[uθ(x)⊤E[ut(xt∣z)∣xt=x]]=Ez,x∼pt(⋅∣z)[uθ(x)⊤ut(x∣z)].第一個等號是 ut 的定義,第二個是 tower property。三項逐項相等(或差一個常數),所以梯度相同。□
那個常數也寫得出來:LCFM−LFM=Ex∼pttrVar[ut(xt∣z)∣xt=x],也就是條件速度繞著邊際速度的變異數。它不含 θ,所以不影響最小值點;但它讓 LCFM 不會收到 0——和 U1.3 那個「loss 停在 E[trVar(ϵ∣xt)]」是同一件事。
這是 conditional trick 的速度版本
U1.2 把這個模式命名成 conditional trick,而且證明過它的 KL 版本:把 x0 放進條件裡,objective 只差一個與參數無關的常數,最佳解仍然是那個算不出來的邊際。這一篇做的是同一件事,只是主角從分布換成了速度場。
三次放在一起看:
| U1.2(KL) | U1.3(MSE) | 這一篇(velocity) |
|---|
| 想要的邊際量 | q(xt−1∣xt) | ∇logpt(x) | ut(x) |
| 好算的條件量 | q(xt−1∣xt,x0) | ϵ(給定 (x0,xt) 就知道) | ut(x∣z)(給定配對就知道) |
| 放進條件裡的 | x0 | x0 | 配對 z |
| 誰去取平均 | KL 的交叉熵項 | L2 回歸 | L2 回歸 |
| 差的那個常數 | 與 θ 無關 | E[trVar(ϵ∣xt)] | E[trVar(ut(xt∣z)∣xt)] |
(表裡前兩欄用的是 U1 的慣例——那裡的 x0 是資料;最右欄是本單元的慣例,配對寫成 z=(x0,x1),x0 是噪聲。)
難算的邊際量 + 好算的條件量 + 一個算得出來的 loss:平均那一步交給 loss 自己做。
書上(U1.2 的參考文獻 [2])把這三個版本並排稱為 diffusion models 的 conditional tricks——它們不是三個各自獨立的技巧,是同一個策略在三種 objective 上的樣子。 U3 還會再用一次。
訓練只換掉目標那一行
repeat
x1 ← 一個 batch 的資料
x0 ← N(0, I)
t ← Uniform[0, 1]
xt ← α_t · x1 + σ_t · x0
loss ← ‖u_θ(xt, t) − (α̇_t · x1 + σ̇_t · x0)‖²
梯度下降
和 U1.2 的五行只差目標那一行。沒有 SDE、沒有 score、沒有 Tweedie,也沒有任何積分要估。
課堂提問Q2
x0 和 x1 是各自獨立抽的,所以配對是隨機的:一個角落的噪聲點很可能被配到對面角落的資料點,路線因此交錯成一團。
這樣學出來的東西還會是對的嗎?
先想一想,再展開看整理後的答案
分布是對的。定理 1 和定理 2 的證明從頭到尾沒有用到「路線不交叉」,也沒有用到 x0 與 x1 獨立——只用到條件路徑的兩個端點對、以及條件期望的性質。所以 p1=pdata 這件事不受影響。
但軌跡會付出代價。網路在一個位置只能輸出一個速度,而交錯的地方有很多條路線經過、方向各不相同,於是它輸出的是那些方向的平均。demo 裡那支明顯比細箭頭短的粗箭頭就是這件事:平均之後不但變短,也不再對齊任何一條原本的直線。
後果是實際軌跡是彎的。彎的軌跡在步數多的時候沒差,步數一少,離散化誤差就跑出來——U2.3 會把這件事算清楚。
順著這個念頭,一個很自然的問題是:既然配對是我們自己挑的,能不能挑一組不那麼交錯的?可以,而且這正是 U3.3 與 U3.4 在做的事。
消化一下
參考文獻
- Lipman, Y., Chen, R. T. Q., Ben-Hamu, H., Nickel, M., Le, M. Flow Matching for Generative Modeling. ICLR 2023. (這一篇的主線:條件路徑、CFM loss 與兩個定理。)
- Albergo, M. S., Vanden-Eijnden, E. Building Normalizing Flows with Stochastic Interpolants. ICLR 2023. (同一個框架的另一條發展線,把 (αt,σt) 當成設計變數。)
- Tong, A., Fatras, K., Malkin, N., Huguet, G., Zhang, Y., Rector-Brooks, J., Wolf, G., Bengio, Y. Improving and Generalizing Flow-Based Generative Models with Minibatch Optimal Transport. TMLR 2024. (Q2 提到的「換一組配對」。)