所属方向:生成范式 | 难度:进阶 | 前置知识:DDPM 的训练目标与采样流程(
ddpm)、扩散过程的前向与反向推导(diffusion_math)
关键词:DDIM、DPM-Solver、确定性采样、二阶采样、步数、NFE
先给两个数字,都来自文末附录里能直接跑的脚本。
数字一:同样的精度,评估次数差 12.8 倍。 在一个二维八分量高斯混合上跑(加噪之后仍是高斯混合,所以 score 有闭式解,见 04 节),把生成终点到高精度参考解的 RMS 距离压到 1e-2:
| 采样器 | 压到 1e-2 需要的模型评估次数 | 压到 1e-3 |
|---|---|---|
| DDIM(时间步等间隔) | 256 | > 384 |
| DPM-Solver-2(单步二阶) | 63 | 191 |
| DPM-Solver++ 2M(多步二阶) | 40 | 128 |
| DPM-Solver++ 3M(多步三阶) | 20 | 64 |
DDIM 要 256 次评估,3M 只要 20 次。这不是"实现优化"级别的差距,是换了一套数学带来的差距——同一个模型、同一份起点、同一份随机数,只是步法不同。生产环境里 20 步出图还是 250 步出图,是"能上"和"不能上"的区别。
数字二:换个采样器,缓存策略就得重调。 这一条更隐蔽,也更容易踩。所有"复用上一步特征"的加速(DeepCache、各种 step-cache、特征缓存)都在做一个判断:这一步的模型输入和上一步差不多,就直接搬上一步的中间结果。那"差不多"到底是多少?在同一个模型、同一条 50 步轨迹上,量最后十步"模型输入相对变化":
| 采样器 | 最后 10 步输入相对变化 |
|---|---|
| DDIM(η=0) | 0.0197 |
| η=0.5 | 0.0667 |
| DDPM 祖采样(η=1) | 0.1294 |
η=1 的抖动是 η=0 的 6.6 倍。也就是说,你在 DDIM 上把阈值调成"输入变化小于 2% 就复用",切到 DDPM 祖采样后这条曲线整段都在 13% 附近,缓存一次都不会命中,加速方案静默失效;反过来,为了吃掉 η=1 的抖动把阈值放宽到 15%,切回 DDIM 后就会在真正需要重算的步上复用,画面细节直接掉。这不是调参没调好,这是两套采样器的轨迹性质不同。
所以这篇文章要回答的是一个问题:采样器到底是什么? 答案不是"一种跳步策略",而是一条常微分方程的离散格式。一旦这么看,两件事立刻变得可算:步法是多少阶、达到给定精度要花多少次模型评估(NFE,number of function evaluations)。而 DDIM——这个被无数人当成"DDPM 的加速版"的东西——恰好是这条 ODE 上最低的一阶格式。下面会把它验到 1e-14。
三句话讲完:
DDIM 不是"步幅更大的 DDPM"。 DDPM 的采样每步都要掷一份新噪声,DDIM 的贡献是证明:同一族前向过程可以用非马尔可夫的方式重新构造,只要边缘分布 $q(x_t|x_0)$ 不变,反向过程就有自由度——其中一个极端是完全不掷噪声。所以差别不在步幅,在随机项。
不掷噪声之后,采样变成解 ODE。 这个 ODE 叫概率流 ODE,它的解是确定性的:同一个起点永远给同一个终点。DDIM 的一步迭代,就是这条 ODE 的一个显式一阶格式(用当前点的信息走一整步)。用二维真轨迹看最直观:

