基于多层欧拉-马鲁亚马方法的扩散模型多项式加速

Polynomial Speedup in Diffusion Models with the Multilevel Euler-Maruyama Method

arXiv: 2603.24594v1

论文信息

标题: Polynomial Speedup in Diffusion Models with the Multilevel Euler-Maruyama Method

作者: Arthur Jacot

发布日期: 2026-03-25

arXiv ID: 2603.24594v1

PDF 链接: 下载 PDF

3 分钟速览

  • 研究问题:这篇论文要解决扩散模型(如 DDPM、DDIM)在采样阶段计算成本过高的问题。传统欧拉-丸山方法需要大量网络评估次数才能达到高精度,论文探索如何在不损失图像质量的前提下,降低计算消耗。
  • 核心方法:提出多层级欧拉-丸山方法,利用多个不同大小和精度的神经网络近似漂移项。在大部分采样步中使用成本低的小网络,仅以低概率触发成本高的大网络,从而在期望上减少总计算量。
  • 关键结果:在 CelebA 64×64 数据集上,ML-EM 方法相对传统欧拉-丸山方法实现了最高约四倍加速(见论文图 1),且理论证明当漂移项处于 “比蒙特卡洛更难” 区间时,算法能以单次最佳网络评估的等效成本求解整个 SDE。
  • 主要局限:方法依赖漂移项满足 HTMC 假设(即网络误差随参数量的下降速率足够慢),且实际加速效果受超参数选择、GPU 批处理方式影响。ML-EM 带来的误差方差较大,需要多次试验选取最佳伯努利采样组合。
  • 适合读者:适合从事扩散模型采样加速、随机微分方程数值解、以及具备深度学习与概率论基础的量化研究人员阅读。

论文背景和研究动机

去噪扩散概率模型是当前图像、视频等多模态生成任务的主流技术,其核心是一个反向随机微分方程。扩散模型的采样过程依赖一个大型深度神经网络(通常是 UNet)在每一步中近似该 SDE 的漂移项或得分函数。这意味着生成的精度越高,所需的网络评估次数就越长,计算负担越重。

已有缓解高计算成本的工作包括 DDIM、概率流 ODE、高阶 Runge-Kutta 求解器、渐进式蒸馏等。这些方法或在求解器精度上做文章,或通过训练新网络直接跳跃多步。本文则从一个全新视角出发,把关注点从 “如何从微分方程离散误差中获得加速” 转换到 “如何从函数近似误差的缩放规律中获得加速”。作者的核心理由是:扩散模型使用的深度网络天然服从误差-参数量规模的幂律缩放,即更精确的网络需要指数级增长的计算成本,当该缩放指数足够大时,存在加速空间,且这种加速是多项式的,在大规模工业场景下更为显著。

核心方法和技术细节

论文聚焦于 SDE

dxt=ft(xt)dt+σtdWtdx_t = f_t(x_t) dt + \sigma_t dW_t

的数值求解,其中漂移项 ftf_t 必须被深度神经网络近似。作者假设存在一系列逐渐精确的近似器 ft1,ft2,…,ftkf_t^1, f_t^2, \dots, f_t^k,满足 ∣ft−ftk∣∞≤2−k|f_t - f_t^k|_\infty \leq 2^{-k},其计算成本为 C(ftk)≤cγ2γkC(f_t^k) \leq c^\gamma 2^{\gamma k}。γ\gamma 由网络测试误差与参数量之间的幂律关系导出,当 γ>2\gamma > 2 时,任务进入 “比蒙特卡洛更难” 区间。

多层级欧拉-丸山方法(ML-EM)的迭代形式如下:在每个离散时间步,漂移项以多重差分的方式更新:

yt+η=yt+η∑k=kmin⁡kmax⁡Bkpk[ftk(yt)−ftk−1(yt)]+η σtZty_{t+\eta} = y_t + \eta \sum_{k=k_{\min}}^{k_{\max}} \frac{B^k}{p_k} \big[f_t^k(y_t) - f_t^{k-1}(y_t)\big] + \sqrt{\eta}\,\sigma_t Z_t

其中 Bk∼Bernoulli(pk)B^k \sim \mathrm{Bernoulli}(p_k) 是独立的伯努利随机变量。其期望等价于仅使用最高精度网络 ftkmax⁡f_t^{k_{\max}} 的标准欧拉-丸山方法,但高成本网络的触发概率 pkp_k 随 kk 指数递减,从而显著降低期望计算量。

