U2.2 15 分鐘閱讀 2026年9月

U2.2 先配對,再平均:Conditional Flow Matching

本篇重用M2.2站內人數變化 = 進 − 出:Continuity Equation·M1.2天氣預報該報幾度:MSE 的最小值是 Conditional Expectation·M1.1從鞋子猜身高:Conditional Expectation

目標寫得出來,可是算不出來

上一篇的結論是:要學的東西是速度場。那 loss 幾乎是自動的——就是一個回歸:

LFM(θ)=Et,  xpt[uθ(x,t)ut(x)2].\mathcal L_{\text{FM}}(\theta)=\mathbb E_{t,\;x\sim p_t}\big[\|u_\theta(x,t)-u_t(x)\|^2\big].

意思是:抽一個時刻 tt、抽一個那個時刻會出現的位置 xx,然後要求網路在那裡輸出的速度,跟正確答案一樣。

麻煩就出現在 ut(x)u_t(x) 這一項——那是「在位置 xx、時刻 tt 的正確速度」,本篇下一節會把它精確定義出來。它不是我們手上有的東西:要知道 xx 這個位置的正確速度,得先知道整條分布路徑 ptp_t 長什麼樣,而 ptp_t 又要有整個 pdatap_{\text{data}} 才定得下來。我們有的只是五萬個資料點

要學的目標本身,是一個我們算不出來的量。

這個處境見過了。U1.2 想要的 q(xt1xt)q(x_{t-1}\mid x_t)U1.3 想要的 logpt\nabla\log p_t,都是同樣寫不出來的邊際量,最後都是靠「條件的量很好算」繞過去的。

一個算不出來的東西,
還能拿它當回歸目標嗎?

把一個整體問題切成一堆小任務

U1.0 的倉庫搬貨可以直接搬過來用。整批貨要「搬得像 B 倉庫的擺法」很難評估,可是一旦每一箱的搬運路線都先寫好,每一箱該往哪走就只是查表。

Flow matching 就是照這個順序做的,分三步。

步驟 1:先抽一組配對。 抽一個起點 x0p0=N(0,I)x_0\sim p_0=\mathcal N(0,I)、抽一個資料點 x1pdatax_1\sim p_{\text{data}},把這一對記成 z=(x0,x1)z=(x_0,x_1)。兩邊各自獨立抽——這是一個選擇,不是必然,U3 會把它鬆開。

步驟 2:給定這組配對,把路線寫死。 例如

xt=αtx1+σtx0,α0=0, α1=1, σ0=1, σ1=0.x_t=\alpha_t x_1+\sigma_t x_0,\qquad \alpha_0=0,\ \alpha_1=1,\ \sigma_0=1,\ \sigma_1=0 .

邊界條件保證 t=0t=0 時人在 x0x_0t=1t=1 時人在 x1x_1。給定 zz,這是一條完全確定的曲線,速度直接微分就有:

ut(xz)=α˙tx1+σ˙tx0.u_t(x\mid z)=\dot\alpha_t x_1+\dot\sigma_t x_0 .

這是條件速度場。它好算到不需要網路——兩個純量乘兩個向量,加起來就是答案。

符號alpha 與 sigma 要滿足什麼,不用滿足什麼

αt\alpha_tσt\sigma_t 是我們自己挑的兩條純量函數,唯一的硬性要求是那四個邊界值(t=0t=0 全是噪聲、t=1t=1 全是資料)與可微。

不需要 αt2+σt2=1\alpha_t^2+\sigma_t^2=1。那是 DDPM 的 VP 路徑才有的關係;最常用的 αt=t, σt=1t\alpha_t=t,\ \sigma_t=1-t 就不滿足它。

步驟 3:把配對平均掉。 一組配對只描述一箱貨;我們要的是整批貨的行為。對 zz 積分:

