MMD loss

Posted by 孙睿 on August 12, 2026

这篇BLOG记录一下机器学习中一个损失函数,MMD(maximum mean discrepancy)的相关知识。

背景介绍

最早接触到这个损失函数是在计算生物学,基因扰动预测任务中,对应的问题场景如下: 给定一个control条件下的细胞表达矩阵$X = (x_1, x_2, …, x_m) \in R^{m\times p}$,其中基因数目为$p$,加入一种扰动后(基因敲除或者化学试剂),测得一个perturb条件下的细胞表达矩阵$Y = (y_1, y_2, …, y_n) \in R^{n\times p}$。希望建立一个扰动预测模型$f: R^{p} \rightarrow R^{p}$,能够预测一个control cell 在加入扰动后的 perturb 状态。

如果 control-perturb 的配对关系已知,即清楚每个control cell $x_i$ 扰动后的 perturb cell $y_j$ 的情况下,这个问题是简单的回归建模,即 \(\min_\theta \sum_{(i,j)\in \mathbb{P}} \lVert f_{\theta}(x_i) - y_j \rVert^2\) 这里$\mathbb{P}$代表全部的配对集合。

但在实际问题中,这个关系是不知道的,这是由于基因测序技术是一种破坏性的方式,需要杀死细胞才能测得其基因表达情况。因此我们无法同时测量得到每个细胞扰动前和扰动后的基因表达向量。前面的$X,Y$是一种分布层面的测量,即将某个细胞系作为control,分为两组,一组直接测量,一组扰动后测量。

在缺乏配对关系的情况下,一个粗暴的方式是将$\mathbb{P}$处理成全部的$m\times n$种 control-perturb 的组合。

\[\min_\theta \sum_{i=1}^m \sum_{j=1}^n \lVert f_{\theta}(x_i) - y_j \rVert^2\]

但是容易验证,在这种情况下,最优的$f_{\theta}(\cdot)$ 实际上是有显示解的,即 \(f_{\theta}(x_i) = \frac{1}{n}\sum_{j=1}^{n} y_j\)

也就是说,只要使用全部扰动后细胞的均值做预测即是最优预测,预测模型完全失去对扰动随机性的建模(扰动是否发生,在不同的细胞间扰动效应是否一致等等)。这也是扰动建模这个研究方向早期(2025.06之前)的一个很明显的弊端,2025年上半年在晶泰科技实习时,我独立发现了领域内的这个系统性偏差并在公司内部的评测中引入了相关的修正策略,2025年下半年领域顶刊 Nature Biotechnology 中也有论文阐述了这个问题。

MMD 简述

在这么长的背景介绍后,终于迎来了最终的问题,既然之前的随机配对的方式存在弊端,那么能否之间建模预测的扰动分布和真实的扰动分布之间的差异,避免前面提到的均值坍塌问题。一个现有的机器学习解决思路,MMD成为了一个相对主流的方案。下面快速的过一遍MMD的主要思路。

  • 给定两个概率分布$P,Q$,怎样量化两个分布的差异?
  • 假如两个分布一致,那么它们在任意的函数映射下的均值都应该一致。具体的,对任意的一个函数$\phi(x)$,都应该有$E_{x\sim P} \phi(x) = E_{x\sim Q} \phi(x)$。
  • 暂不讨论严格证明的,我们找到了一个满足要求的映射$\phi(x)$将样本映射到一个希尔伯特空间$\mathcal{H}$,这个映射使得对任意的$P,Q$, 定义$\mu_P = E_{x\sim P} \phi(x), \mu_Q = E_{x\sim Q} \phi(x)$,都满足 $\lVert \mu_P - \mu_Q \rVert_{\mathcal{H}}^2 >0 \;\;\text{if}\;\; P\neq Q$。
  • 定义MMD为$MMD(P,Q)=\lVert \mu_P - \mu_Q \rVert_{\mathcal{H}}^2$,通过优化这个目标,可以使得两个分布$P,Q$更加接近(分布接近可能存在多种度量方式,例如MMD度量下的分布接近可能在wasserstein度量下并不接近,这个例子会在后面讨论)。
  • 考虑实际计算,$P,Q$形式未知,但可以给出经验估计$\mu_P = \frac{1}{m}\sum_{i=1}^m \phi(x_i), \mu_Q = \frac{1}{n}\sum_{i=1}^n \phi(y_i)$,我们需要计算的是$\lVert \frac{1}{m}\sum_{i=1}^m \phi(x_i) - \frac{1}{n}\sum_{i=1}^n \phi(y_i) \rVert^2_{\mathcal{H}} $
  • 再回到前面的$\phi(x)$选取,这个映射如果选择核函数对应的高维映射,前面的范数计算将会非常容易,这里不再展开核函数相关的讨论,之要知道存在$k(x,y) = \langle \phi(x), \phi(y) \rangle_{\mathcal{H}}$ 即可。