论文中最关键的选择是 pk=min⁡{C 2−(1+γ/2)k, 1}p_k = \min\{C\,2^{-(1+\gamma/2)k},\,1\}。当 γ>2\gamma > 2 时,期望计算复杂度可以控制在 O(ϵ−γ)\mathcal{O}(\epsilon^{-\gamma}),与单次最高精度网络评估成本相同(见论文定理 1 和附录 B 的详细推导)。为了进一步接近实际,作者还提出一种自适应方法,将 pkp_k 参数化为时间依赖的 sigmoid 函数 pk(t)=σ(αklog⁡(t+δ)+βk)p_k(t) = \sigma(\alpha_k \log(t+\delta) + \beta_k),通过前向梯度估计和随机高斯向量的近似,在不进行昂贵反向传播的情况下更新参数,从而找到计算成本与精度之间的最佳折衷。

创新点和贡献

本文的主要创新在于将多层级蒙特卡洛思想引入 SDE 求解器,并特别针对扩散模型的训练与采样带来双重收益。首先,它将关注点从离散化阶提升至函数近似误差率,并证明 HTMC 区间内 ML-EM 可以达到与单次大网络评估等价的效率,即打破了传统思路中 “NFE 必须随精度增长” 的固有认知。

其次,ML-EM 不依赖特定的采样器或蒸馏流程,可以与 DDIM、概率流 ODE 以及高阶求解器等方法叠加使用,具有高度通用性。同时,它利用了实践中常见的多尺度网络,例如,在训练大模型之前的超参搜索过程中自然得到的一系列小型网络即可直接充当多层级近似器,无需重新训练,极大降低了应用门槛。

第三,自适应学习 pk(t)p_k(t) 的方法将问题转化为标量参数优化,并使用前向传播和隐式重参数化梯度进行无偏估计。这一层设计避开了倒推整个 SDE 图所需的内存开销,使方案具备实际可操作性。

实验结果分析

实验在 CelebA 64×64 数据集上完成。作者训练了 5 个逐级增大的 UNet,除最大模型 f5f^5 外,中级网络 f3f^3 和小型网络 f1f^1 用于 ML-EM 采样(见论文第 4 节)。以 f5f^5 配合 1000 步的 DDPM 或 DDIM 采样作为 “真实样本”,衡量其他配置产生的均方误差与计算时间。

固定概率组合(pk=CTk−1p_k = C T_k^{-1} 或 pk=CTk−0.9p_k = C T_k^{-0.9})已展现出相对于传统欧拉-丸山方法的明显优势,而自适应学习概率的 ML-EM 在 DDPM 场景中给出了最佳折衷:最多将生成时间缩短至四分之一,或在相同时间内使 MSE 降低约 10 倍(见论文图 1(a))。有趣的是,ML-EM 对 DDIM 的效果虽不如 DDPM 突出,但仍表现出一定加速。从生成图像质量看,ML-EM 有效抑制了低步数下常见的偏色和对比度问题,再次表明其主要提升来自于对近似误差的统计管理,而非纯粹克服离散化噪声。

论文还估测该任务的 γ≈2.5\gamma \approx 2.5,确认 CelebA 场景已落入 HTMC 区间(图 2)。作者指出,在更大规模数据集和更大网络场景下,γ\gamma 往往更高,ML-EM 的相对加速比可能会进一步扩大。

实践建议

对于希望在扩散模型实际部署中应用 ML-EM 的工程团队,可参考以下要点:

  1. 多尺度网络复用:不必专门为 ML-EM 额外训练模型。在超参数搜索或逐步扩展实验阶段,保留不同尺寸、不同精度的 UNet,可以直接作为多层级近似器。
  2. 超参数选择:当无法精确测量 γ\gamma 时,可采用 pk∝Tk−1p_k \propto T_k^{-1} 的简单设定,这是 β=γ\beta=\gamma 的特例,仍能保证最优误差率(见论文第 3 节关于 β\beta 的讨论)。只需调整全局缩放因子 CC 控制误差-成本平衡。
  3. 自适应训练注意事项:若使用自适应方法学习 pk(t)p_k(t),应避免在梯度估计中将伯努利变量跨样本共享,否则会增加方差。另外,前向梯度估计中的随机高斯向量维度与参数总数匹配,复杂度可控。
  4. GPU 批处理策略:生产环境中可采用跨批次共享伯努利掩码的方式,即一次决策对整个批次应用或跳过网络评估,以充分利用 GPU 的并行计算能力。作者在论文第 4 节验证了该策略的实际加速效果。
  5. 与 DDIM 及高阶求解器结合:ML-EM 独立于采样方式和离散阶数,理论上可与 Runge-Kutta 类方法或概率流 ODE 一并使用。在噪声较小或确定性轨迹的生成任务中,ML-EM 仍有助益,但加速幅度可能低于纯 SDE 场景,实践中值得权衡。