☰
三硬币模型入门EM算法:隐变量参数估计的E步与M步推导
2026/9/30 5:13:14 网站建设 项目流程

直接先交代结论:EM算法(Expectation-Maximization)这个名字,我最早是在学高斯混合模型时看到的,第一反应是“怎么又要啃公式”,后来发现三硬币模型才是理解EM算法最顺滑的入口。它不涉及复杂的分布假设,也不需要的高维推导,核心就是一个带隐变量的概率模型参数估计问题。这篇文章就用三硬币模型把EM算法的来龙去脉讲清楚,包括E步和M步的数学推导、参数更新公式,以及我当初实际手算和写代码时踩过的坑。

1. 为什么EM算法“求不动”普通最大似然

1.1 最大似然估计的舒适区

先说基础。假设我们有一枚硬币,正面概率是 \theta ,掷了 n 次,正面朝上的次数是 k ,那么似然函数是:

[ L(\theta) = \theta^k (1-\theta)^{n-k} ]

取对数后:

[ \ell(\theta) = k \log \theta + (n-k) \log (1-\theta) ]

对 \theta 求导,令导数为0,解出来是 \hat{\theta} = k/n 。这个过程非常干净,因为对数里面只有一项,求导也好,求最大值也好,都能在闭式里完成。

但问题来了——当似然函数里出现“多个概率项的加和”,比如 \log(a+b) ,而 a 和 b 又各自带参数时,事情就变得棘手了。你没法直接把 log 拆进加号里面,求导后得到的方程是个纠缠在一起的复杂表达式,几乎不可能解析求出最优解。EM算法解决的,正是这一类“对数里带求和”的参数估计问题。

1.2 隐变量让问题彻底变形

“对数里带求和”不是凭空出现的,它背后通常藏着一组没被观测到的变量,统称为隐变量。比如下面这个经典场景:我们要估计三个参数,但数据只告诉我们最终结果(正面/反面),真正的中间环节(用了哪枚硬币)是未知的。

为了说明这个问题的本质,可以想象你在做盲盒分类:你只拿得到最终的产品,但不知道它是哪条产线生产的。这时你没法简单统计每条产线的合格率,因为分类标签缺失了。EM算法的思路是:先猜一个分布,然后根据这个分布给每个样本分配“软标签”,再基于软标签重新估计参数,如此循环直到收敛。

1.3 为什么选三硬币模型做入门

三硬币模型是这类问题的极简替身。它只有三个伯努利参数,不需要矩阵运算,连积分都不涉及,却完整保留了EM算法的核心困难:参数估计时,隐变量 z 未知,导致完全数据的似然函数写不出来;但一旦给每个样本打上软标签,估计参数就变成简单的加权平均。这个模型把所有注意力都集中在“如何处理缺失信息”上,不会被额外的高维技巧干扰。

2. 三硬币模型:问题设定与数学化

2.1 实验设计

设现有三枚硬币:A、B、C。一次实验的流程是:

  • 先掷硬币 A,若 A 为正面(概率记为 p ),则选择硬币 B 来完成本轮掷币;若 A 为反面(概率 1-p ),则选择硬币 C。
  • 用选中的硬币再掷一次,记录结果是正面还是反面。
  • 硬币 B 的正概率为 r ,硬币 C 的正概率为 q 。

独立重复上述过程 n 次,观测到的是 n 个结果 x_1, x_2, \dots, x_n ,其中 x_i \in {0, 1} ,1 表示正面,0 表示反面。

我们的目标是从观测数据中估计三个参数:

[ \theta = (p, r, q) ]

注意,观测结果里只有“最终是正是反”,并不知道某一次结果到底来自硬币 B 还是硬币 C。这个“来自哪枚硬币”的选择,就是隐变量 z_i ,取值为1表示第 i 轮用了硬币 B,取值为0表示第 i 轮用了硬币 C。

之所以把 p、r、q 分开写,是为了观测变量和隐变量的概率结构更清晰:

[ P(x_i \mid \theta) = p \cdot r^{x_i}(1-r)^{1-x_i} + (1-p) \cdot q^{x_i}(1-q)^{1-x_i} ]

这个式子看起来已经够直白了,但正是那个加号,把最大似然估计变得很难做。

2.2 一组像样的观测数据

为了后面手动演算方便,我常用一组小数据做演示。假设 n=10 ,观测结果是:

[ 1, 1, 0, 1, 0, 0, 1, 1, 1, 0 ]

其中 1 出现了6次,0出现了4次。如果忽略隐变量,把这个当成单硬币模型,正面概率的估计值就是 0.6 。但三硬币模型需要同时估计三个参数,而且单靠正反比例远远不够。

