Skip to main content
QUICK REVIEW

[论文解读] Training generative neural networks via Maximum Mean Discrepancy optimization

Gintare Karolina Dziugaite, Daniel M. Roy|arXiv (Cornell University)|May 14, 2015
Generative Adversarial Networks and Image Synthesis参考文献 8被引用 222
一句话总结

本文提出 MMD 网络,一种通过非参数两样本检验最小化生成数据与真实数据分布之间最大均值差异(MMD)来训练深度生成神经网络的方法。与依赖对抗训练的 GAN 不同,MMD 网络使用 MMD 的无偏经验估计作为可微分目标,实现稳定优化,并具备理论泛化界。

ABSTRACT

We consider training a deep neural network to generate samples from an unknown distribution given i.i.d. data. We frame learning as an optimization minimizing a two-sample test statistic---informally speaking, a good generator network produces samples that cause a two-sample test to fail to reject the null hypothesis. As our two-sample test statistic, we use an unbiased estimate of the maximum mean discrepancy, which is the centerpiece of the nonparametric kernel two-sample test proposed by Gretton et al. (2012). We compare to the adversarial nets framework introduced by Goodfellow et al. (2014), in which learning is a two-player game between a generator network and an adversarial discriminator network, both trained to outwit the other. From this perspective, the MMD statistic plays the role of the discriminator. In addition to empirical comparisons, we prove bounds on the generalization error incurred by optimizing the empirical MMD.

研究动机与目标

  • 通过用非参数两样本检验替代判别器,开发一种替代对抗训练的稳定方法,用于深度生成模型。
  • 将生成模型学习问题表述为最小化生成数据与真实数据分布之间的经验 MMD。
  • 提供优化经验 MMD 而非真实总体 MMD 所产生的泛化误差的理论界。
  • 证明基于 MMD 的优化可避免 GAN 中常见的训练不稳定性,同时保持优异的样本质量。

提出的方法

  • 将生成模型训练问题表述为最小化数据分布与生成器输出分布之间的 MMD。
  • 使用基于核两样本检验推导出的 MMD 无偏估计量作为训练目标。
  • 通过在经验 MMD 上进行梯度下降来优化生成器参数,将 MMD 视为可微分损失函数。
  • 应用 McDiarmid 不等式与 Rademacher 复杂度界,推导经验 MMD 估计量的泛化误差界。
  • 使用具有特征性质的通用再生核希尔伯特空间(RKHS)核,确保仅当分布相等时 MMD 为零。
  • 通过在各种矩条件下的估计误差界,建立基于 MMD 的训练的理论收敛保证。

实验结果

研究问题

  • RQ1基于 MMD 的优化能否作为深度生成模型对抗训练的稳定、非对抗性替代方法?
  • RQ2经验 MMD 估计量的泛化误差如何随样本大小与核性质变化?
  • RQ3在分布差异方面,基于 MMD 的训练收敛性可提供哪些理论保证?
  • RQ4在合成数据与真实数据上,MMD 网络在样本质量与训练稳定性方面与对抗网络相比表现如何?
  • RQ5在何种条件下,基于 MMD 的目标可确保生成器学习到真实数据分布?

主要发现

  • MMD 网络框架通过用闭式 MMD 统计量替代对抗判别器,实现稳定训练,避免了 GAN 中常见的模式崩溃与训练不稳定性。
  • 理论分析表明,经验 MMD 估计量的泛化误差衰减速度为 $ O(M^{-1/2}) $(当 $ p < 2 $ 时)、$ O(M^{-1/2} ext{log}^{3/2}(M)) $(当 $ p = 2 $ 时)和 $ O(M^{-1/p}) $(当 $ p > 2 $ 时),其中 $ M $ 为样本大小。
  • 当生成器族足够丰富且核具有特征性质时,在非参数极限下近似误差为零,确保 MMD = 0 意味着分布相等。
  • 在合成与真实数据上的实验结果表明,MMD 网络成功将生成器输出分布与真实数据分布对齐,表现为训练迭代过程中 MMD 值持续下降。
  • 该方法在样本质量方面与 GAN 具有竞争力,且无需交替更新判别器与生成器,简化了训练动力学。
  • 理论分析证实,经验 MMD 估计量以高概率集中在真实 MMD 周围,其尾部界随样本大小呈指数衰减。

更好的研究,从现在开始

从阅读论文到最终审阅,大幅缩短您的研究时间。

无需绑定信用卡

本解读由 AI 生成,并经人工编辑审核。