U7.6 12 分鐘閱讀 2026年9月

U7.6 實作:MeanFlow、半群檢查,與交叉問題的最終回收

本篇重用M0.3漂流的溫度計:全導數、微分穿過積分與 JVP·M5.3只看得到樣本,怎麼調機器:對 Generator 微分·M6.1兩個城市的人口分佈差多少:W₂ 距離

沿用同一組 toy

資料、網路骨架、exact_score 都沿用 U1.5 建立的 toy.py(two moons,n=2000n=2000)。這一篇用 U2 的 FM 慣例:t=0t=0 噪聲、t=1t=1 資料、zt=(1t)x0+tx1z_t=(1-t)x_0+t\,x_1、條件速度 x1x0x_1-x_0。網路 model(z, r, t):三層 MLP,寬 128,rrtt 各自做 sinusoidal embedding 後串接——比 U2.5 的網路多一個時間輸入,其餘不變。

需要從前幾個單元帶進來的東西:U3.6 的 reflow×2 模型與 W2W_2 函數、U6.6 的 consistency training 一步生成器,以及那一張「三種一步方法」的 W2W_2 圖——這一篇要在上面疊新的線。

步驟 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.2jvp 一次同時給出 uθu_\theta 與沿軌跡的全導數 ddtuθ=tu+(zu)v\tfrac{d}{dt}u_\theta=\partial_t u+(\partial_z u)\,v,切向量的三個分量就是「zz 以速度 vv 動、rr 不動、tt 以 1 動」;target 是 identity 的右邊,detach 是 stop-gradient;v 在兩處(顯式那一項與切向量)都用條件速度,U7.2 Q1 說過這在期望上精確。psamep_{\text{same}} 那一行是訓練的錨——那些樣本沒有 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 就是 z1=z0+uθ(z0,1,0)z_1=z_0+u_\theta(z_0,1,0)

檢查點:(a) 先只用 r=tr=t 的樣本訓(p_same=1),它應該表現得和 U2.5 的 FM 一模一樣——這是確認網路與 embedding 沒接錯;(b) 開回 p_same=0.25,loss 不會像 FM 那樣平滑下降,早期會抖(bootstrapping 的目標在動),幾千步後穩定;(c) 一步樣本應該已經落在月牙上,不像 FM 一步那樣糊成一團。

步驟 2:三種一步方法在同一張圖(圖 a)