MMD的具体计算公式可以处理为,先把范数平方按内积展开,再代入经验估计:

\[\begin{aligned} MMD(P,Q) &= \lVert \mu_P - \mu_Q \rVert^2_{\mathcal{H}} \\ &= \langle \mu_P, \mu_P \rangle_{\mathcal{H}} + \langle \mu_Q, \mu_Q \rangle_{\mathcal{H}} - 2\langle \mu_P, \mu_Q \rangle_{\mathcal{H}} \end{aligned}\]

其中每一项内积都可以借助核函数 $k(x,y) = \langle \phi(x), \phi(y) \rangle_{\mathcal{H}}$ 直接展开为对样本对的求和:

\[\begin{aligned} \langle \mu_P, \mu_P \rangle_{\mathcal{H}} &= \frac{1}{m^2}\sum_{i=1}^{m}\sum_{j=1}^{m} \langle \phi(x_i), \phi(x_j) \rangle_{\mathcal{H}} = \frac{1}{m^2}\sum_{i,j=1}^{m} k(x_i, x_j) \\ \langle \mu_Q, \mu_Q \rangle_{\mathcal{H}} &= \frac{1}{n^2}\sum_{i,j=1}^{n} k(y_i, y_j) \\ \langle \mu_P, \mu_Q \rangle_{\mathcal{H}} &= \frac{1}{mn}\sum_{i=1}^{m}\sum_{j=1}^{n} k(x_i, y_j) \end{aligned}\]

整理后得到基于核函数的MMD经验估计公式:

\[MMD(P,Q) = \frac{1}{m^2}\sum_{i,j=1}^{m} k(x_i, x_j) + \frac{1}{n^2}\sum_{i,j=1}^{n} k(y_i, y_j) - \frac{2}{mn}\sum_{i=1}^{m}\sum_{j=1}^{n} k(x_i, y_j)\]

可以看到,最终公式中只涉及核函数在样本对上的取值,避开了显式构造高维映射$\phi(x)$,这也是MMD在实际中易于计算的关键。常用的核函数$k(x,y)$有高斯核,$k(x,y) = \exp(-\gamma\lVert x-y \rVert^2)$。

REMARK 1

以上是个人学习MMD中的一些思考,从具体的问题出发,引出分布差异度量这一需求。为了定义两个分布的差异,可以把分布映射到一个新的度量空间中,使用均值之间的度量差异刻画分布差异。这个映射很关键,因为希望对于任意不同的分布进行映射后都能基于均值定义差异。

LLM给出的回答中,高斯核对应的映射$\phi(x)$是可以满足这一需求的,在找到这一映射后最终导出最后的计算结果。

REMARK 2

需要注意的是,分布之间的相似性并不像欧式空间中的向量距离那样显然:在一种度量下可能认为$P_1$比$P_2$更接近$Q$,换一种度量结论可能完全相反。下面用三个简单的离散分布构造一个例子来说明,以下case借助AI生成,未进行严格的计算校验。


第一步:设定分布

在实数轴 $\mathbb{R}$ 上考虑三个分布,其中 $Q$ 为目标分布。为了计算方便,使用狄拉克分布 $\delta_c$(表示在点 $c$ 处概率为 1 的分布)。

  • 目标分布 $Q$:集中在原点。 \(Q = \delta_0\)
  • 分布 $P_1$:在 $x=1$ 和 $x=-1$ 处各有一个点,权重各 $0.5$。 \(P_1 = 0.5\delta_1 + 0.5\delta_{-1}\)
  • 分布 $P_2$:在 $x=0.5$ 处有一个点。 \(P_2 = \delta_{0.5}\)

EMD(推土机距离)的计算