pt(x)=pt(xz)p(z)dz,ut(x):=ut(xz)pt(xz)p(z)pt(x)dz=E[ut(xtz)  xt=x].p_t(x)=\int p_t(x\mid z)\,p(z)\,dz,\qquad u_t(x):=\int u_t(x\mid z)\,\frac{p_t(x\mid z)\,p(z)}{p_t(x)}\,dz =\mathbb E\big[u_t(x_t\mid z)\ \big|\ x_t=x\big].

用白話說:站在 xx 這個位置往外看,會有很多組配對的路線都經過這裡;把它們在這裡的速度,按「有多可能是它」加權平均起來,就是邊際速度。權重 pt(xz)p(z)/pt(x)p_t(x\mid z)p(z)/p_t(x) 就是後驗 p(zxt=x)p(z\mid x_t=x),所以離 xx 很遠的那些路線權重接近零,不會來干擾。

補充給定配對的位置是完全確定的,那密度是什麼意思

z=(x0,x1)z=(x_0,x_1) 時,xtx_t 完全被 zz 決定,所以 pt(z)p_t(\cdot\mid z) 嚴格說是一個 delta。想要一個真正的密度,可以改成只條件在資料點上、取 z=x1z=x_1:這時

pt(xx1)=N(αtx1, σt2I),ut(xx1)=α˙tx1+σ˙tσt(xαtx1).p_t(x\mid x_1)=\mathcal N\big(\alpha_t x_1,\ \sigma_t^2 I\big),\qquad u_t(x\mid x_1)=\dot\alpha_t x_1+\frac{\dot\sigma_t}{\sigma_t}\big(x-\alpha_t x_1\big).

兩種寫法給出同一個 utu_t——由 tower property,先對 x0x_0 平均再對 x1x_1 平均,跟一次對 (x0,x1)(x_0,x_1) 平均是一樣的。下面的 demo 為了讓權重看得見,把 delta 給了一點寬度。

課堂提問Q1

既然邊際速度就是一個加權平均,那最直接的做法應該是先把這個平均估出來:對每一個 (x,t)(x,t) 抽很多組配對、算它們的權重、加權平均,拿這個估計值當回歸目標。

這樣做會遇到什麼?

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

先看它到底要算什麼。權重是 wipt(xzi)w_i\propto p_t(x\mid z_i),這是一個 self-normalized importance sampling:抽 MM 組配對,估計值是 u^=iwiut(xzi)\hat u=\sum_i w_i\,u_t(x\mid z_i)

問題在於有效樣本數wiw_ix(αtx1,i+σtx0,i)2\|x-(\alpha_t x_{1,i}+\sigma_t x_{0,i})\|^2 決定,而在 dd 維裡這個距離的變動是隨 dd 累加的,所以權重會集中在極少數幾組配對上——有效樣本數 1/iwi21/\sum_i w_i^2 迅速掉到個位數,甚至掉到 1。下面的 demo 在 2D 就已經看得到這件事:一個位置真正「有權重」的只有十幾條。要在 d=3072d=3072 得到一個堪用的估計,MM 得大到不可能。

這個平均我們根本不需要自己算。

回頭想想 U1.2 的結論——L2L_2 回歸的最佳解就是條件期望。如果我們把 loss 寫成「拿條件速度當目標」:

LCFM(θ)=Et,  z,  xpt(z)[uθ(x,t)ut(xz)2],\mathcal L_{\text{CFM}}(\theta)=\mathbb E_{t,\;z,\;x\sim p_t(\cdot\mid z)}\big[\|u_\theta(x,t)-u_t(x\mid z)\|^2\big],

那它的最小值就是 E[ut(xtz)xt=x]\mathbb E[u_t(x_t\mid z)\mid x_t=x],而這正好是我們要的 ut(x)u_t(x)。目標每一次都只用一組配對算,噪聲很大,但那個噪聲的平均值是對的——回歸會把它平掉。

所以整件事的順序是:不要估平均,把平均交給 loss

互動 demo:條件速度疊成邊際速度。 拖著 xx 在圖上走,看有哪幾條配對路線會經過它、它們各自的速度指向哪裡。細箭頭是條件速度(每一條都好算),粗箭頭是加權平均後的邊際速度——注意兩者的長度比例,以及「有效參與的條數」這個數字。demo 用的是最簡單的直線路徑 αt=t, σt=1t\alpha_t=t,\ \sigma_t=1-t,也就是下一篇的主角。