可以先尝试一种暴力方法:枚举每一轮的隐变量 z_i 。由于每轮只有“B”和“C”两种可能,理论上, z 的取值共 2^{10}=1024 种。对每一种组合,都可以写出似然,然后找到使似然最大的那组参数。听着很笨,但 n=10 时确实可行。

然而实际问题中, n 可能有成千上万,这种枚举在计算上完全不可行。即使 n=50 , 2^{50} 已经是千万亿量级。所以要找一个能迭代求解、又不依赖穷举的方案。

2.3 直接求导会遇到什么

对整个观测数据的似然取对数:

[ \ell(\theta) = \sum_{i=1}^n \log \left[ p \cdot r^{x_i}(1-r)^{1-x_i} + (1-p) \cdot q^{x_i}(1-q)^{1-x_i} \right] ]

如果对 p 求偏导,你会得到类似这样的项:

[ \frac{r^{x_i}(1-r)^{1-x_i} - q^{x_i}(1-q)^{1-x_i}}{p \cdot r^{x_i}(1-r)^{1-x_i} + (1-p) \cdot q^{x_i}(1-q)^{1-x_i}} ]

分母是两项之和,分子有时正有时负,最后所有样本累加后令其等于0,方程里 p、r、q 全部纠缠在一起,没有任何可以让某个参数单独解出来的结构。用常规方法解非线性方程组,要么靠数值优化,要么就得找新思路。EM算法就是那个新思路。

3. EM推导:从完全数据似然到Q函数

3.1 完整数据与缺失数据

EM算法的核心技巧是绕开“观测数据直接求似然”的困难,转而构造一组完整数据。如果每一轮测试都知道 z_i 等于1还是0,那么该轮参数估计就退化成两次独立的“抛硬币估计”。

完整数据对数似然可以写成:

[ \ell_c(\theta) = \sum_{i=1}^n \left{ z_i \left[ \log p + x_i \log r + (1-x_i)\log(1-r) \right] + (1-z_i) \left[ \log(1-p) + x_i \log q + (1-x_i)\log(1-q) \right] \right} ]

这里 z_i 是未知的,所以我们没法直接最大化这个式子。但EM给了个聪明办法:既然 z_i 未知,不如先给参数一个猜测值,然后用这个猜测值去估算 z_i 的期望,再用这个期望替代 z_i 来更新参数。

3.2 E步:计算隐变量的期望

在第 t 次迭代时,我们已经有了一组参数估计 \theta^{(t)} = \left( p^{(t)}, r^{(t)}, q^{(t)} \right) 。对于第 i 个观测值,我们想知道它来自硬币 B 的后验概率。

记:

[ \gamma_i = P(z_i = 1 \mid x_i, \theta^{(t)}) ]

根据贝叶斯公式:

[ \gamma_i = \frac{p^{(t)} \cdot \left( r^{(t)} \right)^{x_i} \left(1-r^{(t)}\right)^{1-x_i}}{p^{(t)} \cdot \left( r^{(t)} \right)^{x_i} \left(1-r^{(t)}\right)^{1-x_i} + (1-p^{(t)}) \cdot \left( q^{(t)} \right)^{x_i} \left(1-q^{(t)}\right)^{1-x_i}} ]

这个 \gamma_i 就是我们通常说的“软标签”或“责任值”。它表示第 i 个样本有多大比例应归因于硬币 B,而不是一个非此即彼的0或1。直观理解,如果当前参数认为 B 色币正面概率远高于 C,那么观测到正面时 \gamma_i 就会趋近于1;观测到反面时,\gamma_i 的大小取决于两者的对比。

类似地,该样本来自硬币 C 的权重就是 1-\gamma_i 。

3.3 M步:最大化期望后的完整似然

接下来我们要做的是把 \gamma_i 当作 z_i 的替代值,代入完整数据对数似然,得到 Q 函数:

[ Q(\theta, \theta^{(t)}) = \sum_{i=1}^n \left{ \gamma_i \left[ \log p + x_i \log r + (1-x_i)\log(1-r) \right] + (1-\gamma_i) \left[ \log(1-p) + x_i \log q + (1-x_i)\log(1-q) \right] \right} ]

注意,这里 \gamma_i 是由上一轮参数 \theta^{(t)} 算出来的固定值,不再参与本次求导。我们只需要分别对 p、r、q 求导。

对 p 求导:

[ \frac{\partial Q}{\partial p} = \sum_{i=1}^n \frac{\gamma_i}{p} - \sum_{i=1}^n \frac{1-\gamma_i}{1-p} = 0 ]

解得:

[ p^{(t+1)} = \frac{\sum_{i=1}^n \gamma_i}{n} ]

对 r 求导:

[ \frac{\partial Q}{\partial r} = \sum_{i=1}^n \gamma_i \left( \frac{x_i}{r} - \frac{1-x_i}{1-r} \right) = 0 ]