EMD 即 1-Wasserstein 距离,定义为最优运输问题的解: \(EMD(P, Q) = W_1(P,Q) = \inf_{\gamma \in \Pi(P,Q)} \int |x-y| \, d\gamma(x,y)\) 其中 $\Pi(P,Q)$ 是所有以 $P,Q$ 为边缘分布的联合分布的集合。对一维分布而言,它还有等价的 CDF 形式,即累积分布函数之差的积分: \(W_1(P,Q) = \int_{-\infty}^{+\infty} |F_P(x) - F_Q(x)| \, dx\)

对于离散点质量分布,最优运输可以直观地理解为:把 $P$ 的质量以最小成本搬到 $Q$ 的位置,总成本 = 搬运量 × 搬运距离。

  • $EMD(P_1, Q)$:需要把 $+1$ 处 $0.5$ 的质量、$-1$ 处 $0.5$ 的质量分别搬到原点。 \(EMD(P_1, Q) = 0.5 \times |1-0| + 0.5 \times |-1-0| = 0.5 + 0.5 = 1\)
  • $EMD(P_2, Q)$:只需把 $+0.5$ 处全部 $1$ 的质量搬到原点。 \(EMD(P_2, Q) = 1 \times |0.5-0| = 0.5\)

因此 $EMD(P_1, Q) = 1 > EMD(P_2, Q) = 0.5$,EMD 判定 $P_2$ 更接近 $Q$。下面会看到,MMD 在核带宽极大时的排序结论会与 EMD 相反。


第二步:MMD 的推导计算

使用高斯核 $k(x,y) = \exp(-\gamma(x-y)^2)$。这里 $\gamma = \frac{1}{2\sigma^2}$,$\gamma$ 越小,代表核带宽 $\sigma$ 越大(感受野越宽)。

根据 MMD 的公式: \(MMD^2(P, Q) = \mathbb{E}_{x,x'\sim P}[k(x,x')] + \mathbb{E}_{y,y'\sim Q}[k(y,y')] - 2\mathbb{E}_{x\sim P, y\sim Q}[k(x,y)]\)