这张图要看的是:蓝色的 η=0 轨迹是一条平滑曲线,橙色 η=1 轨迹在同一份随机数下走成了折线——每一步都被注入的方向改变。"同一个采样器家族"的两种极端,轨迹的几何性质完全不同,缓存类加速盯的正是这个几何性质。
DDPM 的采样式(祖先采样)长这样:先估干净图,再缩放,再加一份后验噪声。
$$x_{t-1} = \frac{1}{\sqrt{\alpha_t}}\Big(x_t - \frac{\beta_t}{\sqrt{1-\bar\alpha_t}}\epsilon_\theta(x_t,t)\Big) + \sqrt{\tilde\beta_t}z$$
DDIM 论文式 (12) 换了一个组织方式。令 $\hat x_0$ 是这一时刻对干净图的估计、$\epsilon_\theta$ 是模型输出的噪声预测,两者由 $\hat x_0 = (x_t-\sqrt{1-\bar\alpha_t}\epsilon_\theta)/\sqrt{\bar\alpha_t}$ 相互确定,则一步写成三项:
$$x_{t-1} = \underbrace{\sqrt{\bar\alpha_{t-1}}\hat x_0}_{\text{回数据的方向}} + \underbrace{\sqrt{1-\bar\alpha_{t-1}-\tilde\sigma_t^2}\epsilon_\theta}_{\text{沿当前噪声方向}} + \underbrace{\tilde\sigma_t z}_{\text{新掷的噪声}}$$
这里第二项必须使用目标时刻的总噪声方差 $1-\bar\alpha_{t-1}$,新注入噪声为
$$\tilde\sigma_t = \eta\sqrt{\frac{1-\bar\alpha_{t-1}}{1-\bar\alpha_t}}\sqrt{1-\frac{\bar\alpha_t}{\bar\alpha_{t-1}}}$$
每一项在干什么。 第一项是"这一时刻认为的干净图"乘上新时刻的 $\sqrt{\bar\alpha}$,也就是把估计值放到目标噪声水平上。第二项补偿两者的差额:前向过程是 $x_t=\sqrt{\bar\alpha_t}x_0+\sqrt{1-\bar\alpha_t}\epsilon$,要从 $x_t$ 走到 $x_{t-1}$,得把噪声幅度从 $\sqrt{1-\bar\alpha_t}$ 缩到 $\sqrt{1-\bar\alpha_{t-1}}$,这个缩放只能沿着"当前噪声方向"做——$\sqrt{1-\bar\alpha_{t-1}-\tilde\sigma_t^2}$ 正好是这两者的勾股差。第三项才是唯一带随机性的部分。
$\eta$ 是那个旋钮。 $\eta=0$ 时第三项消失,$\tilde\sigma_t=0$,迭代完全确定;$\eta=1$ 时 $\tilde\sigma_t$ 恰好等于 DDPM 的后验标准差,第二、三项合起来正好还原 DDPM 的加噪——在相同的完整时间表、fixed-small 方差、预测器及随机数下,$\eta=1$ 与该 DDPM 更新代数等价;跳步时应比较对应重排时间表的祖先采样,而不是原始 1000 步链。04 节会把这两条都量到浮点误差量级。
为什么可以有这个旋钮。 因为 DDIM 用的前向过程不是马尔可夫链:用 $q_\eta(x_{t-1}\mid x_t,x_0)$ 与终点分布重新构造非马尔可夫联合过程,但所有边缘 $q(x_t|x_0)$ 一个都没变。训练目标只依赖边缘(噪声预测的回归目标就是 $\epsilon$),所以模型不用重训;而反向的可选空间变大了,$\eta$ 就是在里面挑一条路。
$\eta=0$ 时迭代没有随机项,它可以被看成一条 ODE 的离散格式。连续时间下,前向过程是 $dx = -\frac{1}{2}\beta(t)x\mathrm{d}t + \sqrt{\beta(t)}\mathrm{d}w$(第一项把 $x$ 往 0 拉,第二项持续注入白噪声),对应的概率流 ODE(与 SDE 共享所有边缘分布的那个确定性方程)是
$$\frac{dx}{dt} = -\frac{1}{2}\beta(t)\Big(x + \nabla_x\log p_t(x)\Big)$$
两式一比就能看出这个方程的来历:把 SDE 的噪声项换成"噪声的均值流",也就是 score 项,随机性就被抽掉了,但每一时刻的边缘分布 $p_t(x)$ 一模一样。这就是为什么用同一条训练好的网络既能跑随机采样(SDE)、也能跑确定性采样(ODE)。
score 与噪声预测的关系是 $\nabla_x\log p_t(x) \approx -\epsilon_\theta(x,t)/\sqrt{1-\bar\alpha_t}$,代进去就得到一个只用模型输出的 ODE。关键一步是换坐标。 令 $\lambda = \log(\alpha_t/\sigma_t) = \frac{1}{2}\log\frac{\bar\alpha_t}{1-\bar\alpha_t}$(半 log-SNR)。
这个坐标为什么叫"半":信噪比本身就是 $\alpha_t^2/\sigma_t^2 = \bar\alpha_t/(1-\bar\alpha_t)$,取对数再除以二,得到的是"振幅比"的对数——也就是 $x_0$ 与 $\epsilon$ 两个分量在 $x_t$ 里的相对尺度的对数。$\lambda$ 越大代表越干净:$t=1000$ 时 $\bar\alpha_t\approx 4\times10^{-5}$,$\lambda\approx-5$;$t=1$ 时 $\bar\alpha_t\approx0.9999$,$\lambda\approx+4.6$。整条轨迹在 $\lambda$ 上只跨了不到 10 个单位,而 $t$ 跨了 1000 个单位——这就是"难度不是按 $t$ 均匀分布"的定量说法。
为什么要换?因为在这条 ODE 里,步长的含义由 $\lambda$ 决定:$t$ 上均匀的一步,落在 $\lambda$ 上的长度可能差几十倍。$\lambda$ 坐标把 ODE 拉成接近常数系数的形式($\alpha_\lambda$、$\sigma_\lambda$ 随 $\lambda$ 的变化是光滑的 sigmoid 型),指数积分才有干净的闭式。换成 $\lambda$ 之后,一步从 $\lambda_s$ 走到 $\lambda_t$(约定 $h=\lambda_t-\lambda_s>0$,因为我们从噪声走到干净)的精确解是
$$x_t = \frac{\sigma_t}{\sigma_s}x_s + \alpha_t\int_{0}^{h} e^{-(h-u)}\hat x_0(\lambda_s+u)du$$
这个式子值得停一下:它说明整步的解由"$\hat x_0$ 沿着 $\lambda$ 的变化曲线"加权积分决定,权重是 $e^{-(h-u)}$——越靠近区间右端(越干净的那端)权重越大。验证一下:如果 $\hat x_0$ 是常数,积分给出 $\alpha_t(1-e^{-h})\hat x_0$,加上第一项正好是 $\alpha_t \hat x_0+\sigma_t \epsilon$,也就是"一步跳到位"的精确解。所以步法阶数的本质,就是用多少个点上的 $\hat x_0$ 去逼近这条曲线。
把 $\hat x_0(\lambda_s+u)$ 在 $u=0$ 处展开,记 $\hat x_0,\hat x_0',\hat x_0''$ 为该点的一至二阶导,逐项积出系数:
$$x_t = \frac{\sigma_t}{\sigma_s}x_s + \alpha_t\Big[J_0\hat x_0 + J_1\hat x_0' + J_2\hat x_0'' + \cdots\Big]$$
$$J_0 = 1-e^{-h},\qquad J_1 = e^{-h}-1+h,\qquad J_2 = \frac{h^2-2h+2-2e^{-h}}{2}$$
小 $h$ 展开看数量级:$J_0\approx h$、$J_1\approx h^2/2$、$J_2\approx h^3/6$。所以导数项每高一阶,整步误差就多一个 $h$。
多步法(multistep)就是"不额外花评估,用前几步的 $\hat x_0$ 做差商估导数":设这一步的 $\lambda$ 为 $\lambda_{s_0}$,前两步在 $\lambda_{s_1}$、$\lambda_{s_2}$,令 $h_i=\lambda_{s_i}-\lambda_{s_{i+1}}$,则
$$D_{1,0} = \frac{h}{h_0}(\hat x_0^{s_0}-\hat x_0^{s_1}) \approx h\hat x_0',\qquad D_1 = D_{1,0} + \frac{r_0}{r_0+r_1}(D_{1,0}-D_{1,1})\approx h\hat x_0'$$
上一步的 $O(h)$ 偏差被组合系数消掉了(这就是"两个差商加权成一阶导"的标准手法),所以 $D_1$ 对 $\hat x_0'$ 是二阶准确的。同理二阶导用二阶差商:
$$D_2 = \frac{D_{1,0}-D_{1,1}}{r_0+r_1}$$
这里有个必须预先说清的细节。 三个点上的二阶差商等于 $\hat x_0''/2$(这是差商定义直接给出的),所以 $D_2\approx h^2\hat x_0''/2$;而上面泰勒系数 $\varphi_3$ 是按 $h^2\hat x_0''$ 的量纲推的。两者一比差一个 2。这个 2 在 05 节会展开:上游实现里单步版带了它、多步版没带,我先用受控实验把"该不该带"量出来,再决定怎么在文章里说。
把 03.1 的 DDIM 迭代($\eta=0$)和"只留 $J_0$"的一阶格式对照:$\eta=0$ 时第二项系数为 $\sqrt{1-\bar\alpha_{t-1}}$,第三项为 0,代入 $\hat x_0$ 的定义整理,得到
$$x_t = \frac{\sigma_t}{\sigma_s}x_s + \alpha_t(1-e^{-h})\hat x_0(\lambda_s)$$
与一阶格式逐项相同。所以"DDIM 是一阶方法"不是比喻,是恒等式——04 节用两套独立实现的代码把它验到 1e-14 以内。
全部代码在文末附录,五个脚本、只依赖 numpy:oracle_gmm.py(实验台与尺子)、ddim_family.py(η 家族)、dpm_solver_lab.py(λ 坐标与高阶格式)、spacing_lab.py(步数摆法)、make_figures.py(配图)。下面的数字是它们的真实输出。
要在二维上量"采样器差多少",需要一个模型误差为零的环境,否则量到的是模型不行,不是步法不行。用八个高斯分量摆在一圈上,加噪之后仍是高斯混合,score 有闭式解——模型换成这个解析 score,误差就只剩离散化。oracle_gmm.py 自己先做两件事:
A. Oracle 自检:解析 score vs 有限差分
t= 1 abar=0.999900 max|解析 - 差分| = 1.57e-08
t= 200 abar=0.659039 max|解析 - 差分| = 4.99e-11
t= 1000 abar=0.000040 max|解析 - 差分| = 2.14e-11
B. 端点检查:t=T 时协方差与 I 的最大差 = 3.834e-05
C. 尺子的噪声地板:SW1 真vs真 ×8 = 0.03055 ± 0.00452
D. 天花板:真样本当生成样本送进去 SW1 = 0.000000
第一项是"我写的解析 score 不是拍脑袋写的"(和有限差分对上);第二项确认 $t=T$ 时分布确实接近标准正态(否则起点选错了);第三项最重要:两批都是真样本时也能量出 0.03 的距离,这是抽样噪声的地板,后面低于该尺度的差别需要更多样本和置信区间确认,不能仅凭单次结果归因;第四项确认尺子无偏置。
ddim_family.py 的核心就是下面这个函数,变量名与 3.1 节一一对应:
def ddim_step(x, abar_cur, abar_tgt, eps, z, eta):
"""从 abar_cur 走到 abar_tgt(噪声变小)。返回 (x_tgt, x0_hat)。"""
a_c, s_c = alpha_sigma(abar_cur) # 当前时刻的 alpha_t / sigma_t
a_t, s_t = alpha_sigma(abar_tgt) # 目标时刻的 alpha_t / sigma_t
x0 = (x - s_c * eps) / a_c # 第一项用的 x0_hat
var = (1 - abar_tgt) / (1 - abar_cur) * (1 - abar_cur / abar_tgt)
sig_tilde = eta * np.sqrt(max(var, 0.0)) # 第三项:新掷噪声的幅度
direc = np.sqrt(max(s_t ** 2 - sig_tilde ** 2, 0.0)) # 第二项:勾股差
return a_t * x0 + direc * eps + sig_tilde * z, x0
最后一行就是 3.1 节那三项,一次加法写完,没有任何按时刻分支的特判。跑个小验证,看 $\tilde\sigma_t$ 在轨迹两端各占多少:
_, ABAR = linear_schedule()
ts, grid = uniform_t_grid(50) # 51 个 abar,grid[0] 最吵、grid[-1] = 1.0
x, zs = paired_noise(50, n=4, seed=4321) # 起点 + 每一步要用的随机数
print("x.shape =", x.shape, "| grid 长度 =", len(grid))
for i in (0, 25, 48):
for eta in (0.0, 1.0):
a_c, s_c = alpha_sigma(grid[i])
a_t, s_t = alpha_sigma(grid[i + 1])
var = (1 - grid[i+1]) / (1 - grid[i]) * (1 - grid[i] / grid[i+1])
sig_tilde = eta * np.sqrt(var)
print(f"step {i:2d} eta={eta:.1f} sigma_tilde = {sig_tilde:.6f}"
f" (该步总噪声幅度 {s_t:.4f})")
真实输出:
x.shape = (4, 2) | grid 长度 = 51
step 0 eta=0.0 sigma_tilde = 0.000000 (该步总噪声幅度 1.0000)
step 0 eta=1.0 sigma_tilde = 0.574284 (该步总噪声幅度 1.0000)
step 25 eta=0.0 sigma_tilde = 0.000000 (该步总噪声幅度 0.9509)
step 25 eta=1.0 sigma_tilde = 0.419845 (该步总噪声幅度 0.9509)
step 48 eta=0.0 sigma_tilde = 0.000000 (该步总噪声幅度 0.0760)
step 48 eta=1.0 sigma_tilde = 0.063819 (该步总噪声幅度 0.0760)
两个数值得盯一下。η=1 时 $\tilde\sigma_t$ 与该步总噪声幅度之比,在起点是 57%(0.574/1.000),到接近干净端的第 48 步反而是 84%(0.0638/0.0760)。 也就是说"加噪声"的代价在轨迹末端最重:那里 $\hat x_0$ 刚要把几个模式分辨开,一份占 84% 的新噪声就砸进去了。这是 4.6 节 η=0 处处最优的直接原因。
再看一步里三项谁大谁小(用 oracle score 走第 25 步):
eta=1 第 25 步三项 RMS: 回数据 0.1742 | 沿噪声方向 0.8694 | 新掷噪声 0.7800
两个噪声项的模长比"回数据"项大四五倍。三项向量的范数反映更新的组成,不能用向量大小推出网络算力花在哪里,而不是往数据方向推——这也解释了为什么"阶数"值钱:阶数讲的正是怎么更准地把这一步搬完。
ddim_family.py 里故意写了两套互不相干的代码:一套照 DDPM Algorithm 2 的原式走 1000 步(一步不能跳),一套照 DDIM 式 (12) 走。喂同样的随机数:
A. 终点逐元素最大差 = 3.109e-14 终点 RMS = 1.3520
3.1e-14 就是浮点累加误差的量级。η=1 的 DDIM 与 DDPM 祖采样是同一个算法,不是"近似"。
dpm_solver_lab.py 里同样对待 DDIM 与一阶指数积分器:
A. max|DDIM(eta=0) - DPM-Solver-1|
S= 20 uniform-t 8.549e-15 uniform-lambda 4.774e-15
S= 50 uniform-t 1.044e-14 uniform-lambda 1.088e-14
两套记号、两套时间网格,输出差在 1e-14。DDIM 是一阶指数积分器,这句到这里可以当结论用了。
有了参考解(多步三阶跑 4000 步,再用 2000 步自查,分辨率 1.16e-08)就能量阶数了。横轴必须是真实 NFE——脚本里用 Counter 包住模型,把调用次数数出来,而不是拿"步数 × 理论阶数"算。

这张图要看的是斜率:DDIM 的误差衰减阶约 1,二阶方法约 2,修正系数后的三阶方法接近 3;误差更小的曲线位于图的下方。纵向的差距就是"同样的评估次数,精度差多少个数量级"。
拟合出来:
| 采样器 | 拟合阶数 p | 备注 |
|---|---|---|
| DDIM(时间步等间隔) | 0.97 | 一阶 |
| DDIM(λ 等间隔) | 0.99 | 一阶,换坐标不改阶数 |
| DPM-Solver-2(单步) | 1.90 | 二阶,每步 2 次评估 |
| DPM-Solver++ 2M(多步) | 1.98 | 二阶,每步 1 次评估 |
| DPM-Solver++ 3M(原式) | 2.07 | 应该是三阶,实测只有二阶 |
| 3M(三阶项系数 ×2 后) | 2.91 | 接近三阶 |
3M 那两行是全篇唯一"和教科书不一致"的地方,见 05 节。
阶数是数值分析的语言,产品要的是"20 步够不够"。同一批配置换成 SW1(切片 Wasserstein-1,越小越好),换几份起点量出抖动:
| NFE | DDIM(t) | DDIM(λ) | Solver-2 | 2M | 3M(原式) |
|---|---|---|---|---|---|
| 10 | 0.1292±0.006 | 0.2180±0.006 | 0.1910±0.007 | 0.0731±0.006 | 0.0680±0.006 |
| 20 | 0.0741±0.004 | 0.1122±0.004 | 0.0568±0.004 | 0.0317±0.006 | 0.0277±0.005 |
| 50 | 0.0420±0.007 | 0.0516±0.007 | 0.0335±0.006 | 0.0328±0.005 | 0.0327±0.005 |
噪声地板是 0.0301±0.0026。请只看 NFE=20 那一行:20 次评估时各采样器之间的差距(0.0277 到 0.1122)远大于抖动;到了 50 次评估,所有方法都贴着地板(0.033~0.052),这张表就再也排不出名次了——不是"大家都一样好",是当前样本量和重复次数难以稳定区分,增加样本与重复实验仍可提高分辨力。要排名次得回到 4.4 的确定性误差表。
ddim_family.py 在 oracle 下扫一遍 η,再在欠拟合模型下扫一遍。所谓欠拟合是让模型以为数据是单个高斯(用真均值真协方差拟合),它在高噪声区几乎是对的、在低噪声区错得离谱——这是人为选择的一种误差结构,并不代表所有真实网络(也正好对应"训练时看到的是加噪数据分布"这件事)。

这张图要看的是两件事:一是曲线从左到右单调下降(步数越多越好,符合预期);二是本玩具大多数设置下 η=0 的 SW1 更低,不能推广为所有模型中确定性采样必胜,至少在有闭式解、且模型差得很有代表性的两种情况下都是错的。右图还给出一个诚实的例外:N=100 时 η=1 的 0.1830 略低于 η=0 的 0.1845,但差值 0.0015 远小于这批量的抖动,不该当成结论。

这张图要看的是:三条线在干净端都会收敛(步长趋于 0),但在中间段整段差 6.6 倍。缓存阈值是按这条线定的,换采样器就是换这条线的量级。同一份数据里 $x_0$ 估计的位移更夸张:η=0 是 0.0371,η=1 是 0.2407(6.5 倍)。

这张图要看的是柱子的相对高度,以及"步数"和"评估次数"不是一回事:DPM-Solver-2 每步要 2 次评估(先踩一步到中点、再走完),所以它的 NFE 是步数的两倍,这也是它虽然二阶却在低 NFE 段不占优的原因。
阶数回答的是"步数变多时误差掉多快",没回答"步数摆在哪"。spacing_lab.py 固定 NFE=20、只改摆法,终点误差如下:
| 摆法 | DDIM(一阶) | 2M(二阶) |
|---|---|---|
| 在 $t$ 上等间隔(本文的时间网格) | 1.107e-01 | 5.285e-02 |
| 在 $\lambda$ 上等间隔 | 1.713e-01 | 2.897e-02 |
| 在 $\lambda$ 上等间隔、落回整数 $t$ | 1.712e-01 | 2.894e-02 |
| Karras / EDM($\rho=7$) | 1.873e-01 | 7.038e-02 |
| Karras($\rho=3$) | 2.655e-01 | 3.397e-01 |
阶数和网格要配套:一阶方法在 $t$ 上等间隔最好,二阶方法在 $\lambda$ 上等间隔最好(差 1.8 倍)。而 EDM 那套 $\rho=7$ 的摆法在这里垫底——它不是不好,是它的 $\sigma$ 区间(EDM 原文是 80→0.002)和 VP 调度(这里只有 158→0.01,注意 EDM 的 $\sigma$ 是 $\sigma_t/\alpha_t$,不是 $\sqrt{1-\bar\alpha_t}$)不是一回事,照搬会把最吵的那一步拉得极长:NFE=20 时它最大的 $\lambda$ 间隔是 1.02,而 $\lambda$ 等间隔只有 0.51。
uniform-t 为什么赢,值得追问一层。 它把 $\lambda$ 的上界截在了 +1.76(NFE=20 时),剩下的干净端交给最后那一步"直接输出 $\hat x_0$"。这条捷径是不是免费的?不是。沿高精度参考轨迹量"在 $\lambda$ 处直接输出 $\hat x_0$ 会偏多少":
| $\lambda$ | −0.23 | +0.74 | +1.71 | +2.67 | +3.64 |
|---|---|---|---|---|---|
| 提前收尾的偏差 | 5.24e-01 | 1.40e-01 | 4.40e-02 | 9.39e-03 | 1.63e-03 |
截得越早白送的偏差越大,而且这份偏差不随步数下降——它只取决于截在哪个 $\lambda$。于是 $\lambda_{max}$ 存在最优值(2M,NFE=20,只改截断点):
| $\lambda_{max}$ | 1.0 | 2.0 | 2.5 | 3.0 | 4.61(走到底) |
|---|---|---|---|---|---|
| 终点误差 | 9.82e-02 | 2.80e-02 | 1.71e-02 | 1.81e-02 | 2.90e-02 |
U 形,最优在 2.5,比"走到底"好 1.7 倍。这条也解释了 uniform-t 的行为为什么还不错:它的截断点随步数自动往右挪(NFE=20 时 1.76、50 时 2.55、100 时 3.1),等于一个廉价的自适应策略——但它是巧合,不是设计。
顺带一个反直觉的观察。 一步的局部误差带着一个 $\alpha_t$ 因子($\delta_i\approx\alpha_{t_i}\frac{h^2}{2}\hat x_0'$),所以误差权重是 $\alpha_t\,|d\hat x_0/d\lambda|$ 而不是单纯的导数。量下来:噪声端 1.02e-01、干净端 1.00e-01,几乎一样重,峰在 $\lambda\approx-0.2$(也就是 $\bar\alpha\approx0.4$ 的中间段)。常听到的"把步数往干净端挪"在本实验里不成立——难度是按 $\lambda$ 均匀铺开的,真正要避免的是"某一段特别长"(Karras 就是栽在这)。
以 huggingface/diffusers 的 DPMSolverMultistepScheduler.step 为准(https://github.com/huggingface/diffusers/blob/main/src/diffusers/schedulers/scheduling_dpmsolver_multistep.py,以 2026-09 时的实现为准,上游会重构),最小实现与生产实现的差距集中在四处。
第一处:算法类型的开关是"预测什么"。 algorithm_type 有四个取值:dpmsolver 用噪声预测($\epsilon$-parameterization),dpmsolver++ 用数据预测($x_0$-parameterization);后缀 -sde-dpmsolver++ 则在同一条 ODE 上加回一个可控的噪声项。本实验包含噪声预测单步 Solver-2 与数据预测多步 Solver++。参数化会改变数值误差和引导稳定性,不能由训练 loss 的参数化权重图证明某条预测曲线总更光滑。
第二处:低步数时的降阶。 多步法开头没有历史可用,第一、二步必须降成一阶/二阶,lower_order_final 控制最后一步是否也降阶,final_sigmas_type="zero" 让最后一步直接落到 $\sigma=0$。这一条在我们的实现里对应"$\lambda_t=+\infty$ 时 $e^{-h}=0$、整式退化成 $x_t=\hat x_0$",不用特判。我专门验证过"开头降阶"是不是 3M 阶数不达标的原因:给第一步多加一次评估换成中点法(warm2s=True),实测阶数从 2.10 变成 2.10——一点没变,所以瓶颈不在起步。
第三处:时间步的排布。 timestep_spacing 有 leading、trailing、linspace(当前 DPMSolverMultistepScheduler 默认) 三种,同一个调度器换一种排布,20 步出图质量差很多。我们把它量化了:DDIM 在 $t$ 上等间隔时,NFE=191 的终点误差是 1.245e-02;换成 $\lambda$ 等间隔反而变差到 1.775e-02,达到 1e-2 分别需要 256 次和 384 次评估。这是"坐标选择比阶数更早起作用"的直接证据:一阶方法的误差由最大的那一段步长决定,均匀 $\lambda$ 的 $h$ 本来就是相等的;它改变了各时间段的评估密度,误差还由场的导数、传播与终端截断决定。所以别把"回到 λ 坐标"当成万能钥匙——它给高阶方法提供了干净的积分形式,但不自动给一阶方法更好的网格。
第四处,也是唯一需要标注存疑的:三阶项的系数。 我们逐行对齐 diffusers 的多步三阶更新后,实测阶数只有 2.07(4.4 节),不是 3。于是做了一个受控实验:不让 $\hat x_0$ 由 $x$ 决定,而是直接规定它是 $\lambda$ 的二次多项式,此时精确解可以用数值积分算出来,任何真正的三阶格式都应该一步算准。把区间长度 $h$ 从 0.4 缩到 0.025:
一阶: 4.17e-02 1.04e-02 2.57e-03 6.39e-04 1.59e-04 → h^2 ✓
二阶: 8.16e-03 1.05e-03 1.32e-04 1.65e-05 2.07e-06 → h^3 ✓
三阶-D2原式: 1.96e-03 2.46e-04 3.06e-05 3.81e-06 4.75e-07 → h^3 ✗
三阶-D2乘2: 4.43e-13 6.23e-14 2.89e-15 1.11e-15 4.44e-16 → 一步精确 ✓
一阶、二阶的收敛阶都对得上公式推导;三阶原式只有 $h^3$(等于二阶),把 $D_2$ 乘 2 之后直接掉到 1e-13~1e-16 的机器精度。我还用非均匀步长($h_0\neq h_1\neq h$)复验过一遍,结论不变:原式 8.3e-04、乘 2 之后 5.4e-14,所以这不是"等间隔时才碰巧"的巧合。
回头逐字核对上游源码(diffusers v0.30.0):单步版 scheduling_dpmsolver_singlestep.py 里写的是 D2 = 2.0 * (D1_1 - D1_0) / (r0 - r1),带了那个 2;多步版 scheduling_dpmsolver_multistep.py 里写的是 D2 = (1.0 / (r0 + r1)) * (D1_0 - D1_1),没带。两个版本共用同一个 $\varphi_3$ 系数 - (alpha_t * ((exp(-h) - 1.0 + h) / h**2 - 0.5)) * D2。
我把这件事的处理方式说清楚:这篇文章只报告我们自己的复现结果——逐行对齐上游的多步三阶实现、实测阶数 2.07;把 $D_2$ 乘 2 后升到 2.91,局部误差检验一步算准。至于这是上游实现的笔误、还是我对多步差商归一化的理解和原作者不同,我不下结论,已记入待决清单单独确认。对读者的实用含义是确定的:不要假设"用了 3M 就是三阶",阶数要自己量。
省了什么。 DDIM 把每步的随机项去掉,换来三件事:同样的步数下误差更小(4.6 节的两个模型都验证了)、轨迹光滑(缓存加速可用)、以及固定初始噪声时可复现的输出。正则连续 ODE 流可逆,有限步 DDIM 或终步投影并不自动保证一一对应。
赔了什么。
solver_type="midpoint" / "heun" 选项;它们不能和作者单步 API 中的 dpmsolver / taylor 选项混用都是在拿稳定性换名义阶数。什么时候不该用确定性采样器。 需要随机性作为正则的场景(例如低步数下用 SDE 采样换取更好的分布覆盖、或者需要"温度"调节多样性),$x_0$ 强约束会导致过平滑;以及任何依赖"步步重掷噪声"来做布朗桥类操作的训练/蒸馏流程。
什么时候值得上高阶。 评估预算在 20~50 次这个区间时最划算(4.8 节:NFE=20 时 3M 已经压到 1e-2,DDIM 需要 256 次);一旦预算到几百次,所有方法都进入渐进区,选最便宜的一阶反而更省心。
误解一:DDIM 就是步幅更大的 DDPM。 不是。DDIM 可以走 $t$ 上的任意子序列(这是它"能跳步"的前提,不是它的内容),它的内容是去掉随机项。反过来也成立:确定性 DDIM(η=0)也可以走 1000 步,仍不同于随机 DDPM;η=1 配合对应完整时间表时才有前文的等价关系。把这两件事分开之后,"为什么 DDIM 的 1000 步也不等于 DDPM"这个问题才有答案。
误解二:步数少了,就把 η 调大一点补回来。 实测反了。oracle 下 η=0 在 N=10/20/50/100 全部最优(N=10 时 0.1321 对 0.1888);欠拟合模型下同样单调(N=10 时 0.2679 对 0.3496)。直觉为什么错:加噪声确实让每步的 x̂₀ 估计被"抹匀"一点,但代价是它也抹掉了本该累积的方向信息;而噪声项还会让下一步的输入抖动变大(6.6 倍那个数字),把误差一路带下去。
误解三:阶数越高越省。 一是不一定真拿到那个阶(4.4 节 3M 实测 2.07);二是低 NFE 段高阶方法要"攒历史",前几步被迫降阶,实测 NFE=10 时 2M(0.0731)确实明显好于 DDIM(0.1292),但 3M(0.0680)比 2M 只好一点点,收益远小于"二阶到三阶"的名义差距。
误解四:SW1/FID 分不出来就没差别。 4.5 节里 50 次评估时全部方法都贴在 0.031 附近,看着"都一样"——那是因为尺子的噪声地板就是 0.0301。指标贴着地板时不代表方法等价,只代表这个指标失效了,此时要换确定性指标(有真值时用 RMS,我们的表就是这么打的)。
误解五:换到 λ 坐标只会更好。 不保证。λ 坐标便于推导指数积分,但选择 λ 均匀网格是另一件事;本文的一阶实验反而更偏好 t 网格:DDIM 在 $\lambda$ 等间隔网格上误差更大(NFE=191 时 1.775e-02,而 $t$ 等间隔是 1.245e-02)。网格与阶数要配套选。
五个脚本都能直接跑,几分钟内出结果。建议按这个顺序试:
cd outputs/fundamentals_files/ddim_samplers/code
/usr/local/bin/python3 oracle_gmm.py # ① 标定尺子(先量地板)
/usr/local/bin/python3 ddim_family.py # ② η 旋钮与缓存抖动
/usr/local/bin/python3 dpm_solver_lab.py # ③ 阶数、预算、三阶项的受控检验
/usr/local/bin/python3 spacing_lab.py # ④ 步数摆法与收尾点
/usr/local/bin/python3 make_figures.py # ⑤ 画图
四个具体的改动实验,附我这里跑出来的结果:
ddim_family.py 的 D 节会打印"输入相对变化",η=0 时最后 10 步是 0.0197,η=0.5 是 0.0667,η=1 是 0.1294。你会看到"加一点点噪声"的代价比想象中大。dpm_solver_lab.py 里 dpm_pp_third 的 d2_scale 从 1.0 改成 2.0:C 节里 3M 的拟合阶数会从 2.07 变成 2.91,F 节的局部误差从 4.75e-07(最后一行)变成 4.44e-16。这一个常数就是"名义三阶"和"实测三阶"的分界。DDIM 的网格从 uniform_t_grid 换成 uniform_lambda_grid:C 节里 DDIM 那两行的阶数不变(0.97 / 0.99),但误差整体上移,达到 1e-2 所需的 NFE 从 256 涨到 384。"同阶但更差"这件事,只有量出来才看得见。spacing_lab.py 里把 grid_uniform_lambda_capped(20, cap) 的 cap 从 4.61 改成 2.5:E2 节那张 U 形表上,2M 的终点误差会从 2.90e-02 掉到 1.71e-02。这个数不用换模型、不用换采样器,只改一行网格。ddpm:DDPM 的训练目标与采样流程。本篇的前置——$\hat x_0$ 与 $\epsilon_\theta$ 的互换关系、$\bar\alpha_t$ 的调度都在那里定义。diffusion_math:扩散过程的前向与反向推导。概率流 ODE 与 SDE 的关系、参数化选择对误差权重的影响,是理解 05 节 algorithm_type 开关的基础。vae_elbo:变分下界与重参数化。$\eta$ 旋钮的合法性来自"边缘分布不变",那套记号在那里建立。cfg:分类器无关引导。引导会放大 ODE 右端的 Lipschitz 常数,采样步法要跟着调,是 06 节"稳定域"的直接后果。flow_matching:流匹配与 Rectified Flow。把"直线路径"当作先验,是这条 ODE 主线的另一个分支。09 节用到的脚本全文如下(oracle_gmm.py、ddim_family.py、dpm_solver_lab.py、spacing_lab.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。
# -*- coding: utf-8 -*-
"""采样器实验的地基:一个 score 有闭式解的二维高斯混合 + 一把「尺子」。
为什么必须先把这两件事钉死:
1. 比较采样器时,最大的干扰项是**模型误差**。同一个噪声网络,在 t=1 和 t=1000
上的误差可以差两个数量级,你根本分不清「这个采样器好」和「这个时间步好学」。
用闭式 score(oracle),模型误差恒为 0,剩下的全是离散化误差。
2. 比较采样器时,第二大的干扰项是**尺子本身的噪声**。切片 Wasserstein 距离
是用有限样本估的,两个「都对」的分布之间也能量出一个正数。这个正数就是
噪声地板,所有小于它的差别都是运气。
所以这个脚本干三件事:给出 oracle、给出尺子、量出地板。后面三个脚本全部
import 它,保证四处用到的目标分布、指标、随机数是同一套。
/usr/local/bin/python3 oracle_gmm.py
依赖:numpy, matplotlib(画图时才需要)
"""
import numpy as np
T = 1000 # 训练时的扩散步数
D = 2 # 数据维度(二维才画得出轨迹,也才跑得起闭式 score)
N_EVAL = 4000 # 评价用的样本数
N_DIRS = 512 # 切片 Wasserstein 的投影方向数
# ══════════════════════════════════════════════════════════════════
# 1. 噪声调度
# ══════════════════════════════════════════════════════════════════
def linear_schedule(T: int = T, b0: float = 1e-4, b1: float = 0.02):
"""DDPM / diffusers 默认的线性 beta 调度。
返回 abar 的长度是 T+1:abar[t] = prod_{s<=t} (1 - beta_s),
abar[0] = 1 表示「完全干净」。这样下标和论文里的 t 直接对齐,
不用在 0-based / 1-based 之间来回换算。
"""
betas = np.linspace(b0, b1, T)
abar = np.concatenate([[1.0], np.cumprod(1.0 - betas)])
return betas, abar
def alpha_sigma(abar_t: float):
"""给定 abar_t,返回 (alpha_t, sigma_t) = (sqrt(abar), sqrt(1-abar))。
这是全文唯一的「坐标」:x_t = alpha_t * x_0 + sigma_t * eps。
注意 alpha 是**信号尺度**,不是 DDPM 论文里的 alpha_t = 1 - beta_t。
"""
return np.sqrt(abar_t), np.sqrt(1.0 - abar_t)
# ══════════════════════════════════════════════════════════════════
# 2. 目标分布:8 个切向拉长的高斯,摆成一个环
# ══════════════════════════════════════════════════════════════════
def make_target(K: int = 8, radius: float = 1.8, tang: float = 0.30, rad: float = 0.05):
"""K 个分量均匀摆在一个半径 radius 的环上,每个分量沿切线方向拉长。
为什么不用各向同性:各向同性的高斯混合,加噪之后很快变成一个圆球,
score 场几乎是线性的,一阶方法就能解到机器精度,分不出高下。
切向拉长之后,环上的 score 场有明显曲率,离散化误差才显出来。
"""
ang = np.linspace(0.0, 2.0 * np.pi, K, endpoint=False)
mu = np.stack([radius * np.cos(ang), radius * np.sin(ang)], axis=1)
Sig = np.zeros((K, D, D))
for k, a in enumerate(ang):
R = np.array([[np.cos(a), -np.sin(a)], [np.sin(a), np.cos(a)]])
Sig[k] = R @ np.diag([tang, rad]) @ R.T
w = np.ones(K) / K
return w, mu, Sig
def sample_true(n: int, w, mu, Sig, rng):
"""从目标分布抽样。"""
k = rng.choice(len(w), size=n, p=w)
L = np.linalg.cholesky(Sig[k]) # (n, D, D)
z = rng.standard_normal((n, D))
return mu[k] + np.einsum("nij,nj->ni", L, z)
# ══════════════════════════════════════════════════════════════════
# 3. Oracle:加噪后分布的闭式 score,以及最优的 eps 预测
# ══════════════════════════════════════════════════════════════════
def _cov_at(abar_t: float, Sig):
"""q(x_t) 的第 k 个分量协方差:C_k = abar * Sig_k + (1 - abar) * I。"""
return abar_t * Sig + (1.0 - abar_t) * np.eye(D)
def logp_and_score(x, abar_t: float, w, mu, Sig):
"""返回 (log q_t(x), score_t(x)),x 形状 (n, D)。
q_t 是 K 个高斯的混合:第 k 个分量的均值是 sqrt(abar) * mu_k,
协方差是 abar * Sig_k + (1 - abar) * I。score 是这个混合分布的
对数梯度 = 各分量 score 的后验加权平均。
"""
C = _cov_at(abar_t, Sig) # (K, D, D)
inv = np.linalg.inv(C)
ld = np.log(np.linalg.det(C))
r = x[:, None, :] - np.sqrt(abar_t) * mu[None] # (n, K, D)
quad = np.einsum("nki,kij,nkj->nk", r, inv, r)
lg = np.log(w)[None, :] - 0.5 * quad - 0.5 * ld[None, :]
m = lg.max(axis=1, keepdims=True)
p = np.exp(lg - m)
Z = p.sum(axis=1, keepdims=True)
logp = (m[:, 0] + np.log(Z[:, 0]))
pi = p / Z # 后验分量权重
comp = -np.einsum("kij,nkj->nki", inv, r) # 每个分量的 score
score = np.einsum("nk,nki->ni", pi, comp)
return logp, score
def eps_star(x, abar_t: float, w, mu, Sig):
"""MMSE 意义下最优的 eps 预测:eps* = -sigma_t * score_t(x)。
因为 x_t = alpha x_0 + sigma * eps,对 x_t 求 log 梯度会得到
grad log q_t = -E[eps | x_t] / sigma,所以最优的噪声预测就是
score 乘一个 -sigma。后面所有采样器都拿它当「模型」。
"""
_, s = logp_and_score(x, abar_t, w, mu, Sig)
return -np.sqrt(1.0 - abar_t) * s
def x0_from_eps(x, abar_t: float, eps):
"""由 x_t 和 eps 反解 x_0:x_0 = (x_t - sigma_t * eps) / alpha_t。"""
a, sg = alpha_sigma(abar_t)
return (x - sg * eps) / a
# ══════════════════════════════════════════════════════════════════
# 4. 尺子:切片 Wasserstein-1 + 平均对数密度
# ══════════════════════════════════════════════════════════════════
def make_dirs(n_dirs: int = N_DIRS, seed: int = 7):
"""预先抽好投影方向,全篇共用同一组。
这一条很重要:每次评价都重新抽方向的话,两组「同样好」的样本也会因为
方向不同量出不同的数,这部分方差完全是无谓的。方向固定之后,
A 和 B 的差别就只来自样本本身。
"""
th = np.random.default_rng(seed).standard_normal((n_dirs, D))
th /= np.linalg.norm(th, axis=1, keepdims=True)
return th
def sliced_w1(X, Y, TH=None, n_dirs: int = N_DIRS, rng=None):
"""切片 W1:在 n_dirs 个随机方向上算一维 W1,再取平均。
一维 W1 有闭式解:把两个样本集投影后排序,逐位相减取绝对值平均。
单位和数据同单位(这里是「平均要挪多远」),比 MMD^2 好解释。
"""
if TH is None:
TH = make_dirs(n_dirs) if rng is None else None
if TH is None:
th = rng.standard_normal((n_dirs, X.shape[1]))
th /= np.linalg.norm(th, axis=1, keepdims=True)
else:
th = TH
a = np.sort(X @ th.T, axis=0)
b = np.sort(Y @ th.T, axis=0)
return float(np.abs(a - b).mean())
def mean_logp(X, w, mu, Sig):
"""生成样本在真实分布下的平均对数密度。"""
lp, _ = logp_and_score(X, 1.0, w, mu, Sig)
return float(lp.mean())
def evaluate(X, Xref, w, mu, Sig, lp_ref, TH):
"""一次评价,返回 (SW1, dlogp)。
dlogp = 生成样本的平均 log 密度 − 真实样本的平均 log 密度。
0 表示完美;负数表示生成样本落在了真实分布的「空区」。
"""
return sliced_w1(X, Xref, TH=TH), mean_logp(X, w, mu, Sig) - lp_ref
# ══════════════════════════════════════════════════════════════════
# 5. 主流程:验 oracle、量尺子的噪声地板
# ══════════════════════════════════════════════════════════════════
def main():
rng = np.random.default_rng(20260927)
betas, abar = linear_schedule()
w, mu, Sig = make_target()
print("=" * 68)
print("A. Oracle 自检:解析 score vs 有限差分")
print("=" * 68)
for t in (1, 50, 200, 500, 1000):
x = sample_true(8, w, mu, Sig, rng)
_, s = logp_and_score(x, abar[t], w, mu, Sig)
h = 1e-5
fd = np.zeros_like(s)
for i in range(D):
e = np.zeros(D)
e[i] = h
lp_p, _ = logp_and_score(x + e, abar[t], w, mu, Sig)
lp_m, _ = logp_and_score(x - e, abar[t], w, mu, Sig)
fd[:, i] = (lp_p - lp_m) / (2 * h)
err = np.abs(fd - s).max()
print(f" t={t:5d} abar={abar[t]:.6f} max|解析 - 差分| = {err:.2e}")
print(" 差分误差在 1e-7 量级 → 解析 score 是对的(不是我拍脑袋写的)")
# 顺便看一眼:oracle eps 在 t=1000 处还剩多少信息
print()
print("=" * 68)
print("B. 端点检查:t=T 时分布离标准正态有多远")
print("=" * 68)
for t in (200, 500, 800, 1000):
C = _cov_at(abar[t], Sig)
dev = np.abs(C - np.eye(D)).max()
print(f" t={t:5d} abar={abar[t]:.3e} "
f"各分量协方差与 I 的最大差 = {dev:.3e}")
print(" 越接近 t=T,加噪后的混合越接近一个标准正态球")
# ── 尺子的噪声地板 ──
print()
print("=" * 68)
print("C. 尺子的噪声地板:两批「都对」的样本之间也能量出距离")
print("=" * 68)
Xref = sample_true(N_EVAL, w, mu, Sig, rng)
lp_ref = mean_logp(Xref, w, mu, Sig)
print(f" 参考样本: {Xref.shape}, 真实样本自己的平均 logp = {lp_ref:.4f}")
TH = make_dirs()
sw_floor, lp_floor = [], []
for r in range(8):
Y = sample_true(N_EVAL, w, mu, Sig, np.random.default_rng(1000 + r))
sw_floor.append(sliced_w1(Y, Xref, TH=TH))
lp_floor.append(mean_logp(Y, w, mu, Sig))
print(f" SW1 真vs真 ×8: 均值 {np.mean(sw_floor):.5f} "
f"标准差 {np.std(sw_floor):.5f} 最大 {np.max(sw_floor):.5f}")
print(f" dlogp 真vs真 ×8: 均值 {np.mean(lp_floor) - lp_ref:+.4f} "
f"标准差 {np.std(lp_floor):.4f}")
print(" → 后面所有表格里,小于这个数的差别都是抽样噪声,不是采样器的功劳")
# ── 一个有参照的好答案长什么样 ──
print()
print("=" * 68)
print("D. 天花板:把真样本直接当生成样本送进去")
print("=" * 68)
print(f" SW1 = {sliced_w1(Xref, Xref, TH=TH):.6f}"
f" dlogp = {mean_logp(Xref, w, mu, Sig) - lp_ref:+.4f}")
print(" (恒等于 0,说明尺子本身没有偏置)")
if __name__ == "__main__":
main()
# -*- coding: utf-8 -*-
"""DDIM 家族:一个 eta 旋钮,把 DDPM 祖采样和 DDIM 串成一条线。
这个脚本回答三个问题:
1. **eta=1 的 DDIM 是不是就是 DDPM 祖采样?** 用两套独立的代码(一套写
DDPM Algorithm 2 的原式,一套写 DDIM 论文式 (12)),喂同样的随机数和
同样的起点,看输出能不能对到最后一位。能,那这条家族关系就不是传说。
2. **eta 到底该调多大?** 在 oracle score 下扫一遍:步数少的时候是不是
「加噪声更好」?在欠拟合模型下再扫一遍,看结论翻不翻。
3. **换采样器为什么缓存策略得重调?** 量每步「模型输入变了多少」:
确定性采样器的轨迹是光滑的,随机采样器每一步都往里砸一份新噪声。
/usr/local/bin/python3 ddim_family.py
依赖:numpy(+ 同目录的 oracle_gmm.py)
"""
import numpy as np
from oracle_gmm import (T, linear_schedule, make_target, sample_true,
eps_star, x0_from_eps, alpha_sigma, sliced_w1,
mean_logp, make_dirs, N_EVAL, N_DIRS)
# ══════════════════════════════════════════════════════════════════
# 0. 时间步网格与模型
# ══════════════════════════════════════════════════════════════════
def uniform_t_grid(S: int, T: int = T):
"""DDIM 论文 / diffusers 默认的「leading」取法:在 t 上等间隔。
返回 S+1 个 abar:起点 abar[T],终点 1.0(完全干净)。
"""
ts = np.round(np.linspace(0.0, T, S + 1)).astype(int)[::-1]
grid = np.array([1.0 if t == 0 else _ABAR[t] for t in ts], dtype=float)
return ts, grid
def paired_noise(S: int, n: int, seed: int = 1234):
"""给所有采样器**同一份**起点和同一批随机数。
这是整套实验最要紧的一条:如果每个配置各抽各的起点,配置之间就多了一份
抽样方差,量出来的差别里有多少是采样器的功劳根本说不清。配对之后,
eta=0 和 eta=1 的终点之差只能来自算法本身。
"""
r = np.random.default_rng(seed + S)
x_init = r.standard_normal((n, 2))
zs = [r.standard_normal((n, 2)) for _ in range(S)]
return x_init, zs
def make_model(kind: str, w, mu, Sig):
"""两种「模型」,用来分离离散化误差和模型误差。
oracle : 用的就是真 score,模型误差 = 0,剩下的全是离散化误差。
gauss : 模型以为数据是一个单高斯(用真均值真协方差拟合),
也就是「容量不足以表达多峰」。高噪声区它几乎是对的,
低噪声区它错得离谱 —— 这正是真实网络的误差结构。
"""
if kind == "oracle":
return lambda x, abar: eps_star(x, abar, w, mu, Sig)
if kind == "gauss":
m = (w[:, None] * mu).sum(axis=0)
S = ((w[:, None, None] * (Sig + np.einsum("ki,kj->kij", mu, mu))).sum(axis=0)
- np.outer(m, m))
w1, mu1, S1 = np.array([1.0]), m[None, :], S[None, :, :]
return lambda x, abar: eps_star(x, abar, w1, mu1, S1)
raise ValueError(kind)
# ══════════════════════════════════════════════════════════════════
# 1. DDIM 的一步(论文式 12,与 diffusers DDIMScheduler.step 逐项对齐)
# ══════════════════════════════════════════════════════════════════
def ddim_step(x, abar_cur, abar_tgt, eps, z, eta):
"""从 abar_cur 走到 abar_tgt(abar_tgt > abar_cur,噪声变小)。
x_tgt = sqrt(abar_tgt) * x0_hat
+ sqrt(sigma_tgt^2 - sigma_tilde^2) * eps ← 指向 x_t 的方向项
+ sigma_tilde * z ← 随机项
sigma_tilde = eta * sqrt( (1-abar_tgt)/(1-abar_cur) * (1 - abar_cur/abar_tgt) )
eta = 0 → 完全没有随机项,DDIM;
eta = 1 → 随机项等于 DDPM 的后验标准差 beta_tilde。
"""
a_c, s_c = alpha_sigma(abar_cur)
a_t, s_t = alpha_sigma(abar_tgt)
x0 = (x - s_c * eps) / a_c
var = (1.0 - abar_tgt) / (1.0 - abar_cur) * (1.0 - abar_cur / abar_tgt)
sig_tilde = eta * np.sqrt(max(var, 0.0))
direc = np.sqrt(max(s_t ** 2 - sig_tilde ** 2, 0.0))
return a_t * x0 + direc * eps + sig_tilde * z, x0
def ddpm_step(x, abar_cur, abar_tgt, eps, z, betas):
"""DDPM Algorithm 2 的原式,故意写成「另一套代码」用来交叉验证。
x_{t-1} = (x_t - beta_t / sqrt(1 - abar_t) * eps) / sqrt(alpha_t)
+ sqrt(beta_tilde_t) * z
只在相邻步(stride = 1)上成立,不能跳步。
"""
abar_prev = abar_tgt
beta_t = 1.0 - abar_cur / max(abar_prev, 1e-12) # = 1 - alpha_t
beta_tilde = (1.0 - abar_prev) / (1.0 - abar_cur) * beta_t
a_t = np.sqrt(max(abar_cur / max(abar_prev, 1e-12), 1e-12))
mean = (x - beta_t / np.sqrt(1.0 - abar_cur) * eps) / a_t
return mean + np.sqrt(max(beta_tilde, 0.0)) * z
def run_ddim(x, grid, model, eta, zs, track=False):
"""沿 grid 走一遍。grid[0] 最吵,grid[-1] = 1.0。
track=True 时额外返回每一步的 (x, x0_hat),D 节量「轨迹抖动」要用。
"""
xs, x0s = [x.copy()], []
for i in range(len(grid) - 1):
eps = model(x, grid[i])
z = zs[i] if zs is not None else 0.0
x, x0 = ddim_step(x, grid[i], grid[i + 1], eps, z, eta)
if track:
xs.append(x.copy())
x0s.append(x0.copy())
return x, (xs, x0s) if track else x0s
# ══════════════════════════════════════════════════════════════════
# 2. 主流程
# ══════════════════════════════════════════════════════════════════
_BETAS, _ABAR = linear_schedule()
def main():
rng = np.random.default_rng(20260927)
w, mu, Sig = make_target()
Xref = sample_true(N_EVAL, w, mu, Sig, rng)
lp_ref = mean_logp(Xref, w, mu, Sig)
ev_rng = np.random.default_rng(7)
model = make_model("oracle", w, mu, Sig)
# ── A. eta=1 的 DDIM 是不是 DDPM 祖采样 ──
print("=" * 70)
print("A. 交叉验证:DDIM(eta=1, stride=1) 是否等于 DDPM Algorithm 2")
print("=" * 70)
S = 1000
ts, grid = uniform_t_grid(S)
n = 400
x_init = rng.standard_normal((n, 2))
zs = [rng.standard_normal((n, 2)) for _ in range(S)]
x_ddpm = x_init.copy()
for i in range(S): # 原式,一步步走完 1000 步
t = ts[i]
x_ddpm = ddpm_step(x_ddpm, grid[i], grid[i + 1],
model(x_ddpm, grid[i]), zs[i], _BETAS)
x_ddim, _ = run_ddim(x_init.copy(), grid, model, eta=1.0, zs=zs)
d = np.abs(x_ddpm - x_ddim).max()
print(f" 两条独立实现,喂同样的 {S} 份随机数")
print(f" 终点逐元素最大差 = {d:.3e} 终点 RMS = {np.sqrt((x_ddim**2).mean()):.4f}")
print(f" → 差 {d:.1e},是浮点累加误差量级:eta=1 就是 DDPM 祖采样,不是「近似」")
# ── B. eta 扫描(oracle score)──
print()
print("=" * 70)
print("B. eta 扫描:步数越少,是不是越该加噪声?(oracle score,配对起点)")
print("=" * 70)
TH = make_dirs()
floors = [sliced_w1(sample_true(N_EVAL, w, mu, Sig,
np.random.default_rng(1000 + r)), Xref, TH=TH)
for r in range(6)]
print(f" 噪声地板 SW1 = {np.mean(floors):.4f} ± {np.std(floors):.4f}"
f"(真实样本对真实样本),越接近它越好")
NS = (10, 20, 50, 100)
print(f" {'eta':>5} " + " ".join(f"{'N=' + str(s):>10}" for s in NS))
for eta in (0.0, 0.25, 0.5, 0.75, 1.0):
row = []
for S in NS:
ts, grid = uniform_t_grid(S)
x_init, zs = paired_noise(S, N_EVAL)
x, _ = run_ddim(x_init, grid, model, eta, zs)
row.append(sliced_w1(x, Xref, TH=TH))
print(f" {eta:5.2f} " + " ".join(f"{v:10.4f}" for v in row))
# ── C. 换个欠拟合的模型,结论翻不翻 ──
print()
print("=" * 70)
print("C. 同样的扫描,但把模型换成「以为数据是单高斯」")
print("=" * 70)
gm = make_model("gauss", w, mu, Sig)
print(f" {'eta':>5} " + " ".join(f"{'N=' + str(s):>10}" for s in NS))
for eta in (0.0, 0.25, 0.5, 0.75, 1.0):
row = []
for S in NS:
ts, grid = uniform_t_grid(S)
x_init, zs = paired_noise(S, N_EVAL)
x, _ = run_ddim(x_init, grid, gm, eta, zs)
row.append(sliced_w1(x, Xref, TH=TH))
print(f" {eta:5.2f} " + " ".join(f"{v:10.4f}" for v in row))
# ── D. 每步「模型输入变了多少」:缓存友好度 ──
print()
print("=" * 70)
print("D. 轨迹抖动:每步「模型输入」变了百分之几(缓存复用看的就是它)")
print("=" * 70)
print(" 缓存类加速(按『这一步的输入和上一步差不多就复用上一步的特征』决策)")
print(" 盯的是这个量。它一旦被噪声项垫住,原阈值就全废了。")
print(f" {'eta':>5} {'最后10步 输入相对变化':>22} {'最后10步 x0_hat位移':>22}")
S = 50
ts, grid = uniform_t_grid(S)
x_init, zs = paired_noise(S, N_EVAL)
out = {}
for eta in (0.0, 0.5, 1.0):
_, (xs, x0s) = run_ddim(x_init.copy(), grid, model, eta, zs, track=True)
rel_x = [float(np.sqrt((((xs[i + 1] - xs[i]) ** 2).sum(1)).mean())
/ np.sqrt(((xs[i] ** 2).sum(1)).mean()))
for i in range(len(xs) - 1)]
d_x0 = [float(np.sqrt((((x0s[i] - x0s[i - 1]) ** 2).sum(1)).mean()))
for i in range(1, len(x0s))]
out[eta] = (rel_x, d_x0)
print(f" {eta:5.1f} {np.mean(rel_x[-10:]):22.4f} {np.mean(d_x0[-10:]):22.4f}")
r0, r1 = out[0.0][0][-10:], out[1.0][0][-10:]
print(f" → eta=1 的输入抖动是 eta=0 的 {np.mean(r1) / np.mean(r0):.1f} 倍;"
f"按 eta=0 调出来的缓存阈值,换个采样器就完全不是一回事")
if __name__ == "__main__":
main()
# -*- coding: utf-8 -*-
"""从 DDIM 走到高阶:把时间轴换成 lambda,阶数就变成可以直接量出来的东西。
这个脚本做四件事:
1. **验证 DDIM 就是一阶 DPM-Solver。** 两套完全不同的代码(一套写 DDIM
论文式 12,一套写 lambda 空间的一阶指数积分器),同样的时间网格、
同样的模型,看输出是不是同一个数。是,那「DDIM 是一阶方法」就不是比喻。
2. **把阶数量出来。** 拿一条超高精度的参考解当真值,量各采样器终点到它的
距离。这个距离没有抽样噪声(起点配对、采样器确定),能干净地跨几个数量级,
于是「误差 ~ NFE^(-p)」里的 p 可以直接拟合出来。
3. **给出实用的步数-质量表。** 阶数是数值分析的语言,产品要的是
「20 步够不够」,所以再给一张 SW1 表(含噪声地板)。
4. **看高阶方法省下的评估次数到底落在哪。**
/usr/local/bin/python3 dpm_solver_lab.py
依赖:numpy(+ 同目录的 oracle_gmm.py / ddim_family.py)
"""
import numpy as np
from oracle_gmm import (T, linear_schedule, make_target, sample_true,
eps_star, alpha_sigma, sliced_w1, make_dirs, N_EVAL)
from ddim_family import uniform_t_grid, make_model, ddim_step
_BETAS, _ABAR = linear_schedule()
_LAM = 0.5 * np.log(np.maximum(_ABAR, 1e-300) / np.maximum(1.0 - _ABAR, 1e-300))
# ══════════════════════════════════════════════════════════════════
# 0. lambda 坐标与网格
# ══════════════════════════════════════════════════════════════════
def lam_of_abar(abar):
"""half log-SNR:lambda = log(alpha) - log(sigma) = 0.5 * log(abar/(1-abar))。"""
return 0.5 * np.log(np.maximum(abar, 1e-300) / np.maximum(1.0 - abar, 1e-300))
def abar_of_lam(lam):
"""lambda 的反函数:abar = e^(2 lam) / (1 + e^(2 lam))。"""
e = np.exp(2.0 * lam)
return e / (1.0 + e)
def nearest_t(abar):
"""把连续的 abar 落到最近的训练时间步上(真实网络只认整数 t)。
注意 abar 是**递减**的,不能用 searchsorted(它只认递增数组),
老老实实找绝对值最小的那个下标。
"""
return int(np.clip(int(np.argmin(np.abs(_ABAR - abar))), 1, T))
def uniform_lambda_grid(S: int, quantize: bool = False):
"""在 lambda 上等间隔取 S 个点,再落回训练网格。
返回 S+1 个 abar:前 S 个是「要评估模型」的位置,最后一个是 1.0
(完全干净,sigma = 0,对应 lambda = +inf,只能用一阶收尾)。
quantize=True 会把每个 lambda 落回整数时间步。落回去之后可能撞车
(两个 lambda 落到同一个 t),撞车会让多步法的差商分母变成 0,
所以强制 t 严格递减。参考解要更细的分辨率,用 quantize=False 直接
吃连续的 abar —— oracle 认得任意 abar,真实网络才需要整数。
"""
l0, l1 = _LAM[T], _LAM[1] # 从最吵到最干净
lams = np.linspace(l0, l1, S)
if not quantize:
return None, np.concatenate([abar_of_lam(lams), [1.0]])
ts = []
for l in lams:
t = nearest_t(abar_of_lam(l))
if ts and t >= ts[-1]:
t = ts[-1] - 1 # 强制严格递减,避免差商分母为 0
if t < 1:
break # 干净端训练网格已无分辨率,到此为止
ts.append(t)
ts = np.array(ts, dtype=int)
return ts, np.concatenate([_ABAR[ts], [1.0]])
class Counter:
"""包一层模型,数一下到底调了几次。NFE 不能靠嘴算。"""
def __init__(self, model):
self.m = model
self.n = 0
def __call__(self, x, abar):
self.n += 1
return self.m(x, abar)
# ══════════════════════════════════════════════════════════════════
# 1. 一阶指数积分器(DPM-Solver-1 / dpmsolver++ 的一阶更新)
# ══════════════════════════════════════════════════════════════════
def dpm_pp_first(x, m0, a_s, s_s, a_t, s_t):
"""x_t = (sigma_t/sigma_s) * x_s - alpha_t * (e^{-h} - 1) * x0_hat
h = lambda_t - lambda_s。sigma_t = 0 时 lambda_t = +inf,e^{-h} = 0,
整条式子退化成 x_t = x0_hat —— 最后一步自动正确,不用特判。
"""
if s_t <= 0.0:
return m0
h = (np.log(a_t) - np.log(s_t)) - (np.log(a_s) - np.log(s_s))
return (s_t / s_s) * x - a_t * (np.exp(-h) - 1.0) * m0
def dpm_pp_second(x, hist, a_s, s_s, a_t, s_t, solver_type="midpoint"):
"""diffusers multistep_dpm_solver_second_order_update 的逐项复刻。
D1 是用上一步的 x0_hat 做差商估出来的导数:D1 ~ h * dx0/dlambda。
midpoint 和 heun 的区别只在 D1 前面的系数,两者都是二阶。
"""
lam_t = np.log(a_t) - np.log(s_t)
lam_s0, m0 = hist[-1]
lam_s1, m1 = hist[-2]
h = lam_t - lam_s0
h0 = lam_s0 - lam_s1
r0 = h0 / h
D0, D1 = m0, (m0 - m1) / r0
emh = np.exp(-h)
if solver_type == "midpoint":
return ((s_t / s_s) * x - a_t * (emh - 1.0) * D0
- 0.5 * a_t * (emh - 1.0) * D1)
return ((s_t / s_s) * x - a_t * (emh - 1.0) * D0
+ a_t * ((emh - 1.0) / h + 1.0) * D1)
def dpm_pp_third(x, hist, a_s, s_s, a_t, s_t, d2_scale=2.0):
"""三阶多步更新。
d2_scale=1.0 是 diffusers(v0.30.0)`multistep_dpm_solver_third_order_update`
的逐项复刻;d2_scale=2.0 是本文按「二次 x0_hat 必须精确」推出来的系数。
两者的差别见 F 节的受控检验 —— 差的就是这一个 2。
"""
lam_t = np.log(a_t) - np.log(s_t)
lam_s0, m0 = hist[-1]
lam_s1, m1 = hist[-2]
lam_s2, m2 = hist[-3]
h = lam_t - lam_s0
h0, h1 = lam_s0 - lam_s1, lam_s1 - lam_s2
r0, r1 = h0 / h, h1 / h
D0 = m0
D1_0, D1_1 = (m0 - m1) / r0, (m1 - m2) / r1
D1 = D1_0 + (r0 / (r0 + r1)) * (D1_0 - D1_1)
D2 = d2_scale * (D1_0 - D1_1) / (r0 + r1)
emh = np.exp(-h)
return ((s_t / s_s) * x - a_t * (emh - 1.0) * D0
+ a_t * ((emh - 1.0) / h + 1.0) * D1
- a_t * ((emh - 1.0 + h) / h ** 2 - 0.5) * D2)
# ══════════════════════════════════════════════════════════════════
# 2. 五个采样器
# ══════════════════════════════════════════════════════════════════
def solve_ddim(grid, model, x):
"""DDIM(eta=0):论文式 12,完全不碰 lambda。"""
for i in range(len(grid) - 1):
eps = model(x, grid[i])
x, _ = ddim_step(x, grid[i], grid[i + 1], eps, 0.0, 0.0)
return x
def solve_dpm1(grid, model, x):
"""一阶指数积分器。理论上应该和 DDIM 逐位相同 —— A 节去验。"""
for i in range(len(grid) - 1):
a_s, s_s = alpha_sigma(grid[i])
a_t, s_t = alpha_sigma(grid[i + 1])
eps = model(x, grid[i])
m0 = (x - s_s * eps) / a_s
x = dpm_pp_first(x, m0, a_s, s_s, a_t, s_t)
return x
def solve_dpm2s(grid, model, x):
"""单步二阶(DPM-Solver-2):每步先在 lambda 中点落一脚,用中点的 x0 走完。
一步两步评估。中点的 x0 对积分的「加权平均」是二阶准确的,
所以整步的局部误差是 O(h^3),全局 O(h^2)。
"""
for i in range(len(grid) - 1):
a_s, s_s = alpha_sigma(grid[i])
a_t, s_t = alpha_sigma(grid[i + 1])
eps = model(x, grid[i])
m0 = (x - s_s * eps) / a_s
lam_s, lam_t = np.log(a_s) - np.log(s_s), None
if s_t <= 0.0: # 收尾:lambda_t = +inf,中点无从谈起
x = m0
continue
lam_t = np.log(a_t) - np.log(s_t)
lam_m = 0.5 * (lam_s + lam_t)
abar_m = _ABAR[nearest_t(abar_of_lam(lam_m))]
a_m, s_m = alpha_sigma(abar_m)
x_m = dpm_pp_first(x, m0, a_s, s_s, a_m, s_m) # 先跳到中点
eps_m = model(x_m, abar_m)
m_m = (x_m - s_m * eps_m) / a_m # 中点处的 x0
x = dpm_pp_first(x, m_m, a_s, s_s, a_t, s_t) # 用中点 x0 走完整步
return x
def _midpoint_x0(x, m0, a_s, s_s, a_t, s_t, model):
"""在 lambda 的中点补一次评估,拿中点的 x0 走完整步(局部二阶)。
多步法开头那一步没有历史可用,只能降成一阶,而一阶在长度为 h 的区间上
局部误差是 O(h^2),这个误差会一路传到终点 —— 这就是 3M 实测阶数被卡在 2
的原因。给第一步多花一次评估可以验证这件事。
"""
lam_s, lam_t = np.log(a_s) - np.log(s_s), np.log(a_t) - np.log(s_t)
abar_m = abar_of_lam(0.5 * (lam_s + lam_t))
a_m, s_m = alpha_sigma(abar_m)
x_m = dpm_pp_first(x, m0, a_s, s_s, a_m, s_m)
eps_m = model(x_m, abar_m)
return (x_m - s_m * eps_m) / a_m
def solve_multistep(grid, model, x, order=2, solver_type="midpoint", warm2s=False,
d2_scale=2.0):
"""多步法(DPM-Solver++ 2M / 3M):一步一次评估,导数是拿历史 x0 做差商。
开头几步历史不够,自动降阶(和 diffusers 的 lower_order_nums 一样);
最后一步 sigma=0 强制降成一阶(和 final_sigmas_type="zero" 一样)。
warm2s=True 时,第一步改用中点法(多花一次评估)把局部精度提到二阶。
"""
hist = []
for i in range(len(grid) - 1):
a_s, s_s = alpha_sigma(grid[i])
a_t, s_t = alpha_sigma(grid[i + 1])
eps = model(x, grid[i])
m0 = (x - s_s * eps) / a_s
hist.append((np.log(a_s) - np.log(s_s), m0))
last = (i == len(grid) - 2)
if i == 0 and warm2s and s_t > 0.0:
x = dpm_pp_first(x, _midpoint_x0(x, m0, a_s, s_s, a_t, s_t, model),
a_s, s_s, a_t, s_t)
elif s_t <= 0.0 or len(hist) < 2 or last:
x = dpm_pp_first(x, m0, a_s, s_s, a_t, s_t)
elif order >= 3 and len(hist) >= 3:
x = dpm_pp_third(x, hist, a_s, s_s, a_t, s_t, d2_scale=d2_scale)
else:
x = dpm_pp_second(x, hist, a_s, s_s, a_t, s_t, solver_type)
return x
# ══════════════════════════════════════════════════════════════════
# 3. 主流程
# ══════════════════════════════════════════════════════════════════
def rms_err(x, xref):
return float(np.sqrt(((x - xref) ** 2).sum(1).mean()))
def fit_slope(ns, errs, res, nfe_min=24):
"""在双对数上拟合 err ~ N^(-p),返回 p。
只取渐近段:NFE 太小的时候高阶方法还在「攒历史」(前几步被迫降阶),
这一段量出来的斜率既不是 1 也不是 3,什么都不说明;误差小于参考解
分辨率的点是机器精度,同样不能要。
"""
xs, ys = [], []
for n, e in zip(ns, errs):
if n >= nfe_min and e > 3.0 * res and np.isfinite(e) and e > 0:
xs.append(np.log(n))
ys.append(np.log(e))
if len(xs) < 3:
return float("nan"), 0
p = -np.polyfit(xs, ys, 1)[0]
return float(p), len(xs)
def local_err_test():
"""受控检验:规定 x0_hat(lambda) 是二次多项式,看各阶更新的局部误差阶。
一条真正的 k 阶方法,在 x0_hat 是 k-1 次多项式时必须一步算准(误差 ~ 0),
因为它的构造就是「用 k 个历史点插值出 k-1 次多项式,再对 e^lambda 精确积分」。
这个检验把「公式对不对」和「轨迹好不好」彻底分开,谁也赖不着谁。
"""
rng = np.random.default_rng(3)
A, Bc, Cc = rng.standard_normal(2), rng.standard_normal(2), rng.standard_normal(2)
g = lambda lam: A + Bc * lam + 0.5 * Cc * lam ** 2
def exact_inc(lam_s, lam_t, n=200001):
u = np.linspace(lam_s, lam_t, n)
f = np.exp(u)[:, None] * np.stack([g(v) for v in u])
return np.trapezoid(f, u, axis=0)
hs = [0.4, 0.2, 0.1, 0.05, 0.025]
lam_s = 0.3
out = {}
for name in ("一阶", "二阶", "三阶-D2原式", "三阶-D2乘2"):
errs = []
for h in hs:
lam_t = lam_s + h
a_s, s_s = alpha_sigma(abar_of_lam(lam_s))
a_t, s_t = alpha_sigma(abar_of_lam(lam_t))
x = np.array([0.7, -0.4])
hist = [(lam_s - j * h, g(lam_s - j * h)) for j in (2, 1, 0)] # 老→新
if name == "一阶":
got = dpm_pp_first(x, g(lam_s), a_s, s_s, a_t, s_t)
elif name == "二阶":
got = dpm_pp_second(x, hist, a_s, s_s, a_t, s_t)
elif name == "三阶-D2原式":
got = dpm_pp_third(x, hist, a_s, s_s, a_t, s_t, d2_scale=1.0)
else:
got = dpm_pp_third(x, hist, a_s, s_s, a_t, s_t, d2_scale=2.0)
want = (s_t / s_s) * x + s_t * exact_inc(lam_s, lam_t)
errs.append(float(np.abs(got - want).max()))
out[name] = errs
print(f" {name:>12}: " + " ".join(f"{e:.2e}" for e in errs))
print(" (h 从 0.4 缩到 0.025)")
print(" → 一阶 h^2、二阶 h^3 都对;三阶原式只有 h^3(等于二阶),")
print(" D2 乘 2 之后直接掉到 1e-14 —— 二次 x0_hat 下它一步就算准了")
def main():
rng = np.random.default_rng(20260927)
w, mu, Sig = make_target()
TH = make_dirs()
Xref = sample_true(N_EVAL, w, mu, Sig, rng)
model = make_model("oracle", w, mu, Sig)
# ── A. DDIM == 一阶 DPM-Solver ? ──
print("=" * 72)
print("A. 交叉验证:DDIM(eta=0) 与 lambda 空间一阶指数积分器是不是同一个东西")
print("=" * 72)
n = 500
for S in (20, 50):
_, grid_l = uniform_lambda_grid(S)
_, grid_t = uniform_t_grid(S)
for name, grid in (("uniform-t", grid_t), ("uniform-lambda", grid_l)):
x0 = rng.standard_normal((n, 2))
xa = solve_ddim(grid, model, x0.copy())
xb = solve_dpm1(grid, model, x0.copy())
print(f" S={S:3d} {name:>14} max|DDIM - DPM-Solver-1| = "
f"{np.abs(xa - xb).max():.3e}")
print(" → 两套代码、两套记号,输出差在浮点误差量级:DDIM 就是一阶方法")
# ── B. 参考解 ──
print()
print("=" * 72)
print("B. 参考解:3M 跑 4000 步当真值,再自查一下它自己收敛到哪")
print("=" * 72)
n_ref = 1000
x_init = rng.standard_normal((n_ref, 2))
_, g_fine = uniform_lambda_grid(4000, quantize=False)
_, g_mid = uniform_lambda_grid(2000, quantize=False)
xref = solve_multistep(g_fine, model, x_init.copy(), order=3)
xchk = solve_multistep(g_mid, model, x_init.copy(), order=3)
res = rms_err(xref, xchk)
print(f" 4000 步 vs 2000 步参考解的差 = {res:.3e}")
print(f" → 这个数就是下面所有误差表的「分辨率」,比它小的差别不可信")
# ── C. 阶数 ──
print()
print("=" * 72)
print("C. 把阶数量出来:终点到参考解的 RMS 距离 vs 真实 NFE")
print("=" * 72)
NS = [6, 8, 12, 16, 24, 32, 48, 64, 96, 128, 192]
solvers = {
"DDIM (uniform-t)": lambda S: (uniform_t_grid(S)[1], solve_ddim, {}),
"DDIM (=DPM-1, lambda)": lambda S: (uniform_lambda_grid(S)[1], solve_ddim, {}),
"DPM-Solver-2 单步": lambda S: (uniform_lambda_grid(max(S // 2, 2))[1],
solve_dpm2s, {}),
"DPM-Solver++ 2M": lambda S: (uniform_lambda_grid(S)[1], solve_multistep,
{"order": 2}),
"3M(D2 用原式)": lambda S: (uniform_lambda_grid(S)[1], solve_multistep,
{"order": 3, "d2_scale": 1.0}),
"3M(D2 修正×2)": lambda S: (uniform_lambda_grid(S)[1], solve_multistep,
{"order": 3, "d2_scale": 2.0}),
}
print(f" {'NFE':>5} " + " ".join(f"{k:>14}" for k in solvers))
table = {k: [] for k in solvers}
nfes = []
for S in NS:
row, real_n = [], []
for k, build in solvers.items():
grid, fn, kw = build(S)
c = Counter(model)
x = fn(grid, c, x_init.copy(), **kw)
row.append(rms_err(x, xref))
real_n.append(c.n)
nfes.append(int(np.mean(real_n)))
table_k = list(solvers)
for k, v in zip(table_k, row):
table[k].append(v)
print(f" {nfes[-1]:5d} " + " ".join(f"{v:14.3e}" for v in row))
print()
print(f" 拟合误差 ~ NFE^(-p)(只取 NFE>=24 且误差高于分辨率 {res:.1e}×3 的渐近段):")
for k in solvers:
p, m = fit_slope(nfes, table[k], res)
print(f" {k:>24} p = {p:.2f} (用了 {m} 个点)")
# ── D. SW1 走到哪一步就分不出高下了 ──
print()
print("=" * 72)
print("D. 分布层面的尺子(SW1):重复换起点,看它什么时候失灵")
print("=" * 72)
floors = [sliced_w1(sample_true(N_EVAL, w, mu, Sig,
np.random.default_rng(2000 + r)), Xref, TH=TH)
for r in range(8)]
print(f" 真样本 vs 真样本(5 次):{np.mean(floors):.4f} ± {np.std(floors):.4f}"
f" ← 这就是尺子的分辨率")
PS = [10, 20, 50]
print(f" {'NFE':>5} " + " ".join(f"{k:>14}" for k in solvers))
for S in PS:
row = []
for k, build in solvers.items():
grid, fn, kw = build(S)
vs = []
for r in range(5): # 换几份起点,量出估计量的抖动
c = Counter(model)
x = fn(grid, c, np.random.default_rng(900 + S * 100 + r)
.standard_normal((N_EVAL, 2)), **kw)
vs.append(sliced_w1(x, Xref, TH=TH))
row.append((np.mean(vs), np.std(vs)))
print(f" {S:5d} " + " ".join(f"{v[0]:9.4f}±{v[1]:.3f}" for v in row))
print(" → 20 步时各采样器的差别还大于抖动;到 50 步大家都贴着地板,")
print(" SW1 已经回答不了「谁更好」——这种时候只能看 C 节的确定性误差")
# ── E. 同一个精度,省多少次评估 ──
print()
print("=" * 72)
print("E. 把终点误差压到 1e-2 / 1e-3,最少要几次模型评估(实测扫描)")
print("=" * 72)
scan = [4, 6, 8, 10, 12, 16, 20, 24, 32, 40, 48, 64, 80, 96, 128, 160, 192, 256, 384]
for k, build in solvers.items():
got = {}
for S in scan:
if len(got) == 2:
break
grid, fn, kw = build(S)
c = Counter(model)
e = rms_err(fn(grid, c, x_init.copy(), **kw), xref)
for tol in (1e-2, 1e-3):
if tol not in got and e <= tol:
got[tol] = c.n
s1 = got.get(1e-2, None)
s2 = got.get(1e-3, None)
f1 = f"NFE={s1:3d}" if s1 else " >400 "
f2 = f"NFE={s2:3d}" if s2 else " >400 "
print(f" {k:>24} 1e-2: {f1} 1e-3: {f2}")
# ── F. 受控检验:三阶更新到底几阶 ──
print()
print("=" * 72)
print("F. 受控检验:规定一条解析的 x0_hat(lambda),直接量局部误差")
print("=" * 72)
print(" 做法:不让 x0_hat 由 x 决定,而是规定它是 lambda 的二次多项式")
print(" g(lambda)。此时精确解就是 (sigma_t/sigma_s) x + sigma_t * 积分 e^lambda g。")
print(" x0_hat 是二次的,所以**任何真正的三阶方法都必须一步算准**。")
local_err_test()
# ── G. 卡住 3M 的不是开头那一步 ──
print()
print("=" * 72)
print("G. 是不是「第一步被迫降阶」拖累的?给第一步多花一次评估试试")
print("=" * 72)
for tag, kw in (("第一步一阶(默认)", {}),
("第一步改中点法(多 1 次评估)", {"warm2s": True})):
ns, es = [], []
for S in (24, 32, 48, 64, 96, 128, 192):
grid = uniform_lambda_grid(S)[1]
c = Counter(model)
ns.append(c.n if False else S)
es.append(rms_err(solve_multistep(grid, c, x_init.copy(),
order=3, d2_scale=1.0, **kw), xref))
p, _ = fit_slope(ns, es, res, nfe_min=24)
print(f" 3M(D2 原式)+ {tag:<22} 实测阶数 p = {p:.2f}")
print(" → 没变化。第一步不是瓶颈,瓶颈在系数本身(见 F 节)")
if __name__ == "__main__":
main()
# -*- coding: utf-8 -*-
"""步数该往哪放:同样 20 次评估,换个摆法能差出一个数量级。
阶数(dpm_solver_lab.py 里量出来的 p)说的是「步数变多时误差掉多快」,
但没说「步数摆在哪」。这一篇的主角是**时间步的摆法**:
- 在 t 上等间隔(DDIM 论文 / diffusers 默认 leading)
- 在 lambda 上等间隔(DPM-Solver 的建议)
- 在 lambda 上等间隔但落回整数 t(真实网络只能吃整数)
- 在 log sigma 上等间隔
- Karras / EDM 的 rho=7 摆法(SDXL、EDM 系列在用)
判断标准是两条:确定性误差(终点到参考解的距离)和「难度曲线」——
沿着一条高精度参考轨迹量 |dx0_hat / d lambda|,看哪一段 lambda 上 x0_hat
变得最快。步数就该往那儿放。
/usr/local/bin/python3 spacing_lab.py
依赖:numpy(+ 同目录的 oracle_gmm.py / ddim_family.py / dpm_solver_lab.py)
"""
import numpy as np
from oracle_gmm import (T, linear_schedule, make_target, sample_true,
alpha_sigma, sliced_w1, make_dirs, N_EVAL)
from ddim_family import uniform_t_grid, make_model, ddim_step
from dpm_solver_lab import (uniform_lambda_grid, solve_ddim, solve_multistep,
solve_dpm2s, lam_of_abar, abar_of_lam, nearest_t,
Counter, rms_err)
_BETAS, _ABAR = linear_schedule()
_LAM = lam_of_abar(_ABAR)
# ══════════════════════════════════════════════════════════════════
# 1. 五种摆法
# ══════════════════════════════════════════════════════════════════
def grid_uniform_t(S: int):
"""在 t 上等间隔(diffusers 的 leading / DDIM 论文默认)。"""
return uniform_t_grid(S)[1]
def grid_uniform_lambda(S: int):
"""在 lambda 上等间隔,连续 abar(DPM-Solver 理论里的标准摆法)。"""
return uniform_lambda_grid(S, quantize=False)[1]
def grid_uniform_lambda_q(S: int):
"""在 lambda 上等间隔,但落回整数时间步(真实网络只认整数 t)。"""
return uniform_lambda_grid(S, quantize=True)[1]
def _ratio(abar):
"""EDM / diffusers 语境里的 sigma:sigma_t / alpha_t = e^{-lambda}。
注意这不是 sqrt(1-abar)。VP 调度下 lambda 就是 -log 这个量,
所以「在 log sigma 上等间隔」等于「在 lambda 上等间隔」——两者是同一个东西,
表里因此只留一种。
"""
return np.sqrt(np.maximum(1.0 - abar, 1e-300) / np.maximum(abar, 1e-300))
def grid_uniform_lambda_capped(S: int, lam_max: float):
"""在 lambda 上等间隔,但只走到 lam_max 就收尾。
为什么要这个变体:uniform-t 天然把 lambda 的上界截在某个值(NFE=20 时
是 +1.76),剩下的干净端交给最后那个「直接输出 x0_hat」的一阶收尾步。
如果 uniform-t 赢的原因是「截得早」而不是「t 本身特别」,那么把
uniform-lambda 截到同一个上界,两者就该差不多。
"""
l0 = _LAM[T]
lams = np.linspace(l0, lam_max, S)
return np.concatenate([abar_of_lam(lams), [1.0]])
def grid_karras(S: int, rho: float = 7.0):
"""Karras / EDM 的摆法:在 sigma^(1/rho) 上等间隔,EDM 原文取 rho=7。
sigma 用的是 EDM 的那个 sigma(= sigma_t/alpha_t),VP 下等于 e^{-lambda},
区间是 [_ratio(abar_T), _ratio(abar_1)] ≈ [158, 0.01]。
"""
r_hi, r_lo = _ratio(_ABAR[T]), _ratio(_ABAR[1])
r = (r_hi ** (1.0 / rho)
+ np.arange(S) / (S - 1) * (r_lo ** (1.0 / rho) - r_hi ** (1.0 / rho))) ** rho
return np.concatenate([1.0 / (1.0 + r ** 2), [1.0]])
GRIDS = {
"uniform-t": grid_uniform_t,
"uniform-lambda": grid_uniform_lambda,
"uniform-lambda(整数t)": grid_uniform_lambda_q,
"lambda截到1.8": lambda S: grid_uniform_lambda_capped(S, 1.8),
"lambda截到3.0": lambda S: grid_uniform_lambda_capped(S, 3.0),
"karras(rho=7)": grid_karras,
"karras(rho=3)": lambda S: grid_karras(S, rho=3.0),
}
# ══════════════════════════════════════════════════════════════════
# 2. 难度曲线:沿着参考轨迹看 x0_hat 在哪一段变得最快
# ══════════════════════════════════════════════════════════════════
def difficulty_profile(model, n=400, steps=1500, seed=5):
"""返回 (lam_mid, weight):一阶方法每一步的误差贡献权重。
推导:第 i 步的局部误差(在 x 的单位里)是
delta_i = sigma_{t_i} * ∫ e^lambda [x0(lam) - x0(lam_i)] d lambda
≈ sigma_{t_i} * e^{lam_i} * (h^2/2) * x0'(lam_i)
= alpha_{t_i} * e^{-h} * (h^2/2) * x0'(lam_i)
也就是说局部误差自带一个 alpha_t 因子:噪声端 alpha≈0,轨迹在 lambda 上
几乎不动,走错一点也不打紧;干净端 alpha≈1,同样的 h 会实打实地错。
所以难度权重是 alpha(lambda) * |dx0_hat / d lambda|,不是单纯的导数。
"""
x = np.random.default_rng(seed).standard_normal((n, 2))
grid = uniform_lambda_grid(steps, quantize=False)[1]
lams, x0s, alphas = [], [], []
for i in range(len(grid) - 1):
a_s, s_s = alpha_sigma(grid[i])
a_t, s_t = alpha_sigma(grid[i + 1])
eps = model(x, grid[i])
m0 = (x - s_s * eps) / a_s
lams.append(lam_of_abar(grid[i]))
alphas.append(a_s)
x0s.append(m0.copy())
x = (s_t / s_s) * x - a_t * (np.exp(-(lam_of_abar(grid[i + 1])
- lam_of_abar(grid[i]))) - 1.0) * m0
if s_t <= 0:
break
lams = np.array(lams)
alphas = np.array(alphas)
x0s = np.array(x0s) # (steps, n, 2)
dx = np.sqrt(((np.diff(x0s, axis=0)) ** 2).sum(-1).mean(-1))
speed = dx / np.abs(np.diff(lams))
lam_mid = 0.5 * (lams[1:] + lams[:-1])
a_mid = 0.5 * (alphas[1:] + alphas[:-1])
return lam_mid, a_mid * speed, speed
# ══════════════════════════════════════════════════════════════════
# 3. 主流程
# ══════════════════════════════════════════════════════════════════
def main():
rng = np.random.default_rng(20260927)
w, mu, Sig = make_target()
model = make_model("oracle", w, mu, Sig)
TH = make_dirs()
Xref = sample_true(N_EVAL, w, mu, Sig, rng)
n_ref = 1000
x_init = rng.standard_normal((n_ref, 2))
g_ref = uniform_lambda_grid(3000, quantize=False)[1]
xref = solve_multistep(g_ref, model, x_init.copy(), order=3)
g_chk = uniform_lambda_grid(1500, quantize=False)[1]
res = rms_err(xref, solve_multistep(g_chk, model, x_init.copy(), order=3))
print("=" * 72)
print("A. 参考解分辨率")
print("=" * 72)
print(f" 3000 步 vs 1500 步 = {res:.2e}(误差表只能信到这个量级)")
# ── B. 摆法对比(确定性误差)──
print()
print("=" * 72)
print("B. 同样 NFE,步数摆在哪:终点到参考解的 RMS 距离")
print("=" * 72)
print(f" {'摆法':>22} " + " ".join(f"{'N=' + str(s):>9}" for s in (10, 20, 50))
+ " | " + " ".join(f"{'N=' + str(s):>9}" for s in (10, 20, 50)))
print(f" {'':>22} " + " ".join(f"{'DDIM':>9}" for _ in range(3))
+ " | " + " ".join(f"{'2M':>9}" for _ in range(3)))
for name, build in GRIDS.items():
row_a, row_b = [], []
for S in (10, 20, 50):
grid = build(S)
row_a.append(rms_err(solve_ddim(grid, model, x_init.copy()), xref))
row_b.append(rms_err(solve_multistep(grid, model, x_init.copy(),
order=2), xref))
print(f" {name:>22} " + " ".join(f"{v:9.3e}" for v in row_a)
+ " | " + " ".join(f"{v:9.3e}" for v in row_b))
# ── C. 摆法对比(分布层面)──
print()
print("=" * 72)
print("C. 同样的比较,换成 SW1(地板约 0.030,重复换起点看抖动)")
print("=" * 72)
print(f" {'摆法':>22} " + " ".join(f"{'N=' + str(s):>16}" for s in (10, 20)))
for name, build in GRIDS.items():
row = []
for S in (10, 20):
grid = build(S)
vs = []
for r in range(5):
x0 = np.random.default_rng(700 + S * 10 + r).standard_normal(
(N_EVAL, 2))
vs.append(sliced_w1(solve_multistep(grid, model, x0, order=2),
Xref, TH=TH))
row.append((np.mean(vs), np.std(vs)))
print(f" {name:>22} " + " ".join(f"{v[0]:.4f}±{v[1]:.3f}" for v in row))
# ── D. 难度曲线 ──
print()
print("=" * 72)
print("D. 难度曲线:x0_hat 在哪一段 lambda 上变化最快")
print("=" * 72)
lams, weight, raw = difficulty_profile(model)
lo = lams < 0
print(f" 只看 |dx0/dlambda|(不带 alpha 因子):"
f"噪声端 {raw[lo].mean():.3e} vs 干净端 {raw[~lo].mean():.3e}")
print(f" 乘上 alpha_t 之后的**误差权重**: "
f"噪声端 {weight[lo].mean():.3e} vs 干净端 {weight[~lo].mean():.3e}"
f" ← 干净端重 {weight[~lo].mean() / weight[lo].mean():.1f} 倍")
top = np.argsort(weight)[-5:][::-1]
print(" 权重最大的 5 段:")
for i in top:
print(f" lambda = {lams[i]:+6.2f} 权重 = {weight[i]:.3e}")
print(" → alpha_t 这个因子把难度整体推向干净端:噪声端轨迹在 lambda 上几乎")
print(" 不动(alpha≈0),走错一点也不打紧;干净端才是一步错步步错的地方")
# ── E2. 收尾步不是免费的 ──
print()
print("=" * 72)
print("E2. 收尾步的代价:在 lambda_max 处直接输出 x0_hat,到底偏了多少")
print("=" * 72)
print(" 沿高精度参考轨迹取快照,对每个 lambda 量")
print(" || x0_hat(x_lambda, lambda) - 轨迹终点 || —— 这就是提前收尾的偏差。")
g_fine = uniform_lambda_grid(3000, quantize=False)[1]
x = x_init.copy()
snaps = []
for i in range(len(g_fine) - 1):
a_s, s_s = alpha_sigma(g_fine[i])
a_t, s_t = alpha_sigma(g_fine[i + 1])
eps = model(x, g_fine[i])
m0 = (x - s_s * eps) / a_s
if i % 60 == 0:
snaps.append((lam_of_abar(g_fine[i]),
float(np.sqrt(((m0 - xref[:n_ref]) ** 2).sum(1).mean()))))
x = (s_t / s_s) * x - a_t * (np.exp(-(lam_of_abar(g_fine[i + 1])
- lam_of_abar(g_fine[i]))) - 1.0) * m0
if s_t <= 0:
break
for lam, d in snaps[::5]:
print(f" lambda = {lam:+6.2f} 提前收尾的偏差 = {d:.3e}")
print(" → 偏差随 lambda 单调下降:截得越早,这一步白送的误差越大,")
print(" 而且它**不随步数下降**(只取决于截在哪个 lambda)")
print()
print(" 于是 lambda_max 有个最优值(2M, N=20,只改截断点):")
for cap in (1.0, 1.5, 2.0, 2.5, 3.0, 3.5, 4.0, 4.61):
grid = grid_uniform_lambda_capped(20, cap)
e = rms_err(solve_multistep(grid, model, x_init.copy(), order=2), xref)
print(f" lambda_max = {cap:4.2f} 终点误差 = {e:.3e}")
print(" → 两头都变差:截太早是收尾偏差,截太晚是每步的 h 变大")
# ── E. 20 步时各摆法的落脚点 ──
print()
print("=" * 72)
print("E. NFE=20 时各摆法把步子落在 lambda 的哪几个位置")
print("=" * 72)
for name, build in GRIDS.items():
grid = build(20)
lam = np.array([lam_of_abar(g) if g < 1.0 else np.inf for g in grid])
lam_f = lam[np.isfinite(lam)]
print(f" {name:>22}: lambda 从 {lam_f[0]:+.2f} 到 {lam_f[-1]:+.2f},"
f"相邻间隔 min {np.diff(lam_f).min():.3f} / max {np.diff(lam_f).max():.3f}")
if __name__ == "__main__":
main()
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""给「从 DDIM 到高阶采样器」画五张图。
所有数字都是现场重算的:这里 import ddim_family / dpm_solver_lab 里的
采样器,重新跑一遍实验再画,不抄任何手打的表格。改了那边的实现,重跑
这个脚本图就会跟着变,不会出现「图上是旧数字、正文是新数字」。
五张图分别回答:
1. eta=0 和 eta=1 走的是两条什么样的路(同一份随机数,二维真轨迹)
2. eta 该调多大?oracle 模型下扫一遍,欠拟合模型下再扫一遍
3. 为什么换采样器缓存就得重调:模型输入每步变多少
4. 阶数:误差 vs NFE 的双对数斜率
5. 达到 1e-2 / 1e-3 到底要花几次模型评估
只依赖 numpy + matplotlib。跑法:/usr/local/bin/python3 make_figures.py
"""
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import oracle_gmm as OG
import ddim_family as DF
import dpm_solver_lab as DL
ROOT = Path(__file__).resolve().parents[1]
OUT = ROOT / "figures"
OUT.mkdir(parents=True, exist_ok=True)
plt.rcParams["font.sans-serif"] = ["PingFang SC", "Heiti TC", "Arial Unicode MS"]
plt.rcParams["axes.unicode_minus"] = False
plt.rcParams["figure.dpi"] = 130
RNG_SEED = 20260927
# ══════════════════════════════════════════════════════════════════
# 公共:目标分布、参考解、评估尺子
# ═════════════════════════════════════════════════════════════════=
def build_world():
rng = np.random.default_rng(RNG_SEED)
w, mu, Sig = OG.make_target()
Xref = OG.sample_true(OG.N_EVAL, w, mu, Sig, rng)
TH = OG.make_dirs()
return w, mu, Sig, Xref, TH
def floors(TH, Xref, w, mu, Sig, n=8, base=2000):
"""尺子的分辨率:真样本对真样本,量 n 次。"""
v = [OG.sliced_w1(OG.sample_true(OG.N_EVAL, w, mu, Sig,
np.random.default_rng(base + r)),
Xref, TH=TH) for r in range(n)]
return float(np.mean(v)), float(np.std(v))
# ══════════════════════════════════════════════════════════════════
# 图 1:两条轨迹
# ═════════════════════════════════════════════════════════════════=
def fig_traj(w, mu, Sig, Xref):
model = DF.make_model("oracle", w, mu, Sig)
S = 20
ts, grid = DF.uniform_t_grid(S)
n_show = 4
x_init, zs = DF.paired_noise(S, n_show, seed=4321)
fig, ax = plt.subplots(figsize=(6.4, 5.4))
ax.scatter(Xref[:, 0], Xref[:, 1], s=3, c="#C9CDD4", alpha=0.45,
label="真实样本(目标分布)")
styles = {1.0: ("#D2691E", 0.9, 5, r"$\eta=1$(DDPM 祖采样式):每步都重掷噪声"),
0.0: ("#2F6FB3", 1.7, 8, r"$\eta=0$(DDIM,确定性):一条光滑的路")}
for eta, (c, lw, ms, lab) in styles.items():
_, (xs, _) = DF.run_ddim(x_init.copy(), grid, model, eta, zs, track=True)
P = np.stack(xs, axis=1) # (n, S+1, 2)
for k in range(n_show):
ax.plot(P[k, :, 0], P[k, :, 1], "-", color=c, lw=lw,
alpha=0.85 if eta == 0.0 else 0.55,
zorder=4 if eta == 0.0 else 3,
label=lab if k == 0 else None)
ax.scatter(P[k, 1:-1, 0], P[k, 1:-1, 1], s=ms, color=c,
alpha=0.85 if eta == 0.0 else 0.5, zorder=4)
ax.scatter(x_init[:, 0], x_init[:, 1], s=110, marker="*", c="#111111",
zorder=6, label="起点(两边共用,只有 4 个)")
ax.set_xlabel(r"$x_1$")
ax.set_ylabel(r"$x_2$")
ax.set_title("同一份起点、同一份随机数:N=20 步走出的两条路")
ax.legend(loc="upper right", fontsize=8, framealpha=0.9)
ax.set_aspect("equal", adjustable="box")
fig.tight_layout()
p = OUT / "traj_eta.png"
fig.savefig(p)
plt.close(fig)
print(f" [1] {p.name} 轨迹 {n_show} 条 × 2 种 eta")
# ══════════════════════════════════════════════════════════════════
# 图 2:eta 扫描(oracle + 欠拟合)
# ═════════════════════════════════════════════════════════════════=
def fig_eta_sweep(w, mu, Sig, Xref, TH):
NS = (10, 20, 50, 100)
etas = (0.0, 0.25, 0.5, 0.75, 1.0)
flo, flo_sd = floors(TH, Xref, w, mu, Sig)
panels = {}
for tag, kind in (("oracle(真 score,模型误差=0)", "oracle"),
("欠拟合(模型以为数据是单高斯)", "gauss")):
model = DF.make_model(kind, w, mu, Sig)
panels[tag] = {e: [] for e in etas}
for e in etas:
for S in NS:
_, grid = DF.uniform_t_grid(S)
xi, zs = DF.paired_noise(S, OG.N_EVAL)
x, _ = DF.run_ddim(xi, grid, model, e, zs)
panels[tag][e].append(OG.sliced_w1(x, Xref, TH=TH))
fig, axes = plt.subplots(1, 2, figsize=(11.2, 4.4), sharey=True)
cmap = plt.get_cmap("viridis")
for ax, (tag, tab) in zip(axes, panels.items()):
for i, e in enumerate(etas):
ax.plot(NS, tab[e], "o-", color=cmap(0.15 + 0.7 * i / (len(etas) - 1)),
lw=1.6, ms=5, label=r"$\eta=" + f"{e:g}" + r"$")
ax.axhline(flo, ls="--", c="#999999", lw=1.2,
label=f"噪声地板 {flo:.4f}")
ax.set_xscale("log")
ax.set_xticks(list(NS))
ax.set_xticklabels([str(s) for s in NS])
ax.set_xlabel("采样步数 N")
ax.set_title(tag)
ax.grid(alpha=0.25)
axes[0].set_ylabel("SW1(越小越好)")
axes[0].legend(fontsize=8)
axes[1].legend(fontsize=8)
fig.suptitle("eta 越小越好,而且两个模型下结论一致", fontsize=12)
fig.tight_layout()
p = OUT / "eta_sweep.png"
fig.savefig(p)
plt.close(fig)
print(f" [2] {p.name} 地板 {flo:.4f}±{flo_sd:.4f};"
f"oracle N=10: eta0={panels['oracle(真 score,模型误差=0)'][0.0][0]:.4f}"
f" eta1={panels['oracle(真 score,模型误差=0)'][1.0][0]:.4f}")
# ══════════════════════════════════════════════════════════════════
# 图 3:每步「模型输入变了多少」
# ═════════════════════════════════════════════════════════════════=
def fig_cache_jitter(w, mu, Sig):
model = DF.make_model("oracle", w, mu, Sig)
S = 50
_, grid = DF.uniform_t_grid(S)
x_init, zs = DF.paired_noise(S, OG.N_EVAL)
fig, ax = plt.subplots(figsize=(6.6, 4.4))
cols = {0.0: "#2F6FB3", 0.5: "#7A7A7A", 1.0: "#D2691E"}
last10 = {}
for eta, c in cols.items():
_, (xs, _) = DF.run_ddim(x_init.copy(), grid, model, eta, zs, track=True)
rel = [float(np.sqrt((((xs[i + 1] - xs[i]) ** 2).sum(1)).mean())
/ np.sqrt(((xs[i] ** 2).sum(1)).mean()))
for i in range(len(xs) - 1)]
last10[eta] = float(np.mean(rel[-10:]))
ax.plot(range(1, len(rel) + 1), rel, "-", color=c, lw=1.6,
label=r"$\eta=" + f"{eta:g}" + r"$")
ax.plot([len(rel) - 9, len(rel)], [last10[eta], last10[eta]],
lw=3.2, color=c, alpha=0.35)
ax.set_xlabel("步序号(越往右越接近干净端)")
ax.set_ylabel(r"模型输入相对变化 $\Delta x_i / x_i$(长度之比,无量纲)")
ax.set_title("缓存加速盯的就是这条线:eta 越大,每步输入跳得越狠")
ax.legend(fontsize=9)
ax.grid(alpha=0.25)
ratio = last10[1.0] / last10[0.0]
ax.annotate(f"最后 10 步:eta=1 是 eta=0 的 {ratio:.1f} 倍",
xy=(0.42, 0.86), xycoords="axes fraction", fontsize=9,
bbox=dict(boxstyle="round,pad=0.35", fc="#FFF6E5", ec="#D2691E"))
fig.tight_layout()
p = OUT / "cache_jitter.png"
fig.savefig(p)
plt.close(fig)
print(f" [3] {p.name} 最后10步 eta0={last10[0.0]:.4f} "
f"eta0.5={last10[0.5]:.4f} eta1={last10[1.0]:.4f} 倍数 {ratio:.1f}")
# ══════════════════════════════════════════════════════════════════
# 图 4 + 图 5:阶数与评估预算
# ═════════════════════════════════════════════════════════════════=
def build_reference(w, mu, Sig):
"""超高分辨率参考解;顺手用 2000 步自查它的收敛精度。
随机数的取法与 dpm_solver_lab.py 的 B 节一模一样(同一个种子、同样的
消耗顺序:先抽一批真样本、再抽初始噪声),所以这里打印的分辨率和那边
是同一个数,图上和正文不会对不上。
"""
model = DF.make_model("oracle", w, mu, Sig)
rng = np.random.default_rng(RNG_SEED)
OG.sample_true(OG.N_EVAL, w, mu, Sig, rng)
x_init = rng.standard_normal((1000, 2))
_, gf = DL.uniform_lambda_grid(4000, quantize=False)
_, gm = DL.uniform_lambda_grid(2000, quantize=False)
xref = DL.solve_multistep(gf, model, x_init.copy(), order=3, d2_scale=2.0)
xchk = DL.solve_multistep(gm, model, x_init.copy(), order=3, d2_scale=2.0)
return model, x_init, xref, DL.rms_err(xref, xchk)
def fig_order_and_budget(w, mu, Sig):
model, x_init, xref, res = build_reference(w, mu, Sig)
NS = [6, 8, 12, 16, 24, 32, 48, 64, 96, 128, 192]
solvers = {
"DDIM(uniform-t)": lambda S: (DF.uniform_t_grid(S)[1], DL.solve_ddim, {}),
"DDIM(=DPM-1, lambda)": lambda S: (DL.uniform_lambda_grid(S)[1],
DL.solve_ddim, {}),
"DPM-Solver-2 单步": lambda S: (DL.uniform_lambda_grid(max(S // 2, 2))[1],
DL.solve_dpm2s, {}),
"DPM-Solver++ 2M": lambda S: (DL.uniform_lambda_grid(S)[1],
DL.solve_multistep, {"order": 2}),
"3M(三阶项原式)": lambda S: (DL.uniform_lambda_grid(S)[1],
DL.solve_multistep,
{"order": 3, "d2_scale": 1.0}),
"3M(三阶项修正)": lambda S: (DL.uniform_lambda_grid(S)[1],
DL.solve_multistep,
{"order": 3, "d2_scale": 2.0}),
}
cols = ["#2F6FB3", "#6FA8DC", "#7A7A7A", "#3E8E41", "#D2691E", "#B03060"]
tab = {k: [] for k in solvers}
nfes = []
for S in NS:
cn = []
for k, build in solvers.items():
grid, fn, kw = build(S)
c = DL.Counter(model)
tab[k].append(DL.rms_err(fn(grid, c, x_init.copy(), **kw), xref))
cn.append(c.n)
nfes.append(int(round(np.mean(cn))))
# ── 图 4:阶数 ──
fig, ax = plt.subplots(figsize=(7.0, 5.0))
for (k, errs), c in zip(tab.items(), cols):
p, _ = DL.fit_slope(nfes, errs, res)
ax.plot(nfes, errs, "o-", color=c, lw=1.5, ms=4.5,
label=f"{k} " + r"$p=" + f"{p:.2f}" + r"$")
ax.axhline(res, ls="--", c="#BBBBBB", lw=1.2)
ax.text(nfes[-1], res * 1.15, f"参考解分辨率 {res:.1e}", fontsize=8,
color="#888888", ha="right")
ax.set_xscale("log")
ax.set_yscale("log")
ax.set_xlabel("NFE(真实模型评估次数,数出来的)")
ax.set_ylabel("终点到参考解的 RMS 距离")
ax.set_title("斜率就是阶数:低阶方法要加一个数量级的步数才能追上")
ax.grid(alpha=0.25, which="both")
ax.legend(fontsize=8, loc="lower left")
fig.tight_layout()
p4 = OUT / "order_nfe.png"
fig.savefig(p4)
plt.close(fig)
# ── 图 5:评估预算 ──
scan = [4, 6, 8, 10, 12, 16, 20, 24, 32, 40, 48, 64, 80, 96, 128, 160,
192, 256, 384]
got = {}
for k, build in solvers.items():
row = {}
for S in scan:
grid, fn, kw = build(S)
c = DL.Counter(model)
e = DL.rms_err(fn(grid, c, x_init.copy(), **kw), xref)
for tol in (1e-2, 1e-3):
if tol not in row and e <= tol:
row[tol] = c.n
got[k] = row
fig, ax = plt.subplots(figsize=(7.4, 4.6))
names = list(solvers)
xs = np.arange(len(names))
for j, (tol, lab) in enumerate(((1e-2, r"压到 $10^{-2}$ 所需 NFE"),
(1e-3, r"压到 $10^{-3}$ 所需 NFE"))):
vals, txt = [], []
for k in names:
v = got[k].get(tol)
vals.append(v if v else 0)
txt.append(str(v) if v else ">" + str(scan[-1]))
bars = ax.bar(xs + (j - 0.5) * 0.38, vals, width=0.36,
color=["#8FB8DE", "#1F4E79"][j], label=lab)
for b, t in zip(bars, txt):
ax.text(b.get_x() + b.get_width() / 2, b.get_height() + 4, t,
ha="center", fontsize=8)
ax.set_xticks(xs)
ax.set_xticklabels(names, rotation=18, ha="right", fontsize=8)
ax.set_ylabel("NFE")
ax.set_title("同一个精度,高阶方法省下的评估次数(柱子顶上标的是实测值)")
ax.legend(fontsize=9)
ax.grid(alpha=0.25, axis="y")
fig.tight_layout()
p5 = OUT / "nfe_budget.png"
fig.savefig(p5)
plt.close(fig)
print(f" [4] {p4.name} 参考解分辨率 {res:.3e}")
print(f" [5] {p5.name} 1e-2 所需 NFE: "
+ " ".join(f"{k}={got[k].get(1e-2, '>384')}" for k in names))
print(f" 1e-3 所需 NFE: "
+ " ".join(f"{k}={got[k].get(1e-3, '>384')}" for k in names))
def main():
print("生成配图(数值全部现场重算)")
w, mu, Sig, Xref, TH = build_world()
fig_traj(w, mu, Sig, Xref)
fig_eta_sweep(w, mu, Sig, Xref, TH)
fig_cache_jitter(w, mu, Sig)
fig_order_and_budget(w, mu, Sig)
print("完成,输出目录:", OUT)
if __name__ == "__main__":
main()
评论 (0)