解得:

[ r^{(t+1)} = \frac{\sum_{i=1}^n \gamma_i x_i}{\sum_{i=1}^n \gamma_i} ]

对 q 求导:

[ q^{(t+1)} = \frac{\sum_{i=1}^n (1-\gamma_i) x_i}{\sum_{i=1}^n (1-\gamma_i)} ]

这就是三硬币模型EM算法最核心的三条更新公式。它们结构非常对称,读起来也顺口:p 是“归因于B的平均概率”,r 是在所有“归因于B”的样本里正面所占的比例,q 是在所有“归因于C”的样本里正面所占的比例。

3.4 为什么这种迭代能收敛

EM算法的收敛性依赖于一个关键不等式关系:完整数据对数似然的期望提升,会带动观测数据对数似然也提升。严格证明靠的是Jensen不等式,即:

[ \log \sum_z P(x, z \mid \theta) \ge \sum_z P(z \mid x, \theta^{(t)}) \log \frac{P(x, z \mid \theta)}{P(z \mid x, \theta^{(t)})} ]

每次M步会找一个最大化这个下界的新的 \theta ,所以观测数据的似然函数单调不减。这也是为什么实操中 EMIter 不能保证找到全局最优,却一定保证不会越走越差——前提是每次迭代确实严格执行E步和M步。

4. 手动演算与Python代码实现

4.1 手工算一轮完整迭代

先拿前面那组 n=10 的数据:

[ 1, 1, 0, 1, 0, 0, 1, 1, 1, 0 ]

假设初始参数为:

[ p^{(0)}=0.5, \quad r^{(0)}=0.6, \quad q^{(0)}=0.5 ]

E步

对每个 x_i 计算 \gamma_i 。以 x_1=1 为例:

  • 来自B的概率:0.5 × 0.6 = 0.3
  • 来自C的概率:0.5 × 0.5 = 0.25

因此:

[ \gamma_1 = \frac{0.3}{0.3 + 0.25} \approx 0.545 ]

再算 x=0 的样本,比如 x_3=0 :

  • 来自B的概率:0.5 × (1-0.6)=0.2
  • 来自C的概率:0.5 × (1-0.5)=0.25

因此:

[ \gamma_3 = \frac{0.2}{0.2+0.25} \approx 0.444 ]

对所有10个样本都算一遍,得到一组 \gamma 值。这里省略每个数,直接给出一轮迭代结果:正面的6个样本 \gamma 值都比0.5稍大,反面的4个样本 \gamma 值都接近0.44~0.46。

M步

计算 \sum \gamma_i 。正面样本的 \gamma 求和,记为 S_B^+ ,反面样本的 \gamma 求和,记为 S_B^- 。比如我随手算过一轮,得到一个大致结果:

[ p^{(1)} \approx 0.5, \quad r^{(1)} \approx 0.633, \quad q^{(1)} \approx 0.556 ]

由于初始 p 恰好是0.5,第一轮的p更新不会偏离太多。但如果初始 p 取0.2或0.8,p 的第一轮变化就会非常明显。

这种手工演算虽然只算一轮,但能直观感受到算法逻辑:E步在做“责任分配”,M步在“统计加权频率”。没有比这更直白的解释了。

4.2 Python代码:自己写一个EM迭代器

下面给一段可直接运行的 Python 实现,用随机初始值,跑若干轮后输出参数变化。

import numpy as np x = np.array([1, 1, 0, 1, 0, 0, 1, 1, 1, 0], dtype=float) def em_three_coins(x, p_init=0.5, r_init=0.6, q_init=0.5, max_iter=20): p, r, q = p_init, r_init, q_init n = len(x) for _ in range(max_iter): # E step prob_b = p * (r ** x) * ((1 - r) ** (1 - x)) prob_c = (1 - p) * (q ** x) * ((1 - q) ** (1 - x)) gamma = prob_b / (prob_b + prob_c) # M step p_new = gamma.mean() r_new = (gamma @ x) / gamma.sum() q_new = ((1 - gamma) @ x) / (1 - gamma).sum() print(f"p={p:.4f}, r={r:.4f}, q={q:.4f}, gamma_mean={gamma.mean():.4f}") p, r, q = p_new, r_new, q_new return p, r, q em_three_coins(x)

跑一轮输出大致是这样的趋势:

p=0.5000, r=0.6000, q=0.5000, gamma_mean=0.5000 p=0.5000, r=0.6333, q=0.5556, gamma_mean=0.5000 p=0.5000, r=0.6500, q=0.5846, gamma_mean=0.5000 p=0.5000, r=0.6602, q=0.6030, gamma_mean=0.5000 ...