為什麼可以拿條件目標代替?

上面那段話裡藏了兩個需要驗證的環節:utu_t 真的推得動 ptp_t 嗎?以及 LCFM\mathcal L_{\text{CFM}} 真的能學到 utu_t 嗎?

定理 1(utu_t 是對的目標)。 若每個 ut(z)u_t(\cdot\mid z) 生成 pt(z)p_t(\cdot\mid z),則步驟 3 定義的 utu_t 生成 ptp_t

定理 2(CFM 能學到它)。 θLFM=θLCFM\nabla_\theta\mathcal L_{\text{FM}}=\nabla_\theta\mathcal L_{\text{CFM}}

第二個比「最小值相同」更強:梯度處處相同,意思是連訓練過程中每一步走的方向都一樣。我們不是在解一個近似問題,而是在跑同一條 gradient descent,只是每一步的目標換成了算得出來的版本。

展開細節定理 1 的證明:continuity equation 是線性的

對每一個 zz,條件路徑滿足

tpt(xz)=(pt(xz)ut(xz)).\partial_t p_t(x\mid z)=-\nabla\cdot\big(p_t(x\mid z)\,u_t(x\mid z)\big).

兩邊乘 p(z)p(z) 再對 zz 積分。左邊直接得到 tpt(x)\partial_t p_t(x)。右邊把散度移到積分外(\nabla 只對 xx 作用):

pt(xz)ut(xz)p(z)dz=(pt(x)ut(x)),-\nabla\cdot\int p_t(x\mid z)\,u_t(x\mid z)\,p(z)\,dz=-\nabla\cdot\big(p_t(x)\,u_t(x)\big),

最後一步用的就是 utu_t 的定義——它是被定義成「讓這個積分等於 ptutp_t u_t」的那個東西。\square

關鍵在於 continuity equation 對 pp線性的,所以「每一箱貨各自守恆」加起來就是「整批貨守恆」。

這個證明需要 pt(z)p_t(\cdot\mid z) 真的是一個密度,所以它走的是上面 <Remark>z=x1z=x_1 的那個版本。訓練時用的 z=(x0,x1)z=(x_0,x_1) 版本是它的推論——由 tower property,先對 x0x_0 平均再對 x1x_1 平均,跟一次對 (x0,x1)(x_0,x_1) 平均得到同一個 utu_t

展開細節定理 2 的證明:展開平方,只有交叉項要處理

把兩個 loss 都展開成三項。uθ2\|u_\theta\|^2 那一項在兩邊完全相同,因為 LCFM\mathcal L_{\text{CFM}}xx 的邊際分布正是 ptp_t。不含 θ\theta 的那一項對梯度沒有貢獻。剩下交叉項:

Expt[uθ(x)ut(x)]=Expt[uθ(x)E[ut(xtz)xt=x]]=Ez,  xpt(z)[uθ(x)ut(xz)].\mathbb E_{x\sim p_t}\big[u_\theta(x)^\top u_t(x)\big] =\mathbb E_{x\sim p_t}\Big[u_\theta(x)^\top\,\mathbb E\big[u_t(x_t\mid z)\mid x_t=x\big]\Big] =\mathbb E_{z,\;x\sim p_t(\cdot\mid z)}\big[u_\theta(x)^\top u_t(x\mid z)\big].

第一個等號是 utu_t 的定義,第二個是 tower property。三項逐項相等(或差一個常數),所以梯度相同。\square

那個常數也寫得出來:LCFMLFM=ExpttrVar[ut(xtz)xt=x]\mathcal L_{\text{CFM}}-\mathcal L_{\text{FM}}=\mathbb E_{x\sim p_t}\operatorname{tr}\operatorname{Var}\big[u_t(x_t\mid z)\mid x_t=x\big],也就是條件速度繞著邊際速度的變異數。它不含 θ\theta,所以不影響最小值點;但它讓 LCFM\mathcal L_{\text{CFM}} 不會收到 0——和 U1.3 那個「loss 停在 E[trVar(ϵxt)]\mathbb E[\operatorname{tr}\operatorname{Var}(\epsilon\mid x_t)]」是同一件事。

