Files
2026-07-13 13:37:02 +08:00

756 lines
39 KiB
HTML
Raw Permalink Blame History

This file contains invisible Unicode characters
This file contains invisible Unicode characters that are indistinguishable to humans but may be processed differently by a computer. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>Flow Matching Quick Reference</title>
<meta name="generator" content="ARIS render-html (academic, v1)">
<meta name="aris:source-path" content="docs/tutorials/flow_matching_tutorial.md">
<meta name="aris:source-sha256" content="ab4358901d204fec676df6ce3b8559be4e69fec3833f764bcfc937b113eea281">
<meta name="aris:generated-at" content="2026-05-19 03:40 UTC">
<!-- MathJax 3 -->
<script>
window.MathJax = {
tex: { inlineMath: [['$', '$'], ['\\(', '\\)']], displayMath: [['$$', '$$'], ['\\[', '\\]']], processEscapes: true },
options: { skipHtmlTags: ['script', 'noscript', 'style', 'textarea', 'pre', 'code'] }
};
</script>
<script src="https://cdn.jsdelivr.net/npm/mathjax@3/es5/tex-mml-chtml.js" async></script>
<!-- highlight.js -->
<link rel="stylesheet" href="https://cdn.jsdelivr.net/gh/highlightjs/cdn-release@11.9.0/build/styles/atom-one-light.min.css">
<script src="https://cdn.jsdelivr.net/gh/highlightjs/cdn-release@11.9.0/build/highlight.min.js"></script>
<script>document.addEventListener('DOMContentLoaded', () => hljs.highlightAll());</script>
<style>
:root {
--bg: #fdfcf7;
--bg-soft: #f4f1ea;
--bg-code: #f8f5ec;
--ink: #1a1a1a;
--ink-soft: #4a4a4a;
--ink-muted: #6b6b6b;
--primary: #1a4a8c;
--primary-soft: #2d6cb8;
--accent: #b8390e;
--warn: #b45309;
--warn-bg: #fef3c7;
--info-bg: #dbeafe;
--good-bg: #d1fae5;
--good: #065f46;
--bad-bg: #fee2e2;
--bad: #991b1b;
--border: #d6d0c0;
--border-soft: #e8e3d5;
}
* { box-sizing: border-box; }
html { scroll-behavior: smooth; }
body {
font-family: "Source Serif Pro", "Source Serif 4", "Crimson Pro", "Georgia", "Songti SC", "STSong", serif;
line-height: 1.65;
color: var(--ink);
background: var(--bg);
margin: 0;
padding: 0;
font-size: 16px;
}
.layout {
max-width: 1280px;
margin: 0 auto;
display: grid;
grid-template-columns: 260px 1fr;
gap: 48px;
padding: 40px 32px;
}
nav.toc {
position: sticky;
top: 24px;
align-self: start;
font-size: 13px;
max-height: calc(100vh - 48px);
overflow-y: auto;
border-right: 1px solid var(--border-soft);
padding-right: 16px;
}
nav.toc h3 {
margin: 0 0 12px;
font-size: 12px;
text-transform: uppercase;
letter-spacing: 0.08em;
color: var(--ink-muted);
font-weight: 600;
}
nav.toc ol { list-style: none; padding: 0; margin: 0; counter-reset: toc; }
nav.toc ol li { margin: 5px 0; counter-increment: toc; }
nav.toc ol li::before { content: counter(toc) ". "; color: var(--ink-muted); margin-right: 4px; }
nav.toc a {
color: var(--ink-soft);
text-decoration: none;
border-bottom: 1px dotted transparent;
}
nav.toc a:hover { color: var(--primary); border-bottom-color: var(--primary); }
nav.toc ul { list-style: none; padding-left: 14px; margin: 3px 0; font-size: 12px; }
nav.toc ul li::before { content: "→ "; color: var(--border); }
main { min-width: 0; }
header.hero {
border-bottom: 3px double var(--primary);
padding-bottom: 24px;
margin-bottom: 32px;
}
header.hero .eyebrow {
color: var(--accent);
font-size: 13px;
text-transform: uppercase;
letter-spacing: 0.12em;
font-weight: 600;
margin-bottom: 8px;
}
header.hero h1 {
font-size: 32px;
line-height: 1.2;
margin: 0 0 12px;
color: var(--ink);
font-weight: 700;
letter-spacing: -0.01em;
}
header.hero .subtitle {
font-size: 16px;
color: var(--ink-soft);
margin: 0 0 8px;
font-style: italic;
}
header.hero .byline {
font-size: 14px;
color: var(--ink-soft);
margin: 0 0 20px;
}
header.hero .byline strong {
color: var(--ink);
font-weight: 600;
}
header.hero .meta {
display: flex;
gap: 20px;
flex-wrap: wrap;
font-size: 12px;
color: var(--ink-muted);
border-top: 1px solid var(--border-soft);
padding-top: 14px;
}
header.hero .meta span strong { color: var(--ink-soft); }
header.hero .meta code {
font-family: "JetBrains Mono", "SF Mono", "Menlo", "Consolas", monospace;
font-size: 11px;
background: var(--bg-soft);
padding: 1px 5px;
border-radius: 3px;
border: 1px solid var(--border-soft);
}
h2 {
font-size: 24px;
margin: 44px 0 14px;
padding-bottom: 8px;
border-bottom: 1px solid var(--border);
color: var(--ink);
font-weight: 700;
}
h2 .num { color: var(--primary); font-weight: 600; margin-right: 8px; }
h3 { font-size: 19px; margin: 28px 0 10px; color: var(--primary); font-weight: 600; }
h4 { font-size: 16px; margin: 20px 0 8px; color: var(--ink); font-weight: 600; }
p { margin: 10px 0; }
ul, ol { padding-left: 22px; margin: 10px 0; }
ul li, ol li { margin: 4px 0; }
ul li::marker { color: var(--primary); }
strong { color: var(--accent); font-weight: 600; }
em { color: var(--ink-soft); }
a { color: var(--primary); }
a:hover { color: var(--accent); }
code:not(.hljs) {
font-family: "JetBrains Mono", "SF Mono", "Menlo", "Consolas", monospace;
font-size: 0.86em;
background: var(--bg-code);
padding: 1px 5px;
border-radius: 3px;
border: 1px solid var(--border-soft);
color: var(--accent);
}
pre {
background: #fafaf6;
border: 1px solid var(--border);
border-left: 4px solid var(--primary);
padding: 0;
overflow-x: auto;
border-radius: 4px;
margin: 14px 0;
}
pre code, pre code.hljs {
background: transparent !important;
display: block;
padding: 14px 18px !important;
font-size: 13px;
line-height: 1.55;
font-family: "JetBrains Mono", "SF Mono", "Menlo", monospace;
color: var(--ink);
}
pre.diagram {
background: #f9f6ed;
border-left: 4px solid var(--accent);
font-size: 12.5px;
line-height: 1.4;
}
.callout {
margin: 16px 0;
padding: 12px 16px;
border-radius: 4px;
border-left: 4px solid;
font-size: 15px;
}
.callout-title {
font-weight: 600;
margin-bottom: 6px;
font-size: 12px;
text-transform: uppercase;
letter-spacing: 0.06em;
}
.callout-info { background: var(--info-bg); border-left-color: var(--primary); }
.callout-info .callout-title { color: var(--primary); }
.callout-warn { background: var(--warn-bg); border-left-color: var(--warn); }
.callout-warn .callout-title { color: var(--warn); }
.callout-good { background: var(--good-bg); border-left-color: var(--good); }
.callout-good .callout-title { color: var(--good); }
.callout-bad { background: var(--bad-bg); border-left-color: var(--bad); }
.callout-bad .callout-title { color: var(--bad); }
table {
width: 100%;
border-collapse: collapse;
margin: 16px 0;
font-size: 14px;
border: 1px solid var(--border);
border-radius: 4px;
overflow: hidden;
}
thead { background: var(--primary); color: white; }
th, td {
text-align: left;
padding: 9px 12px;
border-bottom: 1px solid var(--border-soft);
vertical-align: top;
}
th { font-weight: 600; font-size: 13px; letter-spacing: 0.02em; }
tr:last-child td { border-bottom: none; }
tbody tr:nth-child(even) { background: var(--bg-soft); }
details.qa, details {
background: white;
border: 1px solid var(--border-soft);
border-radius: 6px;
margin: 10px 0;
padding: 0;
}
details summary {
cursor: pointer;
padding: 10px 14px;
font-weight: 600;
font-size: 14px;
color: var(--primary);
list-style: none;
user-select: none;
}
details summary::-webkit-details-marker { display: none; }
details summary::before {
content: "▸ ";
margin-right: 4px;
display: inline-block;
transition: transform 0.15s;
}
details[open] summary::before { transform: rotate(90deg); }
details[open] summary { border-bottom: 1px solid var(--border-soft); }
details > :not(summary) { padding: 10px 14px; }
details p:first-of-type { margin-top: 8px; }
mjx-container[display="true"] { margin: 12px 0 !important; }
footer.aris-footer {
margin-top: 60px;
padding-top: 20px;
border-top: 1px solid var(--border);
font-size: 12px;
color: var(--ink-muted);
}
footer.aris-footer a { color: var(--ink-muted); border-bottom: 1px dotted var(--border); }
@media (max-width: 900px) {
.layout { grid-template-columns: 1fr; gap: 20px; padding: 20px 16px; }
nav.toc {
position: static;
max-height: none;
border-right: none;
border-bottom: 1px solid var(--border-soft);
padding-right: 0;
padding-bottom: 14px;
}
header.hero h1 { font-size: 24px; }
h2 { font-size: 20px; }
}
@media print {
nav.toc { display: none; }
.layout { grid-template-columns: 1fr; padding: 0; }
body { background: white; }
header.hero { border-bottom-color: var(--ink); }
}
</style>
</head>
<body>
<div class="layout">
<nav class="toc">
<h3>Contents</h3>
<ol>
<li><a href="#0-tldr">§0 TL;DR</a>
</li>
<li><a href="#1-基本设定--直觉">§1 基本设定 &amp; 直觉</a>
</li>
<li><a href="#2-flow-matching-loss">§2 Flow Matching Loss</a>
<ul>
<li><a href="#21-marginal-flow-matching理论上">2.1 Marginal Flow Matching(理论上)</a></li>
<li><a href="#22-conditional-flow-matching实际可训练">2.2 Conditional Flow Matching(实际可训练)</a></li>
<li><a href="#23-关键定理lipman-et-al-2023-theorem-2">2.3 关键定理(Lipman et al. 2023, Theorem 2</a></li>
</ul>
</li>
<li><a href="#3-三种-conditional-path-选择">§3 三种 Conditional Path 选择</a>
<ul>
<li><a href="#31-rectified-flow最简最稳最常用">3.1 Rectified Flow:最简、最稳、最常用</a></li>
<li><a href="#32-vp-path与-ddpm-同族">3.2 VP path(与 DDPM 同族)</a></li>
<li><a href="#33-ve-path与-smldedm-同族">3.3 VE path(与 SMLD/EDM 同族)</a></li>
</ul>
</li>
<li><a href="#4-训练代码框架pytorch">§4 训练代码框架(PyTorch</a>
<ul>
<li><a href="#41-probability-path-抽象">4.1 Probability Path 抽象</a></li>
<li><a href="#42-vector-field-网络教学版-mlp实际用-u-net--dit">4.2 Vector Field 网络(教学版 MLP;实际用 U-Net / DiT</a></li>
<li><a href="#43-cfm-loss">4.3 CFM Loss</a></li>
<li><a href="#44-minimal-training-loop">4.4 Minimal Training Loop</a></li>
</ul>
</li>
<li><a href="#5-ode-采样">§5 ODE 采样</a>
</li>
<li><a href="#6-与-diffusion--score-matching-的关系">§6 与 Diffusion / Score Matching 的关系</a>
<ul>
<li><a href="#61-velocity--score--noise-prediction-互换必考">6.1 Velocity ↔ Score ↔ Noise prediction 互换(必考)</a></li>
<li><a href="#62-几种-path-与-diffusion-的对应表">6.2 几种 path 与 diffusion 的对应表</a></li>
<li><a href="#63-为什么-rectified-flow-训练--采样相对稳">6.3 为什么 Rectified Flow 训练 / 采样相对&quot;&quot;</a></li>
</ul>
</li>
<li><a href="#7-高级话题">§7 高级话题</a>
<ul>
<li><a href="#71-reflowliu-et-al-2022-iclr">7.1 ReflowLiu et al. 2022, ICLR</a></li>
<li><a href="#72-conditional-flow-matching-cfg">7.2 Conditional Flow Matching (CFG)</a></li>
<li><a href="#73-logit-normal-tsd3-默认">7.3 Logit-normal $t$SD3 默认)</a></li>
</ul>
</li>
<li><a href="#8-完整可运行示例2d-toy">§8 完整可运行示例(2D toy</a>
</li>
</ol>
</nav>
<main>
<header class="hero">
<div class="eyebrow">Interview Prep · Flow Matching / Diffusion Generative Modeling</div>
<h1>Flow Matching Quick Reference</h1>
<p class="subtitle">Conditional Flow Matching · Rectified Flow / VP / VE · 训练 + 采样 + SD3 / FLUX 实战</p>
<p class="byline">By <strong>Ruofeng Yang (杨若峰), Shanghai Jiao Tong University</strong></p>
<div class="meta">
<span><strong>Source:</strong> <code>docs/tutorials/flow_matching_tutorial.md</code></span>
<span><strong>SHA256:</strong> <code>ab4358901d20</code></span>
<span><strong>Rendered:</strong> 2026-05-19 03:40 UTC</span>
</div>
</header>
<h2 id="0-tldr">§0 TL;DR</h2>
<div class="callout callout-info"><div class="callout-title">5 句话搞定 Flow Matching</div><p>一页拿下核心要点(详见后文 §1–§4 推导)。</p></div>
<ol><li><strong>目标</strong>:学一个 vector field $v_\theta(t, x)$,使 ODE $\dot{x}_t = v_\theta(t, x_t)$ 把 $x_0 \sim p_0$(噪声)演化到 $x_1 \sim p_1$(数据)。</li><li><strong>训练 (CFM)</strong>$\mathcal{L}_\text{CFM}(\theta) = \mathbb{E}_{t, z, x_t \sim p_t(\cdot|z)} \|v_\theta(t, x_t) - u_t(x_t|z)\|^2$<strong>simulation-free</strong>(不用解 ODE 算 loss)。</li><li><strong>关键定理</strong>$\nabla_\theta \mathcal{L}_\text{FM} = \nabla_\theta \mathcal{L}_\text{CFM}$——所以学 conditional vector field 等价于学 marginal 的(Lipman et al. 2023)。</li><li><strong>最简版 (Rectified Flow / OT-CFM)</strong>$x_t = (1-t)x_0 + tx_1$target $u_t = x_1 - x_0$。SD3 / FLUX / Lumina 都用这个。</li><li><strong>采样</strong>:从 $x_0 \sim p_0$ 出发,用 ODE solver (Euler / Heun / RK4) 积分到 $t=1$。</li></ol>
<h2 id="1-基本设定--直觉">§1 基本设定 &amp; 直觉</h2>
<p>给定数据分布 $p_1$("目标")和简单先验 $p_0$(一般 $\mathcal{N}(0, I)$),我们想构造一族<strong>概率路径</strong> $\{p_t\}_{t \in [0,1]}$ 从 $p_0$ 平滑过渡到 $p_1$。</p>
<div class="callout callout-warn"><div class="callout-title">Convention(全文统一)</div><p>后续公式记号见下表。</p></div>
<ul><li>$x_0 \sim p_0 = \mathcal{N}(0, I)$(噪声端)—— $t=0$</li><li>$x_1 \sim p_1$(数据端)—— $t=1$</li><li>采样方向:从 $t=0$ 积分到 $t=1$noise → data</li><li>注:不同论文 convention 不同——Lipman et al. 2023 用 $x_0$=data, $x_1$=noiseLiu et al. 2022 (Rectified Flow) 用 $x_0$=noise, $x_1$=data(本文采用)。SD3 论文也是 noise→data 但记号略有差异。<strong>面试时第一句先 disambiguate</strong></li></ul>
<p>一族<strong>时变 vector field</strong> $u_t : [0,1] \times \mathbb{R}^d \to \mathbb{R}^d$ 通过 ODE $\dot{x}_t = u_t(x_t)$ 把粒子从 $p_0$ 推到 $p_1$。由<strong>连续性方程</strong></p>
<p>$$\boxed{\;\frac{\partial p_t}{\partial t} + \nabla \cdot (p_t\, u_t) = 0\;}$$</p>
<p>我们的目标:找一个神经网络 $v_\theta(t, x) \approx u_t(x)$。</p>
<pre class="diagram"><code>
p_0 (噪声) p_t (中间) p_1 (数据)
●●●●● → ● ● ● → ████
v_θ(t, x)
─────────→
dx/dt = v_θ</code></pre>
<p>对比 diffusion</p>
<ul><li><strong>Diffusion (SDE)</strong>$dx = f(x, t) dt + g(t) dW$,训练用 score matching $s_\theta \approx \nabla \log p_t$</li><li><strong>Flow matching (ODE)</strong>$dx = v_\theta(t, x) dt$<strong>无随机项</strong>,训练直接回归 vector field</li><li>两者通过 <strong>probability flow ODE</strong> 关联:$v = f - \frac{1}{2} g^2 \nabla \log p_t$(见 §6</li></ul>
<h2 id="2-flow-matching-loss">§2 Flow Matching Loss</h2>
<h3 id="21-marginal-flow-matching理论上">2.1 Marginal Flow Matching(理论上)</h3>
<p>如果我们知道 $u_t$marginal vector field),直接回归即可:</p>
<p>$$\mathcal{L}_\text{FM}(\theta) = \mathbb{E}_{t \sim \mathcal{U}[0,1],\; x \sim p_t} \left\| v_\theta(t, x) - u_t(x) \right\|^2$$</p>
<p><strong>问题</strong>$u_t(x)$ 是 marginal,由所有 conditional path 加权(积分),<strong>无法直接采样</strong></p>
<h3 id="22-conditional-flow-matching实际可训练">2.2 Conditional Flow Matching(实际可训练)</h3>
<p>引入 <strong>conditioning variable</strong> $z$(如 $z = x_1$,或 $z = (x_0, x_1)$)。选 conditional path $p_t(x | z)$ 和 conditional vector field $u_t(x | z)$,使得对 $z$ 边缘后等于 marginal</p>
<p>$$p_t(x) = \int p_t(x | z) q(z)\, dz, \quad u_t(x) = \int u_t(x|z) \frac{p_t(x|z) q(z)}{p_t(x)} dz$$</p>
<p>那么 <strong>Conditional FM loss</strong></p>
<p>$$\boxed{\;\mathcal{L}_\text{CFM}(\theta) = \mathbb{E}_{t,\; z \sim q,\; x \sim p_t(\cdot|z)} \left\| v_\theta(t, x) - u_t(x|z) \right\|^2\;}$$</p>
<p>每一项都<strong>可采样、可计算</strong>。$x \sim p_t(\cdot|z)$ 一般是闭式可采样(如下面的线性插值)。</p>
<h3 id="23-关键定理lipman-et-al-2023-theorem-2">2.3 关键定理(Lipman et al. 2023, Theorem 2</h3>
<div class="callout callout-good"><div class="callout-title">梯度等价定理</div><p>在 $p_t$ 与 $u_t$ 适当正则、$p_t > 0$ 等假设下:</p></div>
<p>$$\nabla_\theta \mathcal{L}_\text{FM}(\theta) = \nabla_\theta \mathcal{L}_\text{CFM}(\theta)$$ 所以 <strong>minimize CFM ≡ minimize FM</strong>。两个 loss 在以上前提下相差一个不依赖 $\theta$ 的常数。</p>
<p><strong>证明草图</strong>:展开二范数 $\|v_\theta\|^2 - 2 v_\theta^\top u_t + \|u_t\|^2$,前两项在两种 loss 下相等(用 $u_t$ 的定义把 marginal 写成 conditional 的加权期望),第三项不依赖 $\theta$,求梯度时消掉。</p>
<div class="callout callout-info"><div class="callout-title">面试加分:marginal vector field 非唯一</div><p>给定 $p_t$,满足连续性方程 $\partial_t p_t + \nabla\cdot(p_t u_t) = 0$ 的 $u_t$ <strong>不唯一</strong>——任意 divergence-free 向量场加上去仍是合法的。CFM 通过 conditional path 自动选了一个"自然的" $u_t$(通常对应 OT map 或 score-based ODE)。这点常被追问 "marginal $u_t$ 唯一吗?"。</p></div>
<h2 id="3-三种-conditional-path-选择">§3 三种 Conditional Path 选择</h2>
<p>conditioning 取 $z = (x_0, x_1)$$x_0 \sim p_0$, $x_1 \sim p_1$。conditional path $p_t(x | x_0, x_1)$ 一般是 Dirac $\delta(x - \psi_t(x_0, x_1))$(确定性插值),conditional vector field 为 $\dot{\psi}_t(x_0, x_1)$。</p>
<table><thead><tr><th>Path</th><th>$x_t = \psi_t(x_0, x_1)$</th><th>Target $u_t$</th><th>用在哪</th></tr></thead><tbody><tr><td><strong>Rectified Flow / OT-CFM</strong></td><td>$(1-t)x_0 + t\, x_1$</td><td>$x_1 - x_0$(常数)</td><td>SD3, FLUX, Lumina, MovieGen</td></tr><tr><td><strong>VP cosine</strong></td><td>$\cos\!\left(\frac{\pi t}{2}\right) x_0 + \sin\!\left(\frac{\pi t}{2}\right) x_1$</td><td>$-\frac{\pi}{2}\sin\!\frac{\pi t}{2}\, x_0 + \frac{\pi}{2}\cos\!\frac{\pi t}{2}\, x_1$</td><td>与 DDPM cosine schedule 同族(在限制下)</td></tr><tr><td><strong>VE</strong></td><td>$x_1 + \sigma(1{-}t)\, x_0$$\sigma$ 递增</td><td>$-\sigma'(1{-}t)\, x_0$</td><td>与 SMLD/EDM 同族(需 prior 方差匹配 $\sigma_{\max}^2$</td></tr></tbody></table>
<h3 id="31-rectified-flow最简最稳最常用">3.1 Rectified Flow:最简、最稳、最常用</h3>
<p>线性插值:$x_t = (1-t) x_0 + t\, x_1$,所以 $\dot{x}_t = x_1 - x_0$ 是<strong>常数</strong>(不依赖 $t$)。</p>
<p>训练目标:</p>
<p>$$\mathcal{L}_\text{RF}(\theta) = \mathbb{E}_{t, x_0, x_1} \|v_\theta(t,\, (1-t)x_0 + t x_1) - (x_1 - x_0)\|^2$$</p>
<p>"OT-CFM" 名字来自:如果 $(x_0, x_1)$ 是 optimal transport coupling(不是独立采样),则学到的 vector field 实现近似 OT map。</p>
<div class="callout callout-good"><div class="callout-title">ReflowRectified Flow 的杀手锏</div><p>用学到的 $v_\theta$ 重新生成 $(x_0, x_1)$ pair(用 $x_0$ 跑 ODE 得到对应 $x_1$),<strong>再训一次</strong>。新的 trajectory 更直,<strong>少步数采样质量大幅提升</strong>,可以做到 1-step / 2-step 生成(InstaFlow 等)。</p></div>
<h3 id="32-vp-path与-ddpm-同族">3.2 VP path(与 DDPM 同族)</h3>
<p>用 $\sigma(t) = \cos\!\frac{\pi t}{2}$(噪声系数), $\alpha(t) = \sin\!\frac{\pi t}{2}$(数据系数),满足 $\sigma^2 + \alpha^2 = 1$variance preserving):</p>
<p>$$x_t = \sigma(t)\, x_0 + \alpha(t)\, x_1, \quad u_t = \sigma'(t)\, x_0 + \alpha'(t)\, x_1$$</p>
<p>边界:$t=0$ 时 $x_t = x_0$(噪声),$t=1$ 时 $x_t = x_1$(数据)。</p>
<p>这种 path 和 DDPM 的 cosine schedule <strong>属于同一族 Gaussian path</strong>(连续极限 + 时间反向)。但严格意义上不能说"完全等价"——DDPM 原文 (Nichol-Dhariwal) 还有 $s=0.008$ offset 等细节,且 DDPM 是 forward noising convention$t=0$ 数据),FM 是反向($t=0$ 噪声)。</p>
<h3 id="33-ve-path与-smldedm-同族">3.3 VE path(与 SMLD/EDM 同族)</h3>
<p>沿用 Lipman et al. 2023 的 conditional VE path</p>
<p>$$p_t(x | x_1) = \mathcal{N}\!\left(x \,\Big|\, x_1,\; \sigma(1-t)^2 I\right)$$</p>
<p>$\sigma(s)$ 在 forward time $s \in [0, 1]$ 上单调递增(如 $\sigma(s) = \sigma_\min (\sigma_\max/\sigma_\min)^s$)。reparameterize 得</p>
<p>$$x_t = x_1 + \sigma(1-t)\, x_0, \quad u_t = -\sigma'(1-t)\, x_0$$</p>
<p>边界:$t=0$ 时 $x_t \approx x_1 + \sigma_\max\, x_0$(大噪声主导),$t=1$ 时 $x_t \approx x_1 + \sigma_\min\, x_0 \approx x_1$(数据)。</p>
<div class="callout callout-warn"><div class="callout-title">VE 部署注意</div><p>prior $p_0$ 严格意义上应是 $\mathcal{N}(0, \sigma_\max^2 I)$(让 $t=0$ 的边缘方差匹配);用 $\mathcal{N}(0, I)$ 时要相应缩放(如 $x_0 \leftarrow \sigma_\max \cdot \tilde{x}_0$)。本文代码示例以教学为主,<strong>实战 VE 用 EDM preconditioning</strong> 更稳。</p></div>
<h2 id="4-训练代码框架pytorch">§4 训练代码框架(PyTorch</h2>
<h3 id="41-probability-path-抽象">4.1 Probability Path 抽象</h3>
<pre><code class="language-python">import math
from dataclasses import dataclass
from typing import Callable, Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
@dataclass
class FlowPath:
&quot;&quot;&quot; Conditional probability path 抽象 &quot;&quot;&quot;
name: str
sample_xt: Callable # (t, x0, x1) -&gt; x_t
target_ut: Callable # (t, x0, x1) -&gt; u_t
def _broadcast_t(t: torch.Tensor, x: torch.Tensor) -&gt; torch.Tensor:
&quot;&quot;&quot; t: [B], x: [B, ...] —— 把 t 扩成可广播 x 的 shape [B, 1, 1, ...] &quot;&quot;&quot;
return t.view(-1, *([1] * (x.dim() - 1)))
def rectified_flow_path() -&gt; FlowPath:
&quot;&quot;&quot; x_t = (1-t)x_0 + t*x_1, u_t = x_1 - x_0 &quot;&quot;&quot;
def sample_xt(t, x0, x1):
tb = _broadcast_t(t, x0)
return (1 - tb) * x0 + tb * x1
def target_ut(t, x0, x1):
return x1 - x0
return FlowPath(&quot;rectified_flow&quot;, sample_xt, target_ut)
def vp_cosine_path() -&gt; FlowPath:
&quot;&quot;&quot; x_t = cos(π t/2) x_0 + sin(π t/2) x_1
t=0: x_t = x_0 (noise); t=1: x_t = x_1 (data) [noise → data 方向] &quot;&quot;&quot;
def sample_xt(t, x0, x1):
tb = _broadcast_t(t, x0)
sig = torch.cos(0.5 * math.pi * tb) # noise coeff
alp = torch.sin(0.5 * math.pi * tb) # data coeff
return sig * x0 + alp * x1
def target_ut(t, x0, x1):
tb = _broadcast_t(t, x0)
d_sig = -0.5 * math.pi * torch.sin(0.5 * math.pi * tb)
d_alp = 0.5 * math.pi * torch.cos(0.5 * math.pi * tb)
return d_sig * x0 + d_alp * x1
return FlowPath(&quot;vp_cosine&quot;, sample_xt, target_ut)
def ve_path(sigma_min: float = 0.01, sigma_max: float = 50.0) -&gt; FlowPath:
&quot;&quot;&quot; VE: x_t = x_1 + σ(1-t) · x_0, σ(s) 在 forward time s 上递增 (log-linear)
t=0: x_t = x_1 + σ_max·x_0 (大噪声); t=1: x_t ≈ x_1 (数据)
注: 严格 VE 需 prior p_0 ~ N(0, σ_max² I),本示例为简化用 N(0, I),
实战需要 EDM-style preconditioning。 &quot;&quot;&quot;
log_min, log_max = math.log(sigma_min), math.log(sigma_max)
def sigma_fwd(s): # 在 forward time s 上递增
return torch.exp(log_min * (1 - s) + log_max * s)
def d_sigma_fwd(s): # dσ/ds = σ · (log σ_max log σ_min)
return sigma_fwd(s) * (log_max - log_min)
def sample_xt(t, x0, x1):
tb = _broadcast_t(t, x0)
return x1 + sigma_fwd(1 - tb) * x0
def target_ut(t, x0, x1):
tb = _broadcast_t(t, x0)
# u_t = d/dt [σ(1-t)] x_0 = -σ&#x27;(1-t) · x_0
return -d_sigma_fwd(1 - tb) * x0
return FlowPath(&quot;ve&quot;, sample_xt, target_ut)</code></pre>
<h3 id="42-vector-field-网络教学版-mlp实际用-u-net--dit">4.2 Vector Field 网络(教学版 MLP;实际用 U-Net / DiT</h3>
<pre><code class="language-python">class SinusoidalTimeEmbed(nn.Module):
&quot;&quot;&quot; 与 Transformer positional embedding 同构的时间编码 &quot;&quot;&quot;
def __init__(self, dim: int):
super().__init__()
self.dim = dim
def forward(self, t: torch.Tensor) -&gt; torch.Tensor:
# t: [B] in [0, 1]
half = self.dim // 2
freqs = torch.exp(-math.log(10000) * torch.arange(half, device=t.device) / half)
args = t[:, None] * freqs[None, :]
return torch.cat([torch.sin(args), torch.cos(args)], dim=-1)
class VectorFieldMLP(nn.Module):
&quot;&quot;&quot; v_θ(t, x) ——— 简化版,用于 2D toy / low-dim 实验
实际生成模型把这里换成 U-Net (image) 或 DiT (high-res / video) &quot;&quot;&quot;
def __init__(self, dim: int, hidden: int = 256, t_dim: int = 128):
super().__init__()
self.t_embed = nn.Sequential(
SinusoidalTimeEmbed(t_dim),
nn.Linear(t_dim, hidden),
nn.SiLU(),
nn.Linear(hidden, hidden),
)
self.net = nn.Sequential(
nn.Linear(dim + hidden, hidden), nn.SiLU(),
nn.Linear(hidden, hidden), nn.SiLU(),
nn.Linear(hidden, dim),
)
def forward(self, t: torch.Tensor, x: torch.Tensor) -&gt; torch.Tensor:
# t: [B], x: [B, dim]
return self.net(torch.cat([x, self.t_embed(t)], dim=-1))</code></pre>
<h3 id="43-cfm-loss">4.3 CFM Loss</h3>
<pre><code class="language-python">def cfm_loss(
model: nn.Module,
path: FlowPath,
x1: torch.Tensor, # [B, ...] 数据样本
x0: Optional[torch.Tensor] = None, # 默认 N(0, I)
t_dist: str = &quot;uniform&quot;, # &quot;uniform&quot; or &quot;logitnormal&quot;
return_components: bool = False,
):
&quot;&quot;&quot;
Conditional Flow Matching loss:
L = E ‖v_θ(t, x_t) - u_t(x_t | x_0, x_1)‖²
&quot;&quot;&quot;
B = x1.shape[0]
device = x1.device
if x0 is None:
x0 = torch.randn_like(x1)
# t 采样
if t_dist == &quot;uniform&quot;:
t = torch.rand(B, device=device)
elif t_dist == &quot;logitnormal&quot;:
# SD3 默认: t = σ(z), z ~ N(0, 1). 更集中在 t≈0.5(最难学的中间区)
t = torch.sigmoid(torch.randn(B, device=device))
else:
raise ValueError(f&quot;unknown t_dist: {t_dist}&quot;)
x_t = path.sample_xt(t, x0, x1)
u_t = path.target_ut(t, x0, x1)
v_pred = model(t, x_t)
loss = F.mse_loss(v_pred, u_t)
if return_components:
return loss, {&quot;v_pred_norm&quot;: v_pred.norm().item(), &quot;u_norm&quot;: u_t.norm().item()}
return loss</code></pre>
<h3 id="44-minimal-training-loop">4.4 Minimal Training Loop</h3>
<pre><code class="language-python">def train_flow_matching(
model: nn.Module,
dataloader, # yields x1 batches
path: FlowPath,
total_steps: int = 50_000,
lr: float = 3e-4,
weight_decay: float = 0.0,
device: str = &quot;cuda&quot;,
log_every: int = 200,
ema_decay: float = 0.9999, # EMA 对生成模型必加
):
model = model.to(device).train()
opt = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)
ema_model = _make_ema(model) # 见下
step = 0
while step &lt; total_steps:
for x1 in dataloader:
x1 = x1.to(device, non_blocking=True)
loss = cfm_loss(model, path, x1, t_dist=&quot;logitnormal&quot;)
opt.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step()
_update_ema(ema_model, model, ema_decay)
if step % log_every == 0:
print(f&quot;[{step:6d}] {path.name} loss = {loss.item():.4f}&quot;)
step += 1
if step &gt;= total_steps: break
return model, ema_model
@torch.no_grad()
def _make_ema(model):
import copy
ema = copy.deepcopy(model).eval()
for p in ema.parameters(): p.requires_grad_(False)
return ema
@torch.no_grad()
def _update_ema(ema, model, decay):
for ep, p in zip(ema.parameters(), model.parameters()):
ep.mul_(decay).add_(p.detach(), alpha=1 - decay)</code></pre>
<h2 id="5-ode-采样">§5 ODE 采样</h2>
<p>训练完 $v_\theta$ 后,从 $x_0 \sim p_0$ 出发,解 ODE $\dot{x}_t = v_\theta(t, x_t)$ 到 $t = 1$。</p>
<pre><code class="language-python">@torch.no_grad()
def euler_sampler(model, x0, steps=50, t_start=0.0, t_end=1.0):
&quot;&quot;&quot; 一阶 Euler:每步 1 NFE,简单但需要较多步数 &quot;&quot;&quot;
x = x0.clone()
ts = torch.linspace(t_start, t_end, steps + 1, device=x0.device)
for i in range(steps):
t = ts[i].expand(x.shape[0])
dt = ts[i + 1] - ts[i]
x = x + dt * model(t, x)
return x
@torch.no_grad()
def heun_sampler(model, x0, steps=50, t_start=0.0, t_end=1.0):
&quot;&quot;&quot; 二阶 Heun (improved Euler / RK2):每步 2 NFE,精度 O(dt²) &quot;&quot;&quot;
x = x0.clone()
ts = torch.linspace(t_start, t_end, steps + 1, device=x0.device)
for i in range(steps):
b = x.shape[0]
t_i, t_next = ts[i], ts[i + 1]
dt = t_next - t_i
v1 = model(t_i.expand(b), x)
x_euler = x + dt * v1
v2 = model(t_next.expand(b), x_euler)
x = x + dt * 0.5 * (v1 + v2)
return x
@torch.no_grad()
def rk4_sampler(model, x0, steps=25, t_start=0.0, t_end=1.0):
&quot;&quot;&quot; 四阶 Runge-Kutta:每步 4 NFE,精度 O(dt⁴)
25 步 × 4 NFE = 100 NFE,但通常比 100 步 Euler 准很多 &quot;&quot;&quot;
x = x0.clone()
ts = torch.linspace(t_start, t_end, steps + 1, device=x0.device)
for i in range(steps):
b = x.shape[0]
t_i, t_next = ts[i], ts[i + 1]
dt = t_next - t_i
k1 = model(t_i.expand(b), x)
k2 = model((t_i + dt / 2).expand(b), x + dt / 2 * k1)
k3 = model((t_i + dt / 2).expand(b), x + dt / 2 * k2)
k4 = model(t_next.expand(b), x + dt * k3)
x = x + dt / 6 * (k1 + 2 * k2 + 2 * k3 + k4)
return x</code></pre>
<div class="callout callout-info"><div class="callout-title">Sampler 选择 cheat sheet</div><p>按 NFE / 质量 trade-off 排序如下。</p></div>
<ul><li><strong>Euler</strong>1 NFE/step,需 ≥50 步才出好图;调试 baseline 用</li><li><strong>Heun / RK2</strong>2 NFE/step~25 步质量已不错;EDM 默认</li><li><strong>RK4</strong>4 NFE/step10-20 步通常等同 Euler 100 步</li><li><strong>Adaptive (Dopri5 / dopri8)</strong>torchdiffeq 提供;自动控制误差,但 NFE 不可控</li><li><strong>Rectified Flow 重训后</strong>:经 1-2 次 reflow1-4 步 Euler 即可达到接近多步质量</li></ul>
<h2 id="6-与-diffusion--score-matching-的关系">§6 与 Diffusion / Score Matching 的关系</h2>
<p>对任意 SDE $dx = f(x, t) dt + g(t) dW$forward),存在对应的 <strong>probability flow ODE</strong>Song et al. 2021):</p>
<p>$$dx = \underbrace{\left[ f(x, t) - \frac{1}{2} g^2(t)\, \nabla_x \log p_t(x) \right]}_{\text{vector field } u_t(x)} dt$$</p>
<p>这个 ODE 在每一时刻的边缘分布 $p_t$ 与 SDE 完全一样。</p>
<div class="callout callout-good"><div class="callout-title">FM ↔ Score Matching 的桥梁(带前提)</div><p>当 FM 的概率路径源自一个非退化的 noising SDE ($g(t) > 0$) 时,学 score $s_\theta \approx \nabla \log p_t$ 与学 vector field $v_\theta \approx u_t$ 是<strong>同一信息的两种参数化</strong></p></div>
<p>$$v_\theta(t, x) = f(x, t) - \tfrac{1}{2} g^2(t)\, s_\theta(t, x)$$ 所以在 VP/VE path 下,FM 可被视为 score matching 在 ODE 视角下的等价参数化。<strong>但对 Rectified Flow / OT-CFM 不成立</strong>(没有标准 SDE 对应),那里 FM 是更一般的 vector-field regression。</p>
<h3 id="61-velocity--score--noise-prediction-互换必考">6.1 Velocity ↔ Score ↔ Noise prediction 互换(必考)</h3>
<p>VP/VE path 下,假设 $x_t = \alpha(t) x_1 + \sigma(t) x_0$$x_0 \sim \mathcal{N}(0, I)$ 是噪声方向),三种主流 prediction target 之间是线性可逆转换:</p>
<p>$$ \begin{aligned} \epsilon\text{-prediction} &:\quad \epsilon_\theta(t, x_t) \approx x_0 \\ x_0\text{-prediction} &:\quad x^0_\theta(t, x_t) \approx x_1 \\ v\text{-prediction (Salimans-Ho)} &:\quad v_\theta(t, x_t) \approx \alpha'(t) x_1 + \sigma'(t) x_0 \\ \text{score} &:\quad s_\theta(t, x_t) \approx -x_0 / \sigma(t) \end{aligned} $$</p>
<p>已知 $x_t$ 和任一 prediction,可代数恢复其他三种。例如 VP 下 $\epsilon$ 与 score 关系:</p>
<p>$$s_\theta(t, x_t) = -\epsilon_\theta(t, x_t) / \sigma(t)$$</p>
<p>这就是为什么 DDPM (学 $\epsilon$) 与 score-based (学 $\nabla \log p_t$) 是<strong>等价参数化</strong>。Flow matching 学 $v = \alpha' x_1 + \sigma' x_0$ 也是其中一种,且在 RF (linear) 下退化为 $v = x_1 - x_0$。</p>
<h3 id="62-几种-path-与-diffusion-的对应表">6.2 几种 path 与 diffusion 的对应表</h3>
<table><thead><tr><th>FM Path</th><th>等价 Diffusion / SDE</th><th>典型 noise schedule</th></tr></thead><tbody><tr><td>VP cosine</td><td>DDPM (cosine)</td><td>$\bar\alpha_t = \cos^2(\pi t/2)$</td></tr><tr><td>VP linear</td><td>DDPM (linear β)</td><td>$\beta_t = \beta_0 + t(\beta_1 - \beta_0)$</td></tr><tr><td>VE</td><td>SMLD / EDM</td><td>$\sigma_t \in [\sigma_\min, \sigma_\max]$ 对数线性</td></tr><tr><td>Rectified Flow</td><td>没有标准非零扩散 noising SDE 对应(退化情形除外)</td><td>路径本身是直线,"最短" path</td></tr></tbody></table>
<h3 id="63-为什么-rectified-flow-训练--采样相对稳">6.3 为什么 Rectified Flow 训练 / 采样相对"稳"</h3>
<ul><li><strong>常数 target</strong>$u_t = x_1 - x_0$ 不显式依赖 $t$(在给定 $x_0, x_1$ 后),数值上易拟合</li><li><strong>直线 path</strong>:少步数 ODE 积分误差小</li><li><strong>Loss conditioning</strong>RF 本身比 native DDPM 训练更平衡;但<strong>不是说不需要 reweighting</strong>——SD3 在 RF 之上仍做 logit-normal $t$ sampling 等 reweighting 并 ablate 出涨点</li><li><strong>Reflow 可压缩 NFE</strong>1-step 生成路线(InstaFlow / SD3-Turbo / Flux-Schnell</li></ul>
<h2 id="7-高级话题">§7 高级话题</h2>
<h3 id="71-reflowliu-et-al-2022-iclr">7.1 ReflowLiu et al. 2022, ICLR</h3>
<p>Rectified Flow 之所以能少步数生成,关键是 <strong>reflow 算法</strong></p>
<ol><li>第一次训练得到 $v_\theta^{(1)}$(用独立配对 $(x_0, x_1) \sim p_0 \otimes p_1$</li><li>用 $v_\theta^{(1)}$ 跑 ODE 生成<strong>coupled</strong> pair $(x_0, x_1^{(1)})$,即 $x_1^{(1)} = \text{ODE}(x_0; v_\theta^{(1)})$</li><li>用 coupled pair 重新训练得到 $v_\theta^{(2)}$,新的 trajectory 更<strong></strong></li><li>重复——Liu et al. 2022 证明在适当假设下,配对的 <strong>convex transport cost</strong> 非增(每次 reflow 不会让总传输成本变差)</li></ol>
<p>"trajectory 变直"是直觉与经验观察;具体严格定理是 transport cost 单调性。实际 1-2 次 reflow 就能做到 4-step 媲美 50-stepInstaFlow / SD3-Turbo / Flux-Schnell)。极限:完全直线 → 1-step 生成($x_1 = x_0 + v_\theta(0, x_0)$)。</p>
<h3 id="72-conditional-flow-matching-cfg">7.2 Conditional Flow Matching (CFG)</h3>
<p>对条件生成(如 text-to-image),模型接收额外条件 $c$:</p>
<p>$$v_\theta(t, x, c)$$</p>
<p>训练时以概率 $p_\text{drop}$(一般 0.1)把 $c$ 替换成空(如 null embedding),得到 <strong>无条件 head</strong></p>
<p>采样时用 <strong>Classifier-Free Guidance</strong></p>
<p>$$v_\text{CFG}(t, x, c) = v_\theta(t, x, \emptyset) + s \cdot \left[v_\theta(t, x, c) - v_\theta(t, x, \emptyset)\right]$$</p>
<p>$s$ 是 guidance scale(一般 1.5-7.5)。$s > 1$ 时放大 conditional 信号,提升文本对齐但损失多样性。</p>
<h3 id="73-logit-normal-tsd3-默认">7.3 Logit-normal $t$SD3 默认)</h3>
<p>SD3 (Esser et al. 2024) 发现,<strong>$t \sim \mathcal{U}[0, 1]$ 不是最优</strong>。中间区域($t \approx 0.5$)的 target 噪声-信号比最难学。改成:</p>
<p>$$t = \sigma(\tau), \quad \tau \sim \mathcal{N}(m, s^2)$$</p>
<p>即 $\tau$ 高斯采样后 sigmoid 映射回 $(0, 1)$,可调 $m, s$ 控制 $t$ 分布偏重哪段。默认 $m = 0, s = 1$ 时 $t$ 集中在 0.5 附近。这是 SD3 论文中 ablation 涨点的关键之一。</p>
<h2 id="8-完整可运行示例2d-toy">§8 完整可运行示例(2D toy</h2>
<p>下面是端到端最小可运行示例:训练 vector field 学习把 $\mathcal{N}(0, I)$ 映到一个 2D 月亮形分布。</p>
<pre><code class="language-python">if __name__ == &quot;__main__&quot;:
# 1) 数据 (target distribution p_1): 2D 月亮形
from sklearn.datasets import make_moons
def sample_moons(n: int) -&gt; torch.Tensor:
X, _ = make_moons(n_samples=n, noise=0.05)
return torch.tensor(X, dtype=torch.float32) * 2.0 # scale
# 2) 模型 + path
model = VectorFieldMLP(dim=2, hidden=128)
path = rectified_flow_path()
# 3) &quot;dataloader&quot; (随机生成)
class MoonDataset:
def __init__(self, batch=512, total=5000):
self.batch = batch; self.total = total
def __iter__(self):
for _ in range(self.total):
yield sample_moons(self.batch)
# 4) 训练
train_flow_matching(
model,
MoonDataset(batch=512, total=2000),
path=path,
total_steps=2000,
lr=3e-4,
device=&quot;cuda&quot; if torch.cuda.is_available() else &quot;cpu&quot;,
log_every=100,
)
# 5) 采样
model.eval()
device = next(model.parameters()).device
x0 = torch.randn(2000, 2, device=device)
x_samples = euler_sampler(model, x0, steps=50)
# 与真实 2D 月亮叠图,可视化检验
import matplotlib.pyplot as plt
real = sample_moons(2000).numpy()
fake = x_samples.cpu().numpy()
plt.scatter(real[:, 0], real[:, 1], alpha=0.3, label=&quot;real&quot;)
plt.scatter(fake[:, 0], fake[:, 1], alpha=0.3, label=&quot;generated&quot;)
plt.legend(); plt.savefig(&quot;flow_matching_moons.png&quot;, dpi=120)</code></pre>
<div class="callout callout-warn"><div class="callout-title">Production 化要补的(教学版未含)</div><p>上线前需补的工程项如下。</p></div>
<ul><li><strong>EMA scheduler</strong>:训练后期 decay 更接近 1(如 0.9999 → 0.99995</li><li><strong>Gradient checkpointing</strong>U-Net / DiT 显存优化</li><li><strong>Mixed precision</strong>fp16 / bf16 + GradScaler</li><li><strong>Latent space</strong>:高分辨率图像在 VAE latent 里跑 FMLDM / SD3 / FLUX</li><li><strong>Conditioning</strong>text encoder (T5 / CLIP) + cross-attention 或 token concat</li><li><strong>Distributed</strong>DDP / FSDP for multi-GPU</li><li><strong>Loss weighting</strong>SD3 用 logit-normal $t$ 已隐式 reweightingEDM 显式 SNR weighting</li></ul>
<p><strong>Flow Matching Quick Reference</strong> · 主要参考:Lipman et al. 2023 (Flow Matching), Liu et al. 2022 (Rectified Flow), Esser et al. 2024 (SD3 / MM-DiT)</p>
<footer class="aris-footer">
Generated by <a href="https://github.com/wanshuiyin/Auto-claude-code-research-in-sleep/blob/main/skills/render-html/SKILL.md">ARIS <code>/render-html</code></a> ·
source path <code>docs/tutorials/flow_matching_tutorial.md</code> ·
SHA256 <code>ab4358901d20</code> ·
generated at 2026-05-19 03:40 UTC.
This is a generated view — edit the source Markdown, then re-render.
</footer>
</main>
</div>
</body>
</html>