U7.6 實作:MeanFlow、半群檢查,與交叉問題的最終回收
本篇重用M0.3漂流的溫度計:全導數、微分穿過積分與 JVP·M5.3只看得到樣本,怎麼調機器:對 Generator 微分·M6.1兩個城市的人口分佈差多少:W₂ 距離
沿用同一組 toy
資料、網路骨架、exact_score 都沿用 U1.5 建立的 toy.py(two moons,)。這一篇用 U2 的 FM 慣例: 噪聲、 資料、、條件速度 。網路 model(z, r, t):三層 MLP,寬 128, 與 各自做 sinusoidal embedding 後串接——比 U2.5 的網路多一個時間輸入,其餘不變。
需要從前幾個單元帶進來的東西:U3.6 的 reflow×2 模型與 函數、U6.6 的 consistency training 一步生成器,以及那一張「三種一步方法」的 圖——這一篇要在上面疊新的線。
步驟 1:MeanFlow 的損失
def meanflow_loss(model, x1, p_same=0.25):
B = x1.shape[0]
x0 = torch.randn_like(x1)
t, r = torch.rand(B), torch.rand(B)
t, r = torch.minimum(t, r), torch.maximum(t, r) # t ≤ r:位置在 t,往資料端跳到 r
r = torch.where(torch.rand(B) < p_same, t, r) # 一部分樣本 r = t:純 flow matching
zt = (1 - t)[:, None] * x0 + t[:, None] * x1
v = x1 - x0 # 條件速度,兩處 v 都用它
u, dudt = torch.func.jvp(
model, (zt, r, t),
(v, torch.zeros_like(r), torch.ones_like(t)), # 切向量 (dz/dt, dr/dt, dt/dt) = (v, 0, 1)
)
target = v - (t - r)[:, None] * dudt # MeanFlow identity 的右邊
return ((u - target.detach()) ** 2).mean() # stop-gradient
逐行對回 U7.2:jvp 一次同時給出 與沿軌跡的全導數 ,切向量的三個分量就是「 以速度 動、 不動、 以 1 動」;target 是 identity 的右邊,detach 是 stop-gradient;v 在兩處(顯式那一項與切向量)都用條件速度,U7.2 Q1 說過這在期望上精確。 那一行是訓練的錨——那些樣本沒有 bootstrapping,就是 U2.5 的 FM 損失。
取樣是定義式反過來:
@torch.no_grad()
def sample(model, n, steps):
z = torch.randn(n, 2)
ts = torch.linspace(0, 1, steps + 1)
for t, r in zip(ts[:-1], ts[1:]): # 每段 t → r
z = z + (r - t) * model(z, r.expand(n), t.expand(n)) # ψ_{t→r}(z) = z + (r−t) u(z, r, t)
return z
steps=1 就是 。
檢查點:(a) 先只用 的樣本訓(p_same=1),它應該表現得和 U2.5 的 FM 一模一樣——這是確認網路與 embedding 沒接錯;(b) 開回 p_same=0.25,loss 不會像 FM 那樣平滑下降,早期會抖(bootstrapping 的目標在動),幾千步後穩定;(c) 一步樣本應該已經落在月牙上,不像 FM 一步那樣糊成一團。
步驟 2:三種一步方法在同一張圖(圖 a)
對步數 ,畫 對步數:
- MeanFlow(本篇);
- consistency training 的一步生成器(U6.6)——多步用「跳到資料端、重新加噪、再跳」;
- reflow×2 加 Euler(U3.6);
- 原始 FM 加 Euler(U2.5)當基準線。
期待看到四種曲線的形狀不同,而且每一種形狀都能用 U7.5 的誤差來源解釋:FM 隨步數一路下降(曲率 + 累積);reflow×2 從較低的起點下降(曲率被壓小);consistency training 一步很好,但多步不見得跟得上 FM 那個斜率——U6.5 說 ODE 的 保證在它身上沒有,實際長什麼樣要量了才知道(U6.6 步驟 6 量過同一條線,把兩張圖放在一起看);MeanFlow 一步與 CT 相近、多步持續改善(多步是合法路徑)。若 MeanFlow 一步明顯輸給 CT,先檢查 p_same 與訓練步數——它的 bootstrapping 需要比 FM 更長的訓練。
課堂提問Q1
步驟 2 那張圖上,四條線的起點(一步)高低不同,斜率(隨步數下降多快)也不同。
這兩件事分別由什麼決定?把它們分開,比直接比誰的線比較低有用得多。
先想一想,再展開看整理後的答案
起點=那個網路一步能做多好,也就是「把整條軌跡壓成一次呼叫」的逼近誤差。FM 的一步最差(它根本不是為一步設計的,一步 Euler 就是假設整條軌跡是直線);reflow×2 好一些(軌跡被拉直過,U3.3);CT 與 MeanFlow 最好(它們學的就是那個一步映射)。這一欄量的是訓練目標對不對。
斜率=多走幾步這件事對這個方法合不合法。 這一欄和起點完全無關,由「學的物件是什麼」決定:
- FM 與 reflow:學的是速度場,多步就是把 ODE 積得更細,U3.2 給 ——所以是一條有斜率的直線。
- MeanFlow:學的是 ,多步是走幾段合法的路,半群條件保證段數變多不會出問題(步驟 3 在量它到底有多成立)。
- CT:學的只有「跳到資料端」。多走一步的唯一辦法是回到資料端再加噪,U6.5 說那個程序沒有 的保證。
所以會看到的圖是:一步的高低看訓練目標,之後的斜率看學的物件撐不撐得住多步。 一個方法可以一步很強、斜率卻是平的(CT),也可以一步不怎麼樣、但斜率很陡(FM)。只報「 的數字」或只報「 的數字」,都會得到相反的結論——這也是讀這一類論文時最容易被帶偏的地方。
(如果你的 MeanFlow 曲線斜率也是平的,先別急著下結論:檢查 p_same 的比例與訓練步數,bootstrapping 沒收斂時多步會先壞掉。)
步驟 3:半群檢查(圖 b)
取 512 個 ,比較「先 再 」與「直接 」的終點距離,再對其他切法(、)重做。畫終點距離的直方圖,並疊上「用 exact_flow_map 算的真實半群誤差」(應該是零到數值精度)。
期待:學到的 的半群誤差不是零——identity 只隱含半群,網路有誤差——但量級遠小於一步樣本與資料的 。這個數字值得記下:它就是 U7.0 條件三在實際模型上的滿足程度。作業 3 的 Shortcut 是把半群當顯式損失,可以比較兩者的半群誤差。
作業
-
DMD 在 toy 上。 學生是一個一步生成器 (MLP)。teacher denoiser 用 U2.5 訓好的 FM 網路換算(FM 慣例下 ,denoiser 給出 score )。fake denoiser
D_fake用 U1.5 那五行訓,資料換成G(z).detach()。生成器的更新照 U7.4 的梯度:x = G(z); t = ...; eps = torch.randn_like(x) zt = (1 - t)[:, None] * eps + t[:, None] * x with torch.no_grad(): g = w(t)[:, None] * (score(D_fake, zt, t) - score(D_real, zt, t)) # ∇_θ KL 的方向 loss_G = (zt * g).sum() / B # 讓 ∂loss/∂zt = g,backward 後 zt 往 s_real − s_fake 走交替更新(每步生成器配 步 fake denoiser,),畫 對訓練步數。回答: 時是否震盪?把 那一項拿掉,樣本跑去哪裡?
-
交叉問題的最終回收。 換到
toy_crossing.py的兩 Gaussian 例子。(a) 用 U6.6 的 consistency training 訓一步生成器,時間網格 ,畫一步樣本,標出落在 等高線外、位於兩個上方 mode 之間的樣本比例;(b) 用作業 1 的 DMD 在同一份資料上訓一步生成器,畫同一張圖。期待:(a) 在 小時大量樣本落在中間——這就是 U2.3 Q1 的交叉、U7.4 Q1 的「平均」—— 大時比例下降但不歸零;(b) 幾乎沒有樣本落在中間。用exact_flow_map對 (a) 的每一顆中間樣本算出它的 真正該去的終點,確認它落在幾個終點的平均位置。寫三句話說明這張圖與 U2.3 那張「條件直線交叉 vs 邊際曲線」的關係。 -
Shortcut。 把步驟 1 改成 U7.3 的 progressive 目標:
model(z, t, d),,四分之三的樣本走 的 FM 損失,其餘走 。不需要jvp。與 MeanFlow 比:一步 、訓練時間、以及步驟 3 的半群誤差——半群當顯式損失後,誤差應該更小;一步品質誰好? -
(選)拿掉 。 把步驟 1 的
target改成只有v(等於對所有 都做 FM 回歸)。訓出來的 學到什麼?一步取樣長什麼樣?用 U7.2 的三角形圖解釋:少了修正向量,網路把弦學成了切線。
參考文獻
- Geng, Z., Deng, M., Bai, X., Kolter, J. Z., He, K. Mean Flows for One-step Generative Modeling. 2025.(步驟 1 的損失與取樣。)
- Yin, T. et al. One-step Diffusion with Distribution Matching Distillation. CVPR 2024.(作業 1 的梯度與交替訓練。)
- Frans, K., Hafner, D., Levine, S., Abbeel, P. One Step Diffusion via Shortcut Models. ICLR 2025.(作業 3。)