Back to blog

Inductive Moment Matching: Few-Step Generation via Single-Stage Training

提出归纳矩匹配(IMM)框架,通过单阶段训练实现一步或少步生成,在ImageNet-256×256上以8步达到1.99 FID,超越扩散模型

Inductive Moment Matching: Few-Step Generation via Single-Stage Training

一、论文概述

项目内容
标题Inductive Moment Matching
作者Linqi Zhou, Stefano Ermon, Jiaming Song
机构Stanford University
论文https://arxiv.org/abs/2503.07565
发布2025-03-10 (v1), 最新 v7 (2025-05-14)
会议ICML

二、核心思想

问题定义

生成模型面临三难困境(Trilemma):

挑战说明
高保真输出生成高质量样本
高效推理减少推理步数
稳定训练避免训练崩溃

现有方法的局限:

方法问题
扩散模型推理步数多,速度慢
扩散蒸馏需要预训练初始化,训练不稳定
一致性模型(CMs)需要仔细调参,存在训练崩溃风险

解决方案概述

Inductive Moment Matching (IMM) 是一种新的生成模型框架,核心创新:

  • 单阶段训练:无需预训练初始化,从头训练
  • 少步生成:支持1步或少步推理
  • 分布级收敛:保证分布级收敛,而非点级一致性
  • 训练稳定:在各种超参数和标准模型架构下保持稳定

Figure 2: 一步采样器定义

三、技术架构

整体框架

IMM基于**随机插值(Stochastic Interpolants)**构建:

  1. 插值定义:连接数据分布 q(x)q(\mathbf{x}) 和先验分布 p(ϵ)p(\epsilon) 的随机过程
  2. 边际分布:qt(xt)q_t(\mathbf{x}_t) 表示时间 tt 处的边际分布
  3. 一步映射:学习从 qtq_t 到 qsq_s (s<ts < t) 的一步映射
  4. 归纳训练:通过数学归纳法训练,保证收敛

核心公式

随机插值定义: qt(xt∣x,ϵ)=N(It(x,ϵ),γt2I)q_t(\mathbf{x}_t | \mathbf{x}, \epsilon) = \mathcal{N}(\mathbf{I}_t(\mathbf{x}, \epsilon), \gamma_t^2 \mathbf{I})

其中:

  • It(x,ϵ)\mathbf{I}_t(\mathbf{x}, \epsilon) 是插值函数
  • γt\gamma_t 是噪声水平
  • 约束:I1(x,ϵ)=ϵ\mathbf{I}_1(\mathbf{x}, \epsilon) = \epsilon, I0(x,ϵ)=x\mathbf{I}_0(\mathbf{x}, \epsilon) = \mathbf{x}

自一致性插值(Self-Consistent Interpolants): xs∼qs∣t(xs∣x,xt)\mathbf{x}_s \sim q_{s|t}(\mathbf{x}_s | \mathbf{x}, \mathbf{x}_t)

满足:qs(xs)=∫qs∣t(xs∣x,xt)qt(xt)dxtq_s(\mathbf{x}_s) = \int q_{s|t}(\mathbf{x}_s | \mathbf{x}, \mathbf{x}_t) q_t(\mathbf{x}_t) d\mathbf{x}_t

IMM目标函数: min⁡θMMD2(qs(1),qs(2))\min_\theta \text{MMD}^2(q_s^{(1)}, q_s^{(2)})

其中:

  • qs(1)q_s^{(1)}:从 qrq_r 一步映射到 qsq_s 的分布
  • qs(2)q_s^{(2)}:从 qtq_t 一步映射到 qsq_s 的分布
  • MMD:最大均值差异(Maximum Mean Discrepancy)

Figure 3: 自一致性插值

简化参数化

DDIM插值: xs=αs∣txt+σs∣tϵ\mathbf{x}_s = \alpha_{s|t} \mathbf{x}_t + \sigma_{s|t} \epsilon

简化目标: L=Es,t[w(s,t)⋅MMD2(q^s(1),q^s(2))]\mathcal{L} = \mathbb{E}_{s,t} \left[ w(s,t) \cdot \text{MMD}^2(\hat{q}_s^{(1)}, \hat{q}_s^{(2)}) \right]

关键设计选择:

组件选择说明
映射函数 r(s,t)r(s,t)r(s,t)=sr(s,t) = s简化计算
时间分布 p(s,t)p(s,t)均匀采样s∼U(0,1)s \sim U(0,1), t∼U(s,1)t \sim U(s,1)
核函数Gaussian kernelk(x,y)=exp⁡(−∥x−y∥2/2σ2)k(\mathbf{x}, \mathbf{y}) = \exp(-\|\mathbf{x}-\mathbf{y}\|^2 / 2\sigma^2)
权重函数 w(s,t)w(s,t)w(s,t)=1w(s,t) = 1均匀权重

