flow matching 伪代码-1
小记训练随机取一条从噪声x0x_0x0到数据x1x_1x1的路径学习这条路径的速度。推理从噪声出发用模型预测的速度场积分 ODE逐步走到数据分布。训练时模型学到的并不是某一个具体样本的“直线”而是在所有条件路径的监督下最终学习一个能够把噪声分布运输到数据分布的整体 velocity field。训练伪代码CFM / Rectified Flow线性路径# model(x_t, t):# 给定当前状态 x_t 和时间 t# 预测此时应该沿着哪个方向移动即速度场 v_theta(x_t, t)## x1:# 从真实数据分布中采样的数据例如真实图像# x1 ~ p_dataforx1indataloader:# ---------------------------------------------------------# 1. 采样起点噪声和时间# ---------------------------------------------------------# 从标准高斯分布采样初始状态# x0 ~ N(0, I)## x0 是生成过程的起点推理时也会从这样的随机噪声开始x0randn_like(x1)# 随机采样时间 t ∈ [0, 1]## t 0 : 完全处于噪声端 x0# t 1 : 完全到达真实数据端 x1## [B, 1, 1, 1] 表示对 batch 中每个样本独立采样一个时间# 后续通过广播机制与 [B, C, H, W] 的图像进行计算trand_uniform([B,1,1,1])# t ~ Uniform(0, 1)# ---------------------------------------------------------# 2. 构造训练时的中间状态 x_t# ---------------------------------------------------------# 在线性路径straight-line path下# 将 x0 和 x1 进行线性插值## x_t (1 - t) * x0 t * x1## 因此# t 0 - x_t x0# t 1 - x_t x1## 可以理解为# x0 -------- x1# x_t## 训练时并不是让模型自己走到 x_t# 而是直接根据已知的 x0、x1 构造任意时间点的 x_t。x_t(1-t)*x0t*x1# ---------------------------------------------------------# 3. 计算 Conditional Flow Matching 的目标速度# ---------------------------------------------------------# 对线性路径## x_t (1 - t)x0 tx1## 对 t 求导## dx_t / dt x1 - x0## 因此真实的目标速度是## target_v x1 - x0## 它表示# 从当前 x_t 出发如果沿着这条 x0 - x1 的直线路径走# 此时应该朝哪个方向移动以及移动多快。## 注意# target_v 与 t 无关但它对应的是当前这条# x0 - x1 条件路径的速度。target_vx1-x0# ---------------------------------------------------------# 4. 模型预测速度并进行优化# ---------------------------------------------------------# 模型只看到# 当前状态 x_t# 当前时间 t## 模型不知道训练时真正采样出来的 x0 和 x1# 它需要学习根据 x_t、t 推断正确的速度方向。pred_vmodel(x_t,t)# 使用 MSE 让模型预测的速度接近目标速度## L ||v_theta(x_t, t) - (x1 - x0)||²## 通过大量不同的 x0、x1、t 训练后# 模型学习整个数据分布上的 velocity field。lossmse_loss(pred_v,target_v)# 反向传播并更新模型参数loss.backward()optimizer.step()optimizer.zero_grad()推理伪代码欧拉法解 ODE# model(x_t, t):# 训练好的速度场模型## 输入# 当前状态 x_t# 当前时间 t## 输出# 当前时间点的速度 v_theta(x_t, t)## N:# ODE 数值积分的步数# N 越大通常积分越精确但推理计算量也越大# ---------------------------------------------------------# 1. 从噪声分布采样初始状态# ---------------------------------------------------------# 生成过程从标准高斯噪声开始## x(0) ~ N(0, I)## 这对应训练时的 x0。## 后续通过学习到的 velocity field# 将这个噪声逐渐运输到数据分布。xrandn([B,C,H,W])# ---------------------------------------------------------# 2. 将连续时间 [0, 1] 离散化# ---------------------------------------------------------# 将 [0, 1] 划分成 N 个小区间## dt 1 / N## 因此## t_0 0# t_1 dt# ...# t_N 1## ODE## dx/dt v_theta(x, t)dt1.0/N# ---------------------------------------------------------# 3. 使用 Euler 方法求解 ODE# ---------------------------------------------------------foriinrange(N):# 当前时间点## t 从 0 逐渐增加到接近 1ti/N# -----------------------------------------------------# 根据当前状态和时间预测 velocity# -----------------------------------------------------## v v_theta(x_t, t)## velocity 可以理解为# “在当前 x_t 和当前时间 t 下# 下一瞬间应该往哪个方向移动”vmodel(x,t)# -----------------------------------------------------# Euler 数值积分# -----------------------------------------------------## ODE## dx/dt v_theta(x, t)## Euler 离散化## x_{tdt}# x_t dt * v_theta(x_t, t)## 每一步都根据当前 velocity 向前走一小步。xxdt*v# ---------------------------------------------------------# 4. 得到最终生成结果# ---------------------------------------------------------# 经过 N 次 ODE 更新后## x ≈ x(1)## 此时状态已经从初始噪声分布逐渐运输到了# 模型学习到的数据分布。x_genx最核心的逻辑其实可以浓缩成一句话训练随机取x0x_0x0、x1x_1x1、ttt↓构造xt(1−t)x0tx1x_t (1-t)x_0 tx_1xt(1−t)x0tx1↓计算真实速度vx1−x0v x_1-x_0vx1−x0↓训练模型学习vθ(xt,t)v_θ(x_t, t)vθ(xt,t)推理x0 N(0,I)x_0 ~ N(0,I)x0N(0,I)↓不断计算vθ(xt,t)v_θ(x_t,t)vθ(xt,t)↓x←xdt⋅vθ(x,t)x ← x dt·v_θ(x,t)x←xdt⋅vθ(x,t)↓最终得到x1x_1x1≈ 数据样本所以 Flow Matching 本质上是在训练一个“速度场”训练阶段告诉模型“在各种xt,tx_t,txt,t下应该往哪里走”推理阶段再把这个速度场当作 ODE 的右端项从噪声一路积分到数据。
上一篇/下一篇内容由系统自动关联
返回资讯列表 →