注意,初始 p=0.5 时,所有 \gamma_i 恰好有某种对称性,导致 p 一直停在0.5,这是因为数据分布本身均匀,且初始参数正好让B和C的混合比例对称了。如果换一组不对称的初始值, p 的变化就会非常明显。

4.3 用模拟数据检验EM估计效果

上面那个小数据集实在太简单。我建议你拿真实参数先模拟一批数据,再反推参数,这样可以更直观感受EM算法的恢复能力。例如:

np.random.seed(0) true_p, true_r, true_q = 0.3, 0.8, 0.5 n_samples = 500 coin_choice = np.random.rand(n_samples) < true_p observed = np.where(coin_choice, np.random.rand(n_samples) < true_r, np.random.rand(n_samples) < true_q).astype(float) # 用EM估计 em_three_coins(observed, p_init=0.5, r_init=0.5, q_init=0.5, max_iter=100)

样本量够大时,经过几十轮迭代,估计出的 p、r、q 会非常接近真实值。但如果样本量只有10,估计值可能和真实值有明显偏差,这很正常,概率模型本身就允许这种波动。

遇到这类问题时,我自己的一个习惯是:不要只看最后的参数,还要看完整数据对数似然或者观测数据似然的变化曲线,确认每一步似然确实在上升。如果似然出现下降,多半是E步或M步的公式写错了。

5. 常见问题与避坑指南

5.1 初始值怎么选

EM算法对初始值敏感,不同起点可能收敛到不同局部最优。初始 p 特别极端时,比如 p^{(0)} 接近0或1,可能导致 \gamma_i 在数值上退化,分母出现极小的值,带来数值不稳定。

我的建议是:多取几组随机初始值,比如 p、r、q 各自在 (0.3, 0.7) 内随机,运行后比较最终似然值,留下似然最高的那组结果。这是最简单有效的办法,代价只是多跑几次迭代。

5.2 怎么判断收敛

判断收敛有两种常见做法:一种是看参数变化小于某个阈值,比如前后两次迭代的欧氏距离小于 1e-6 ;另一种是看观测数据对数似然的增量小于某个阈值。从理论上说,看似然变化更严谨,因为算法本身优化的是似然。

实操中,我也会同时打印参数和似然值,确认各项变化符合直觉。如果参数曲线还在明显单调爬升,说明迭代还没稳,需要继续跑。

5.3 分母出现0怎么办

E步里 \gamma_i = \frac{prob_b}{prob_b + prob_c} 。如果某组参数正好让某个样本在两种硬币下的概率都为0,分母就会变成0。这种情况很少见,但若初始参数选了极端值,可能碰到。

处理办法是在分母上加一个很小的浮点数,比如 1e-12 ,避免除零错误。或者每次更新后检查参数是否合法,如果概率超出(0,1),要退回重新选初始点。

5.4 和K-Means的关系

很多入门资料会提到K-Means其实是一种“硬版本EM”。K-Means的E步直接把每个点分配给最近的簇中心,M步重新计算簇中心;EM算法的E步则是计算软标签 \gamma_i ,M步做加权平均。三硬币模型里的 \gamma_i 虽然只在0和1之间浮动,但本质上和K-Means的簇归属概率是一样的思路。如果你熟悉K-Means,再回头看EM,会感觉两者有一种非常自然的递进关系。

5.5 为什么不能直接把对数求和的项拆开

我需要再强调一下: \log(a+b) 没办法直接拆成 \log a + \log b 。这是三硬币模型、高斯混合模型、甚至HMM中所有EM推导的共同障碍。EM之所以能“绕过去”,是因为它不直接处理观测数据对数似然,而是构建了一个包含了隐变量的完全数据对数似然,再对其期望做最大化。完全数据似然里没有加和障碍,所以每个参数都能独立估出来。

如果你在推导时卡住,一定要仔细检查自己是不是试图把 \log 拆进加号里了,这是初学者最容易犯的错。

6. 最后的实操心得

根据我自己的推导经验,三硬币模型表面上只需要三个参数,但步骤极为密集,稍不注意就会在E步或M步的某个符号上出错。我的建议是至少完整手算一次,哪怕只有5个样本,也能体会到每一行公式对应的实际意义。手算结束后再用代码复现,把代码里的每一步和手算子步骤一一对应起来。

代码复现时,先用固定种子生成模拟数据,再用已知真实参数的模拟结果来验证你实现是否正确。如果真实参数是0.3、0.8、0.5,而估出来接近0.5、0.4、0.9,那就要警惕是不是正反面搞反了,或者p与q、r的对应搞混了。

这算是我个人学EM算法时最重要的体会:不要试图一步到位理解所有原理,先盯住“软标签—加权平均—再软标签”这个循环,三硬币模型把这个循环压缩到了最简。弄懂了这一圈,高斯混合模型、隐马尔可夫模型里的EM算法再看就顺畅多了。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询