与一致性模型的关系

关键发现:一致性模型(CMs)是IMM的特例

  • CMs使用单粒子、一阶矩匹配
  • IMM使用多粒子、全矩匹配
  • 这解释了CMs训练不稳定的原因:单粒子估计方差大

四、核心创新

创新点说明优势
归纳训练范式通过数学归纳法训练保证分布级收敛
自一致性插值定义满足边际保持的插值理论保证
多粒子MMD使用M个粒子估计MMD训练稳定
单阶段训练无需预训练初始化简化流程
灵活步数支持1步到多步生成推理灵活

五、实验结果

CIFAR-10结果

方法族方法FID↓步数↓参数量
GANBigGAN6.951112M
GANGigaGAN3.451569M
GANStyleGAN-XL2.301166M
扩散NCSN++2.34200025.6M
扩散DDPM++2.28100025.6M
流匹配FM2.31100025.6M
蒸馏DMD2.19125.6M
蒸馏SEDD2.10125.6M
从头训练IMM (2步)1.98225.6M

关键发现:

  • IMM在2步生成中达到1.98 FID,超越所有从头训练的方法
  • 与蒸馏方法相当,但无需预训练
  • 使用pushforward采样器

ImageNet-256×256结果

方法架构FID↓步数↓
DiTDiT-XL2.27250
SiTDiT-XL2.06250
VARVAR-2B2.0610
IMMDiT-XL1.998
IMMDiT-XL1.9016

关键发现:

  • 8步IMM (1.99 FID) 超越DiT和SiT的250步结果
  • 16步IMM (1.90 FID) 超越VAR-2B的10步结果
  • 推理速度提升30×以上

Figure 6: FID与步数关系

训练稳定性

Figure 4: FID收敛曲线

稳定性验证:

测试结果
Fourier embedding (scale=16)稳定收敛(CMs在此设置下不稳定)
位置编码稳定收敛
不同粒子数MM越大越稳定

缩放行为

Figure 7: 缩放行为

缩放特性:

  • FID与训练/推理计算量强相关
  • 模型越大、步数越多,样本质量越高
  • 支持DiT-S/B/L/XL不同规模

Figure 8: 样本质量随模型规模和步数变化

消融实验

流调度和参数化:

参数化CIFAR-10ImageNet
id/cos3.77-
id/FM--
sEDM/cos2.083.75
sEDM/FM1.983.40
eFM2.063.18

关键发现:

  • Identity参数化性能较差
  • OT-FM调度+Euler参数化在大规模上更优
  • 小规模sEDM/FM最优,大规模eFM最优

六、相关工作

扩散模型与流匹配

  • 扩散模型:DDPM, NCSN++, DDPM++
  • 流匹配:FM, OT-FM
  • 随机插值:统一扩散和流匹配框架

扩散蒸馏

  • DMD:分布匹配蒸馏
  • SEDD:离散扩散蒸馏
  • 渐进蒸馏:逐步减少步数

一致性模型

  • CMs:点级一致性,训练不稳定
  • 伪目标:使用时间依赖的一致性目标
  • IMM的改进:分布级收敛,多粒子估计

生成矩匹配

  • GMMN:生成矩匹配网络
  • MMD GAN:使用MMD的GAN
  • IMM的区别:结合插值框架,支持少步生成

七、总结

核心贡献

  1. 提出IMM框架:新的少步生成模型,单阶段训练
  2. 理论保证:证明分布级收敛,统一多种现有方法
  3. 训练稳定:在各种超参数和架构下保持稳定
  4. SOTA性能:
    • CIFAR-10: 1.98 FID (2步)
    • ImageNet-256×256: 1.99 FID (8步)

技术影响

  • 少步生成新范式:无需蒸馏,从头训练
  • 理论统一:将一致性模型、GAN等统一到IMM框架
  • 实用性强:支持灵活步数,训练稳定
  • 可扩展:支持不同模型规模和分辨率

局限性

  1. 计算开销:MMD估计需要多个粒子,增加计算量
  2. 超参数敏感:粒子数M、核带宽等需要调整
  3. 架构依赖:在DiT上效果好,其他架构待验证
  4. 分辨率限制:主要评估256×256,更高分辨率待探索

八、参考资源

论文

相关工作

  • DMD: Distribution Matching Distillation
  • SEDD: Discrete Diffusion Distillation
  • VAR: Visual Autoregressive Modeling
  • GMMN: Generative Moment Matching Networks