mirror of
https://github.com/rohitg00/ai-engineering-from-scratch.git
synced 2026-10-02 01:54:39 +08:00
feat(phase-08/13): flow matching and rectified flows
1-D flow matching with straight-line interpolant and Euler inference at 1/2/4/8/20 steps. Shows 4-step matches 20-step quality on a toy mixture. Covers rectified flow reflow and the SD3 / Flux.1 / AudioCraft 2 switchover.
This commit is contained in:
@@ -0,0 +1,59 @@
|
||||
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 900 500" font-family="Georgia, 'Times New Roman', serif">
|
||||
<defs>
|
||||
<marker id="arrow" viewBox="0 0 10 10" refX="9" refY="5" markerWidth="6" markerHeight="6" orient="auto">
|
||||
<path d="M0,0 L10,5 L0,10 z" fill="#1a1a1a"/>
|
||||
</marker>
|
||||
<style>
|
||||
.box { fill: #faf6ef; stroke: #1a1a1a; stroke-width: 1.5; }
|
||||
.hot { fill: #fff1d6; stroke: #c0392b; stroke-width: 1.5; }
|
||||
.cold { fill: #eaf4ff; stroke: #2c5f8c; stroke-width: 1.5; }
|
||||
.label { font-size: 14px; font-weight: 600; fill: #1a1a1a; }
|
||||
.content { font-size: 12px; fill: #333; }
|
||||
.mono { font-size: 12px; fill: #333; font-family: 'Menlo', monospace; }
|
||||
.caption { font-size: 11px; fill: #555; font-style: italic; }
|
||||
.title { font-size: 16px; font-weight: 700; fill: #1a1a1a; }
|
||||
</style>
|
||||
</defs>
|
||||
|
||||
<text x="450" y="28" text-anchor="middle" class="title">flow matching: train on a straight line</text>
|
||||
|
||||
<!-- DDPM vs FM comparison -->
|
||||
<text x="225" y="70" text-anchor="middle" class="label">DDPM: curved path</text>
|
||||
<rect x="50" y="90" width="350" height="170" class="box"/>
|
||||
|
||||
<!-- data cluster -->
|
||||
<circle cx="95" cy="210" r="18" fill="#c0392b" opacity="0.4"/>
|
||||
<text x="95" y="240" text-anchor="middle" class="caption">x_0 (data)</text>
|
||||
<!-- noise -->
|
||||
<circle cx="365" cy="130" r="18" fill="#2c5f8c" opacity="0.4"/>
|
||||
<text x="365" y="110" text-anchor="middle" class="caption">x_T ~ N(0, I)</text>
|
||||
|
||||
<path d="M 95,210 Q 160,160 200,200 T 280,150 T 365,130" fill="none" stroke="#1a1a1a" stroke-width="1.5"/>
|
||||
<text x="225" y="190" text-anchor="middle" class="caption">1000-step SDE</text>
|
||||
<text x="225" y="253" text-anchor="middle" class="caption">DDIM collapses to ~50 steps</text>
|
||||
|
||||
<text x="665" y="70" text-anchor="middle" class="label">flow matching: straight line</text>
|
||||
<rect x="490" y="90" width="380" height="170" class="box"/>
|
||||
|
||||
<circle cx="535" cy="210" r="18" fill="#c0392b" opacity="0.4"/>
|
||||
<text x="535" y="240" text-anchor="middle" class="caption">x_0</text>
|
||||
<circle cx="835" cy="130" r="18" fill="#2c5f8c" opacity="0.4"/>
|
||||
<text x="835" y="110" text-anchor="middle" class="caption">x_1 ~ N(0, I)</text>
|
||||
|
||||
<line x1="535" y1="210" x2="835" y2="130" stroke="#1a1a1a" stroke-width="2"/>
|
||||
<text x="685" y="190" text-anchor="middle" class="mono">x_t = t · x_1 + (1-t) · x_0</text>
|
||||
<text x="685" y="253" text-anchor="middle" class="caption">2-8 Euler steps at inference</text>
|
||||
|
||||
<!-- loss -->
|
||||
<rect x="50" y="290" width="820" height="70" class="hot"/>
|
||||
<text x="460" y="315" text-anchor="middle" class="label">loss = || v_θ(x_t, t) - (x_1 - x_0) ||²</text>
|
||||
<text x="460" y="340" text-anchor="middle" class="caption">simulation-free training; target is a constant vector along the straight line</text>
|
||||
|
||||
<!-- rectified flow -->
|
||||
<rect x="50" y="380" width="820" height="110" class="cold"/>
|
||||
<text x="460" y="402" text-anchor="middle" class="label">rectified flow: iteratively straighten</text>
|
||||
<text x="460" y="425" text-anchor="middle" class="caption">1. train v_1 with random (x_0, x_1) pairs</text>
|
||||
<text x="460" y="445" text-anchor="middle" class="caption">2. sample paired (x_0, x_1) by integrating v_1 -> pairs are ODE-matched</text>
|
||||
<text x="460" y="465" text-anchor="middle" class="caption">3. retrain v_2 on those pairs -> genuinely straighter flow</text>
|
||||
<text x="460" y="482" text-anchor="middle" class="caption">2 reflow iterations enable 1-step inference (SDXL-Turbo, SD3-Turbo, Flux-schnell)</text>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 3.5 KiB |
@@ -0,0 +1,148 @@
|
||||
import math
|
||||
import random
|
||||
|
||||
|
||||
def tanh(v):
|
||||
return [math.tanh(x) for x in v]
|
||||
|
||||
|
||||
def tanh_grad(h):
|
||||
return [1 - x * x for x in h]
|
||||
|
||||
|
||||
def matmul(W, x):
|
||||
return [sum(w * xi for w, xi in zip(row, x)) for row in W]
|
||||
|
||||
|
||||
def add(a, b):
|
||||
return [x + y for x, y in zip(a, b)]
|
||||
|
||||
|
||||
def randn_matrix(rows, cols, rng, scale=0.3):
|
||||
return [[rng.gauss(0, scale) for _ in range(cols)] for _ in range(rows)]
|
||||
|
||||
|
||||
def init_net(in_dim, hidden, out_dim, rng):
|
||||
return {
|
||||
"W1": randn_matrix(hidden, in_dim, rng),
|
||||
"b1": [0.0] * hidden,
|
||||
"W2": randn_matrix(hidden, hidden, rng),
|
||||
"b2": [0.0] * hidden,
|
||||
"W3": randn_matrix(out_dim, hidden, rng),
|
||||
"b3": [0.0] * out_dim,
|
||||
}
|
||||
|
||||
|
||||
def forward(x, t, net):
|
||||
inp = [x, t, t * t, math.sin(2 * math.pi * t), math.cos(2 * math.pi * t)]
|
||||
pre1 = add(matmul(net["W1"], inp), net["b1"])
|
||||
h1 = tanh(pre1)
|
||||
pre2 = add(matmul(net["W2"], h1), net["b2"])
|
||||
h2 = tanh(pre2)
|
||||
out = add(matmul(net["W3"], h2), net["b3"])
|
||||
return out[0], {"inp": inp, "h1": h1, "h2": h2}
|
||||
|
||||
|
||||
def backward(target, out, cache, net):
|
||||
grads = {k: None for k in net}
|
||||
for p in net:
|
||||
if isinstance(net[p][0], list):
|
||||
grads[p] = [[0.0] * len(net[p][0]) for _ in net[p]]
|
||||
else:
|
||||
grads[p] = [0.0] * len(net[p])
|
||||
d_out = 2 * (out - target)
|
||||
grads["b3"][0] += d_out
|
||||
for j in range(len(cache["h2"])):
|
||||
grads["W3"][0][j] += d_out * cache["h2"][j]
|
||||
d_h2 = [net["W3"][0][j] * d_out for j in range(len(cache["h2"]))]
|
||||
d_pre2 = [d_h2[j] * tanh_grad(cache["h2"])[j] for j in range(len(cache["h2"]))]
|
||||
for j in range(len(cache["h2"])):
|
||||
grads["b2"][j] += d_pre2[j]
|
||||
for k in range(len(cache["h1"])):
|
||||
grads["W2"][j][k] += d_pre2[j] * cache["h1"][k]
|
||||
d_h1 = [sum(net["W2"][j][k] * d_pre2[j] for j in range(len(cache["h2"])))
|
||||
for k in range(len(cache["h1"]))]
|
||||
d_pre1 = [d_h1[j] * tanh_grad(cache["h1"])[j] for j in range(len(cache["h1"]))]
|
||||
for j in range(len(cache["h1"])):
|
||||
grads["b1"][j] += d_pre1[j]
|
||||
for k in range(len(cache["inp"])):
|
||||
grads["W1"][j][k] += d_pre1[j] * cache["inp"][k]
|
||||
return grads
|
||||
|
||||
|
||||
def apply(net, grads, lr):
|
||||
for k, v in net.items():
|
||||
if isinstance(v[0], list):
|
||||
for i in range(len(v)):
|
||||
for j in range(len(v[i])):
|
||||
v[i][j] -= lr * grads[k][i][j]
|
||||
else:
|
||||
for i in range(len(v)):
|
||||
v[i] -= lr * grads[k][i]
|
||||
|
||||
|
||||
def sample_data(rng):
|
||||
return rng.gauss(-2.0, 0.3) if rng.random() < 0.5 else rng.gauss(2.0, 0.3)
|
||||
|
||||
|
||||
def train(net, steps, lr, rng):
|
||||
for _ in range(steps):
|
||||
x0 = sample_data(rng)
|
||||
x1 = rng.gauss(0, 1)
|
||||
t = rng.random()
|
||||
x_t = t * x1 + (1 - t) * x0
|
||||
target = x1 - x0
|
||||
pred, cache = forward(x_t, t, net)
|
||||
grads = backward(target, pred, cache, net)
|
||||
apply(net, grads, lr)
|
||||
|
||||
|
||||
def sample(net, num_steps, rng):
|
||||
x = rng.gauss(0, 1)
|
||||
dt = 1.0 / num_steps
|
||||
for i in range(num_steps):
|
||||
t = 1.0 - i * dt
|
||||
v, _ = forward(x, t, net)
|
||||
x -= dt * v
|
||||
return x
|
||||
|
||||
|
||||
def histogram(samples, lo=-5.0, hi=5.0, bins=30):
|
||||
width = (hi - lo) / bins
|
||||
counts = [0] * bins
|
||||
for s in samples:
|
||||
if lo <= s < hi:
|
||||
counts[int((s - lo) / width)] += 1
|
||||
peak = max(counts) or 1
|
||||
height = 6
|
||||
rows = []
|
||||
for r in range(height, 0, -1):
|
||||
thr = peak * r / height
|
||||
rows.append("".join("#" if c >= thr else " " for c in counts))
|
||||
rows.append("-" * bins)
|
||||
return "\n".join(rows)
|
||||
|
||||
|
||||
def main():
|
||||
rng = random.Random(31)
|
||||
net = init_net(in_dim=5, hidden=24, out_dim=1, rng=rng)
|
||||
|
||||
print("=== training flow matching on two-mode mixture ===")
|
||||
train(net, steps=6000, lr=0.01, rng=rng)
|
||||
|
||||
print()
|
||||
for num_steps in [1, 2, 4, 8, 20]:
|
||||
samples = [sample(net, num_steps, rng) for _ in range(500)]
|
||||
m = sum(samples) / len(samples)
|
||||
pos = sum(1 for s in samples if s > 0)
|
||||
print(f"=== {num_steps}-step Euler integration ===")
|
||||
print(histogram(samples))
|
||||
print(f"mean {m:+.2f}, left mode = {500 - pos}, right mode = {pos}")
|
||||
print()
|
||||
|
||||
print("takeaway: straight-line flow matching lets Euler work at 4-8 steps.")
|
||||
print(" DDPM needs 20+ for similar quality in this toy.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,163 @@
|
||||
# Flow Matching & Rectified Flows
|
||||
|
||||
> Diffusion models take 20-50 sampling steps because they walk a curved path from noise to data. Flow matching (Lipman et al., 2023) and rectified flow (Liu et al., 2022) trained straight paths. Straighter paths mean fewer steps mean faster inference. Stable Diffusion 3, Flux.1, and AudioCraft 2 all switched to flow matching in 2024.
|
||||
|
||||
**Type:** Build
|
||||
**Languages:** Python
|
||||
**Prerequisites:** Phase 8 · 06 (DDPM), Phase 1 · Calculus
|
||||
**Time:** ~45 minutes
|
||||
|
||||
## The Problem
|
||||
|
||||
DDPM's reverse process is a 1000-step stochastic walk from `N(0, I)` back to the data distribution. DDIM collapsed it to 20-50 deterministic steps. You want fewer steps — ideally one. The blocker is that the ODE solving the reverse process is stiff; the path is curved.
|
||||
|
||||
If you could train the model such that the path from noise to data was a *straight line*, a single Euler step from `t=1` to `t=0` would work. Flow matching builds this directly: define a straight-line interpolation from `x_1 ∼ N(0, I)` to `x_0 ∼ data`, train a vector field `v_θ(x, t)` to match its time derivative, integrate at inference.
|
||||
|
||||
Rectified flow (Liu 2022) goes further: iteratively straighten the paths with a reflow procedure that produces a progressively closer-to-linear ODE. After two reflow iterations, a 2-step sampler matches 50-step DDPM quality.
|
||||
|
||||
## The Concept
|
||||
|
||||

|
||||
|
||||
### Straight-line flow
|
||||
|
||||
Define:
|
||||
|
||||
```
|
||||
x_t = t · x_1 + (1 - t) · x_0, t ∈ [0, 1]
|
||||
```
|
||||
|
||||
where `x_0 ~ data` and `x_1 ~ N(0, I)`. The time derivative along this straight line is constant:
|
||||
|
||||
```
|
||||
dx_t / dt = x_1 - x_0
|
||||
```
|
||||
|
||||
Define a neural vector field `v_θ(x_t, t)` and train it to match this derivative:
|
||||
|
||||
```
|
||||
L = E_{x_0, x_1, t} || v_θ(x_t, t) - (x_1 - x_0) ||²
|
||||
```
|
||||
|
||||
This is the **conditional flow matching** loss (Lipman 2023). Training is simulation-free: you never unroll the ODE. Just sample `(x_0, x_1, t)` and regress.
|
||||
|
||||
### Sampling
|
||||
|
||||
At inference, integrate the learned vector field *backwards* in time:
|
||||
|
||||
```
|
||||
x_{t-Δt} = x_t - Δt · v_θ(x_t, t)
|
||||
```
|
||||
|
||||
Start at `x_1 ~ N(0, I)`, Euler-step down to `t=0`.
|
||||
|
||||
### Rectified flow (Liu 2022)
|
||||
|
||||
Straight-line flow works but the learned paths are *not actually straight* — they curve because many `x_0`s can map to the same `x_1`. Rectified flow's reflow step:
|
||||
|
||||
1. Train flow model v_1 with random pairings.
|
||||
2. Sample N pairs `(x_1, x_0)` by integrating v_1 from `x_1` to its landing `x_0`.
|
||||
3. Train v_2 on those paired examples. Because the pairs are now "ODE-matched", the straight-line interpolant between them is genuinely flatter.
|
||||
4. Repeat.
|
||||
|
||||
In practice 2 reflow iterations get you to near-linear, enabling 2-4 step inference. SDXL-Turbo, SD3-Turbo, LCM are all distilled-from-flow-matching models.
|
||||
|
||||
### Why this won for images in 2024
|
||||
|
||||
Three reasons:
|
||||
|
||||
1. **Simulation-free training** — no ODE unrolling during training, trivial to implement.
|
||||
2. **Better loss geometry** — straight paths have consistent signal-to-noise, whereas DDPM ε-loss has bad SNR at edges of the schedule.
|
||||
3. **Faster inference** — 4-8 steps at SDXL-Turbo quality; 1 step with consistency distillation.
|
||||
|
||||
## Flow matching vs DDPM — the exact connection
|
||||
|
||||
Flow matching with a Gaussian-conditional path is diffusion *with a specific noise schedule*. Pick the `x_t = α(t) x_0 + σ(t) x_1` schedule and flow matching recovers Stratonovich-reformulated diffusion with `v = α'·x_0 - σ'·x_1`. The two are algebraically equivalent for Gaussian paths.
|
||||
|
||||
What flow matching added: the *clarity* of the target (a plain velocity), a cleaner loss, and the license to experiment with non-Gaussian interpolants.
|
||||
|
||||
## Build It
|
||||
|
||||
`code/main.py` implements 1-D flow matching on a two-mode Gaussian mixture. The vector field `v_θ(x, t)` is a tiny MLP trained with the straight-line target. At inference, integrate 1, 2, 4, and 20 Euler steps and compare sample quality.
|
||||
|
||||
### Step 1: training loss
|
||||
|
||||
```python
|
||||
def train_step(x0, net, rng, lr):
|
||||
x1 = rng.gauss(0, 1)
|
||||
t = rng.random()
|
||||
x_t = t * x1 + (1 - t) * x0
|
||||
target = x1 - x0
|
||||
pred = net_forward(x_t, t)
|
||||
loss = (pred - target) ** 2
|
||||
# backprop + update
|
||||
```
|
||||
|
||||
### Step 2: multi-step inference
|
||||
|
||||
```python
|
||||
def sample(net, num_steps):
|
||||
x = rng.gauss(0, 1)
|
||||
for i in range(num_steps):
|
||||
t = 1.0 - i / num_steps
|
||||
dt = 1.0 / num_steps
|
||||
x -= dt * net_forward(x, t)
|
||||
return x
|
||||
```
|
||||
|
||||
### Step 3: compare step counts
|
||||
|
||||
Expect the 4-step sampler to already match the 20-step quality — a big deal for latency.
|
||||
|
||||
## Pitfalls
|
||||
|
||||
- **Time parameterization.** Flow matching uses `t ∈ [0, 1]` with `t=0` at data, `t=1` at noise. DDPM uses `t ∈ [0, T]` with `t=0` at data, `t=T` at noise. Same direction, different scale. Papers get this wrong constantly.
|
||||
- **Schedule choice.** Rectified flow's straight line is "the" flow-matching schedule, but you can use cosine or logit-normal t-sampling (SD3 does this) for better scale coverage.
|
||||
- **Reflow cost.** Generating the paired dataset for reflow is a full inference pass per sample. Only do reflow when you really need 1-2 step inference.
|
||||
- **Classifier-free guidance still applies.** Just swap ε for v in the linear combination: `v_cfg = (1+w) v_cond - w v_uncond`.
|
||||
|
||||
## Use It
|
||||
|
||||
| Use case | 2026 stack |
|
||||
|----------|-----------|
|
||||
| Text-to-image, best quality | Flow matching: SD3, Flux.1-dev |
|
||||
| Text-to-image, 1-4 steps | Distilled flow matching: Flux.1-schnell, SD3-Turbo, SDXL-Turbo |
|
||||
| Real-time inference | Consistency distillation from a flow-matched base (LCM, PCM) |
|
||||
| Audio generation | Flow matching: Stable Audio 2.5, AudioCraft 2 |
|
||||
| Video generation | Flow matching mixed with diffusion (Sora, Veo, Stable Video) |
|
||||
| Science / physics (particle trajectories, molecules) | Flow matching + equivariant vector field |
|
||||
|
||||
Whenever a paper says "faster than diffusion" in 2025-2026, it is almost always flow matching + distillation.
|
||||
|
||||
## Ship It
|
||||
|
||||
Save `outputs/skill-fm-tuner.md`. Skill takes a diffusion-style model spec and converts it to a flow-matching training config: schedule choice, time sampling distribution (uniform / logit-normal), optimizer, reflow plan, target step count, eval protocol.
|
||||
|
||||
## Exercises
|
||||
|
||||
1. **Easy.** Run `code/main.py` and compare 1-step vs 20-step MSE vs the true data distribution.
|
||||
2. **Medium.** Switch from uniform `t` sampling to logit-normal (concentrates sampling at mid-t). Does the model quality improve?
|
||||
3. **Hard.** Implement one reflow iteration: generate paired (x_0, x_1) by integrating the first model, train a second model on the pairs, and compare 1-step sample quality.
|
||||
|
||||
## Key Terms
|
||||
|
||||
| Term | What people say | What it actually means |
|
||||
|------|-----------------|-----------------------|
|
||||
| Flow matching | "Straight-line diffusion" | Train `v_θ(x, t)` to match `x_1 - x_0` along an interpolant. |
|
||||
| Rectified flow | "Reflow" | Iterative procedure that straightens learned flows. |
|
||||
| Velocity field | "v_θ" | Output of the model — the direction to move `x_t`. |
|
||||
| Straight-line interpolant | "The path" | `x_t = (1-t)·x_0 + t·x_1`; trivial target derivative. |
|
||||
| Euler sampler | "1st order ODE solver" | Simplest integrator; works well when paths are straight. |
|
||||
| Logit-normal t | "SD3 sampling" | Concentrate `t` sampling toward mid-values where gradients are strongest. |
|
||||
| Consistency distillation | "1-step sampler" | Train a student to map any `x_t` directly to `x_0`. |
|
||||
| CFG with velocity | "v-CFG" | `v_cfg = (1+w) v_cond - w v_uncond`; same trick, new variable. |
|
||||
|
||||
## Further Reading
|
||||
|
||||
- [Liu, Gong, Liu (2022). Flow Straight and Fast: Learning to Generate and Transfer Data with Rectified Flow](https://arxiv.org/abs/2209.03003) — rectified flow.
|
||||
- [Lipman et al. (2023). Flow Matching for Generative Modeling](https://arxiv.org/abs/2210.02747) — flow matching.
|
||||
- [Esser et al. (2024). Scaling Rectified Flow Transformers for High-Resolution Image Synthesis](https://arxiv.org/abs/2403.03206) — SD3, rectified flow at scale.
|
||||
- [Albergo, Vanden-Eijnden (2023). Stochastic Interpolants](https://arxiv.org/abs/2303.08797) — general framework that covers FM + diffusion.
|
||||
- [Song et al. (2023). Consistency Models](https://arxiv.org/abs/2303.01469) — 1-step distillation of diffusion / flow.
|
||||
- [Sauer et al. (2023). Adversarial Diffusion Distillation (SDXL-Turbo)](https://arxiv.org/abs/2311.17042) — turbo variant.
|
||||
- [Black Forest Labs (2024). Flux.1 models](https://blackforestlabs.ai/announcing-black-forest-labs/) — flow matching in production.
|
||||
@@ -0,0 +1,20 @@
|
||||
---
|
||||
name: fm-tuner
|
||||
description: Convert a diffusion training plan into a flow-matching / rectified-flow config.
|
||||
version: 1.0.0
|
||||
phase: 8
|
||||
lesson: 13
|
||||
tags: [flow-matching, rectified-flow, diffusion]
|
||||
---
|
||||
|
||||
Given a diffusion-style training plan (data, compute, schedule, target step count, quality bar), output a flow-matching equivalent:
|
||||
|
||||
1. Schedule + interpolant. Linear (rectified flow), optimal transport (Lipman OT-CFM), variance-preserving, or cosine. One-sentence reason.
|
||||
2. Time sampling. Uniform, logit-normal (SD3), or mode-weighted. Warn when uniform sampling at 1000 Hz wastes capacity at endpoints.
|
||||
3. Target. Velocity v = x_1 - x_0 (rectified flow) or alpha'(t)x_1 + sigma'(t)x_0 (CFM). State which.
|
||||
4. Optimizer + lr warmup. Include AdamW with beta2 = 0.95 for stability at transformer scale.
|
||||
5. Reflow plan. Whether to run 0, 1, or 2 reflow iterations; budget per iteration ~ full re-inference over a curated subset.
|
||||
6. Step counts. Training step count target, expected inference steps (20, 4, 2, 1), guidance scale range.
|
||||
7. Eval. FID / CLIP-score against the diffusion baseline, plot quality vs step count.
|
||||
|
||||
Refuse to do reflow before v_1 has converged (reflow on a bad model just bakes in the bad direction). Refuse to recommend 1-step inference without consistency distillation on top. Flag any flow-matching model that targets > 20 step inference - if you need that many steps, you wasted the reformulation.
|
||||
Reference in New Issue
Block a user