1. 计算 $MMD^2(P_1, Q)$

  • 第一项($P_1$ 内部): \(\mathbb{E}[k(x,x')] = 0.5^2 k(1,1) + 0.5^2 k(-1,-1) + 2 \times 0.5^2 k(1,-1)\) 因为 $k(x,x)=1$,且 $k(1,-1) = \exp(-4\gamma)$,所以: \(= 0.25(1) + 0.25(1) + 0.5\exp(-4\gamma) = 0.5 + 0.5\exp(-4\gamma)\)
  • 第二项($Q$ 内部): \(\mathbb{E}[k(y,y')] = k(0,0) = 1\)
  • 第三项(交叉项): \(\mathbb{E}[k(x,y)] = 0.5 k(1,0) + 0.5 k(-1,0) = 0.5\exp(-\gamma) + 0.5\exp(-\gamma) = \exp(-\gamma)\)

组合起来: \(MMD^2(P_1, Q) = 1.5 + 0.5\exp(-4\gamma) - 2\exp(-\gamma)\)

2. 计算 $MMD^2(P_2, Q)$

  • 第一项($P_2$ 内部):$k(0.5, 0.5) = 1$
  • 第二项($Q$ 内部):$k(0,0) = 1$
  • 第三项(交叉项):$k(0.5, 0) = \exp(-0.25\gamma)$

组合起来: \(MMD^2(P_2, Q) = 2 - 2\exp(-0.25\gamma)\)


第三步:数值分析

把第二步推导出的两个公式画出来,直观观察两条曲线随 $\gamma$ 的变化(对数横轴,$\gamma$ 越小代表核带宽 $\sigma$ 越大)。两条曲线在 $\gamma \approx 0.235$ 处交叉:交叉点左侧(大带宽区)$MMD(P_1,Q) < MMD(P_2,Q)$,正是下面要讨论的反例区;右侧(小带宽区)$MMD(P_1,Q) > MMD(P_2,Q)$,与几何直觉/EMD 一致:

MMD(P1,Q)与MMD(P2,Q)随核参数γ的变化

接下来看两个具体的情况:

情况 A:$\gamma$ 较大(小带宽,感受野窄)

假设 $\gamma = 2$(即 $\sigma \approx 0.5$,核函数很尖锐,只关注极近距离)。

  • $MMD^2(P_1, Q) = 1.5 + 0.5\exp(-8) - 2\exp(-2) \approx 1.5 + 0 - 2(0.135) = 1.23$
  • $MMD^2(P_2, Q) = 2 - 2\exp(-0.5) \approx 2 - 2(0.606) = 0.788$ 此时,$MMD(P_1, Q) > MMD(P_2, Q)$,MMD 与直觉一致,认为 $P_2$ 更接近 $Q$。

情况 B:$\gamma \to 0$(极大带宽,感受野极宽)

假设 $\gamma = 0.1$(即 $\sigma \approx 2.23$,核函数非常平缓)。

  • $MMD^2(P_1, Q) = 1.5 + 0.5\exp(-0.4) - 2\exp(-0.1) \approx 1.5 + 0.335 - 1.810 = 0.025$
  • $MMD^2(P_2, Q) = 2 - 2\exp(-0.025) \approx 2 - 1.951 = 0.049$ 此时,$MMD(P_1, Q) < MMD(P_2, Q)$,MMD 认为 $P_1$ 比 $P_2$ 更接近 $Q$,与推土机距离的结论完全相反。

第四步:反转的原因

为什么大带宽下会出现这种反转?用泰勒展开来看。

当 $\gamma \to 0$ 时,高斯核可以近似为 $k(x,y) = \exp(-\gamma(x-y)^2) \approx 1 - \gamma(x-y)^2$。 将这个近似代入 MMD 公式,经过代数化简(利用方差和均值的性质),可以得到近似公式: \(MMD^2(P, Q) \approx 2\gamma (\mu_P - \mu_Q)^2\) (其中 $\mu_P, \mu_Q$ 分别是分布 $P$ 和 $Q$ 的均值)

这说明当核带宽极大($\gamma \to 0$)时,MMD 退化成了仅仅比较两个分布的均值(一阶矩)。

回到上面的分布:

  • $Q$ 的均值是 $0$。
  • $P_2$ 的均值是 $0.5$,与 $Q$ 的均值差是 $0.5$。代入近似公式:$MMD^2 \approx 2\gamma (0.5)^2 = 0.5\gamma$。
  • $P_1$ 的均值是 $0.5 \times 1 + 0.5 \times (-1) = 0$,与 $Q$ 的均值差是 $0$。代入近似公式:$MMD^2 \approx 2\gamma (0)^2 = 0$(实际由更高阶的 $O(\gamma^2)$ 项决定)。

$P_1$ 虽然在空间上分散在 $1$ 和 $-1$ 处(搬运成本高),但它的重心(均值)落在了原点。当高斯核的带宽极大时,只能看到两个分布的重心是否重合,于是判定 $P_1$ 和 $Q$ 极其相似。而 $P_2$ 虽然离原点更近,但它的重心在 $0.5$,在大带宽下这个重心的偏移会被 MMD 捕捉到。


小结

回到最初的问题:$MMD(P_1, Q) > MMD(P_2, Q)$ 能说明 $P_2$ 比 $P_1$ 更接近 $Q$ 吗?

  1. 在几何/搬运意义上(Wasserstein):不能。如上面的例子所示,$P_1$ 的 MMD 可以比 $P_2$ 更小,但 $P_1$ 的搬运成本显然更高。
  2. 在 RKHS 均值嵌入意义上(MMD 本身):能。MMD 忠实地反映了在当前核带宽 $\gamma$ 下,$P_1$ 的均值嵌入确实比 $P_2$ 的均值嵌入离 $Q$ 更近。
  3. 工程上,核带宽($\sigma$ 或 $\gamma$)的选择是 MMD 应用的关键:
    • 带宽太大($\gamma \to 0$),MMD 退化为均值比较,丢失所有高阶形状信息。
    • 带宽太小($\gamma \to \infty$),MMD 退化为逐点匹配,对噪声极其敏感。
    • 在实际应用中(如深度学习的 MMD Loss),通常使用多核 MMD(Multiple Kernel MMD),将多个不同 $\gamma$ 的高斯核线性组合,避免模型被单一带宽的度量偏好影响。
{/% if page.mathjax %}{/% include mathjax_support.html %}{/% endif %}