對步數 {1,2,4,8,32}\{1,2,4,8,32\},畫 W2W_2 對步數:

  • 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 的 O(1/N)O(1/N) 保證在它身上沒有,實際長什麼樣要量了才知道(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.2O(1/N)O(1/N)——所以是一條有斜率的直線。
  • MeanFlow:學的是 ψts\psi_{t\to s},多步是走幾段合法的路,半群條件保證段數變多不會出問題(步驟 3 在量它到底有多成立)。
  • CT:學的只有「跳到資料端」。多走一步的唯一辦法是回到資料端再加噪,U6.5 說那個程序沒有 O(1/N)O(1/N) 的保證。

所以會看到的圖是:一步的高低看訓練目標,之後的斜率看學的物件撐不撐得住多步。 一個方法可以一步很強、斜率卻是平的(CT),也可以一步不怎麼樣、但斜率很陡(FM)。只報「N=1N=1 的數字」或只報「N=32N=32 的數字」,都會得到相反的結論——這也是讀這一類論文時最容易被帶偏的地方。

(如果你的 MeanFlow 曲線斜率也是平的,先別急著下結論:檢查 p_same 的比例與訓練步數,bootstrapping 沒收斂時多步會先壞掉。)

步驟 3:半群檢查(圖 b)

取 512 個 z0z_0,比較「先 0120\to\tfrac12121\tfrac12\to1」與「直接 010\to1」的終點距離,再對其他切法(01410\to\tfrac14\to103410\to\tfrac34\to1)重做。畫終點距離的直方圖,並疊上「用 exact_flow_map 算的真實半群誤差」(應該是零到數值精度)。

期待:學到的 ψθ\psi_\theta 的半群誤差不是零——identity 只隱含半群,網路有誤差——但量級遠小於一步樣本與資料的 W2W_2。這個數字值得記下:它就是 U7.0 條件三在實際模型上的滿足程度。作業 3 的 Shortcut 是把半群當顯式損失,可以比較兩者的半群誤差。

作業

  1. DMD 在 toy 上。 學生是一個一步生成器 Gθ:R2R2G_\theta:\mathbb R^2\to\mathbb R^2(MLP)。teacher denoiser 用 U2.5 訓好的 FM 網路換算(FM 慣例下 zt=(1t)ϵ+txz_t=(1-t)\epsilon+t\,x,denoiser DE[xzt]D\approx\mathbb E[x\mid z_t] 給出 score s=tD(zt,t)zt(1t)2s=\dfrac{t\,D(z_t,t)-z_t}{(1-t)^2})。fake denoiser D_fakeU1.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 走

    交替更新(每步生成器配 kk 步 fake denoiser,k{1,2,5}k\in\{1,2,5\}),畫 W2W_2 對訓練步數。回答:k=1k=1 時是否震盪?把 sfake-s_{\text{fake}} 那一項拿掉,樣本跑去哪裡?

  2. 交叉問題的最終回收。 換到 toy_crossing.py 的兩 Gaussian 例子。(a) 用 U6.6 的 consistency training 訓一步生成器,時間網格 N{2,4,8,32}N\in\{2,4,8,32\},畫一步樣本,標出落在 p1p_1 等高線外、位於兩個上方 mode 之間的樣本比例;(b) 用作業 1 的 DMD 在同一份資料上訓一步生成器,畫同一張圖。期待:(a) 在 NN 小時大量樣本落在中間——這就是 U2.3 Q1 的交叉、U7.4 Q1 的「平均」——NN 大時比例下降但不歸零;(b) 幾乎沒有樣本落在中間。用 exact_flow_map 對 (a) 的每一顆中間樣本算出它的 zz 真正該去的終點,確認它落在幾個終點的平均位置。寫三句話說明這張圖與 U2.3 那張「條件直線交叉 vs 邊際曲線」的關係。

  3. Shortcut。 把步驟 1 改成 U7.3 的 progressive 目標:model(z, t, d)d{0,18,14,12,1}d\in\{0,\tfrac18,\tfrac14,\tfrac12,1\},四分之三的樣本走 d=0d=0 的 FM 損失,其餘走 s(x,t,2d)=12[s(x,t,d)+s(x,t+d,d)]s(x,t,2d)=\tfrac12[s(x,t,d)+s(x',t+d,d)]。不需要 jvp。與 MeanFlow 比:一步 W2W_2、訓練時間、以及步驟 3 的半群誤差——半群當顯式損失後,誤差應該更小;一步品質誰好?

  4. (選)拿掉 dudt\tfrac{du}{dt} 把步驟 1 的 target 改成只有 v(等於對所有 (r,t)(r,t) 都做 FM 回歸)。訓出來的 uθ(z,r,t)u_\theta(z,r,t) 學到什麼?一步取樣長什麼樣?用 U7.2 的三角形圖解釋:少了修正向量,網路把弦學成了切線。

想一想

步驟 1 的 jvp 切向量是 (v, 0, 1)。若誤寫成 (0, 0, 1),算出來的是什麼?

想一想

步驟 2 的圖上,consistency training 從一步到八步 W2W_2 幾乎不動,MeanFlow 卻持續下降。原因是:

參考文獻

  1. Geng, Z., Deng, M., Bai, X., Kolter, J. Z., He, K. Mean Flows for One-step Generative Modeling. 2025.(步驟 1 的損失與取樣。)
  2. Yin, T. et al. One-step Diffusion with Distribution Matching Distillation. CVPR 2024.(作業 1 的梯度與交替訓練。)
  3. Frans, K., Hafner, D., Levine, S., Abbeel, P. One Step Diffusion via Shortcut Models. ICLR 2025.(作業 3。)