這是 conditional trick 的速度版本

U1.2 把這個模式命名成 conditional trick,而且證明過它的 KL 版本:把 x0x_0 放進條件裡,objective 只差一個與參數無關的常數,最佳解仍然是那個算不出來的邊際。這一篇做的是同一件事,只是主角從分布換成了速度場。

三次放在一起看:

U1.2(KL)U1.3(MSE)這一篇(velocity)
想要的邊際量q(xt1xt)q(x_{t-1}\mid x_t)logpt(x)\nabla\log p_t(x)ut(x)u_t(x)
好算的條件量q(xt1xt,x0)q(x_{t-1}\mid x_t,x_0)ϵ\epsilon(給定 (x0,xt)(x_0,x_t) 就知道)ut(xz)u_t(x\mid z)(給定配對就知道)
放進條件裡的x0x_0x0x_0配對 zz
誰去取平均KL 的交叉熵項L2L_2 回歸L2L_2 回歸
差的那個常數θ\theta 無關E[trVar(ϵxt)]\mathbb E[\operatorname{tr}\operatorname{Var}(\epsilon\mid x_t)]E[trVar(ut(xtz)xt)]\mathbb E[\operatorname{tr}\operatorname{Var}(u_t(x_t\mid z)\mid x_t)]

(表裡前兩欄用的是 U1 的慣例——那裡的 x0x_0資料;最右欄是本單元的慣例,配對寫成 z=(x0,x1)z=(x_0,x_1)x0x_0噪聲。)

難算的邊際量 + 好算的條件量 + 一個算得出來的 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

x0x_0x1x_1各自獨立抽的,所以配對是隨機的:一個角落的噪聲點很可能被配到對面角落的資料點,路線因此交錯成一團。

這樣學出來的東西還會是對的嗎?

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

分布是對的。定理 1 和定理 2 的證明從頭到尾沒有用到「路線不交叉」,也沒有用到 x0x_0x1x_1 獨立——只用到條件路徑的兩個端點對、以及條件期望的性質。所以 p1=pdatap_1=p_{\text{data}} 這件事不受影響。

軌跡會付出代價。網路在一個位置只能輸出一個速度,而交錯的地方有很多條路線經過、方向各不相同,於是它輸出的是那些方向的平均。demo 裡那支明顯比細箭頭短的粗箭頭就是這件事:平均之後不但變短,也不再對齊任何一條原本的直線。

後果是實際軌跡是彎的。彎的軌跡在步數多的時候沒差,步數一少,離散化誤差就跑出來——U2.3 會把這件事算清楚。

順著這個念頭,一個很自然的問題是:既然配對是我們自己挑的,能不能挑一組不那麼交錯的?可以,而且這正是 U3.3U3.4 在做的事。

消化一下

想一想

邊際速度 ut(x)u_t(x) 是:

想一想

定理 2 說 θLFM=θLCFM\nabla_\theta\mathcal L_{\text{FM}}=\nabla_\theta\mathcal L_{\text{CFM}}。那兩個 loss 的數值呢?

想一想

下列哪一個改動會讓這個框架失效?

參考文獻

  1. Lipman, Y., Chen, R. T. Q., Ben-Hamu, H., Nickel, M., Le, M. Flow Matching for Generative Modeling. ICLR 2023. (這一篇的主線:條件路徑、CFM loss 與兩個定理。)
  2. Albergo, M. S., Vanden-Eijnden, E. Building Normalizing Flows with Stochastic Interpolants. ICLR 2023. (同一個框架的另一條發展線,把 (αt,σt)(\alpha_t,\sigma_t) 當成設計變數。)
  3. 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 提到的「換一組配對」。)