做机器学习的人几乎绕不开 EM 算法。高斯混合模型、隐马尔可夫模型、概率潜在语义分析,只要你跟“隐变量”打交道,EM 就会出现在你面前。但很多教材讲 EM,一上来就是琴生不等式、期望最大化、Q 函数推导,公式推完你只会背结论,根本不知道算法在算什么。三硬币模型恰好是打破这种困境的最佳切片:它把“隐变量”具象成“到底用了哪枚硬币”,把抽象的 Q 函数变成一组可以拿纸笔手算的加权平均。这篇文章我会从最大似然估计卡壳的地方讲起,完整拆解三硬币模型的 E 步和 M 步,再给你一份能直接跑的 Python 实现,最后把我在实操里踩过的初值坑、收敛坑、对称解坑全部摊开说。
1. 为什么学 EM 必须先看三硬币模型
1.1 从一个“缺了一半数据”的问题说起
最大似然估计大家都会:拿到一组观测样本,假设它们独立同分布地来自某个参数化分布,然后最大化对数似然。这个流程的关键前提是四个字——数据完整。
拿三硬币模型来说:有硬币 A、B、C,先抛 A 决定本轮用 B 还是 C,再抛选中的硬币记录正反面。如果给你的是完整数据,不仅告诉你每轮结果,还告诉你每轮用的是 B 还是 C,那么参数估计简单得令人发指:π 就是“用 B 的轮数占总轮数的比例”,p 就是“B 硬币的正面频率”,q 就是“C 硬币的正面频率”。直接数数,甚至不需要求导。
但现实往往是残缺的。你可能只拿到结果序列,至于每轮用的是 B 还是 C,这个信息彻底丢了。你需要估计三个参数,手里却只有一半信息。这就是隐变量的本质:它参与了数据生成,却未被观测到。
生活里的同类问题很常见。你拿到 100 份匿名问卷,每份答案是“是”或“否”,但问卷上没有填写人属于哪个小组。现在你要估计两个小组各自回答“是”的比例,以及样本来自第一组的比例——这就是三硬币模型的抽象。隐变量就是“这个人属于哪个组”,观测变量就是“他回答是还是否”。
1.2 观测似然为什么难啃
如果忽略隐变量,直接对观测数据的似然做最大似然估计,你会发现一个尴尬的问题:对数里面出现了求和号。
观测似然可以写成每个样本的贡献的乘积,而每个样本的贡献是“选 B 的概率乘 B 投出该结果的概率,加上选 C 的概率乘 C 投出该结果的概率”。这导致每项都是两项相加,取对数后就变成了 log(a+b)。log(a+b) 没法简洁地拆开,求导的时候每一项分母都同时带着 p 和 q 两个参数,互相纠缠,根本得不到闭式解。
这就是 EM 算法存在的动机。既然观测似然不好直接优化,完整数据似然又非常好解,那不如换个思路:我先根据当前参数,猜出缺失信息(隐变量)的分布,把观测数据“补全”,然后在完整数据的框架下更新参数。猜不能瞎猜,要用后验概率来猜,这是 E 步;猜完更新参数,这是 M 步。如此循环,直到收敛。
三硬币模型就是演示这套思路的最小舞台。它只有三个参数、一个二值隐变量,任何一步计算都能手算验算,非常适合作为 EM 的入门切片。
2. 三硬币模型:设定、符号与数据生成
2.1 问题设定与符号约定
先把模型参数约定清楚,后面所有推导都基于这套符号:
- 硬币 A:决定本轮用 B 还是 C,正面朝上的概率是 π
- 硬币 B:正面朝上的概率是 p
- 硬币 C:正面朝上的概率是 q
- 每一轮的数据生成过程:先抛 A,若正面则选 B,否则选 C;再抛所选硬币,记录结果 1 表示正面、0 表示反面;重复 n 轮,得到观测序列 y₁, y₂, …, yₙ
再定义隐变量:zᵢ = 1 表示第 i 轮选了 B,zᵢ = 0 表示第 i 轮选了 C。z 是我们永远看不到的缺失信息。
参数空间很简单:π、p、q 都在 [0, 1] 区间内。需要注意的是,观测变量的边缘分布并不是一个标准伯努利,而是两个伯努利的混合:
P(yᵢ = 1 | θ) = π·p + (1 - π)·q
这个式子本身就是理解三硬币模型的一把钥匙:任何一个观测结果,既可能来自 B,也可能来自 C,具体“归功于谁”取决于 π、p、q 的当前取值。
2.2 完整数据似然与观测似然:差不只一个求和号
如果 zᵢ 全部已知,完整数据似然长这样:
L_complete(θ) = ∏ᵢ [π·p^{yᵢ}(1-p)^{1-yᵢ}]^{zᵢ} · [(1-π)·q^{yᵢ}(1-q)^{1-yᵢ}]^{1-zᵢ}
这个函数对每个参数单独看,形式就是标准的伯努利似然,求导后能直接得到闭式解。也就是说,如果把隐变量当作已知,参数估计就是数数问题。
但如果 zᵢ 未知,观测似然必须把所有可能的隐变量取值全部求和:
L_obs(θ) = ∏ᵢ [π·p^{yᵢ}(1-p)^{1-yᵢ} + (1-π)·q^{yᵢ}(1-q)^{1-yᵢ}]
差别看似只是括号里多了一项,但就是因为这个求和号,直接求导这条路被堵死了。EM 算法的思路,就是建立一条“绕过 log-sum 障碍”的迭代路径:利用完整数据似然的结构,逐步逼近观测似然的最大值。
3. E 步与 M 步:手动推导一次你就懂
3.1 E 步:猜隐变量的后验概率
E 步的任务只有一个:在给定当前参数 θ⁽ᵗ⁾ 和观测 yᵢ 的条件下,计算 zᵢ = 1 的后验概率。根据贝叶斯公式直接写:
μᵢ⁽ᵗ⁾ = P(zᵢ = 1 | yᵢ, θ⁽ᵗ⁾) = π·p^{yᵢ}(1-p)^{1-yᵢ} / [π·p^{yᵢ}(1-p)^{1-yᵢ} + (1-π)·q^{yᵢ}(1-q)^{1-yᵢ}]
观察这个式子,你会发现 μᵢ 对不同的观测结果自然产生了不同的行为:
- 如果 yᵢ = 1,μᵢ 度量的是“看到正面时,这轮更像是 B 干的还是 C 干的”
- 如果 yᵢ = 0,μᵢ 度量的是“看到反面时,这轮更像是 B 干的还是 C 干的”
μᵢ 不是 0 就是 1 的硬判断,而是 0 到 1 之间的软分配。这是 EM 和 K-Means 这一类“硬聚类”算法的根本区别:EM 允许“这轮有 70% 的可能是 B,30% 的可能是 C”。
我经常用这么一句话概括 E 步:先按当前参数,把每轮观测的“功劳”按后验概率分给 B 和 C,再数数。数数之前要先分功劳,这就是 E 步存在的意义。
3.2 M 步:带着软计数更新参数
在完整数据下,三个参数的最大似然估计非常直观:
- π 的估计:选 B 的轮数比例,等于 Σzᵢ / n
- p 的估计:B 硬币的正面频率,等于 Σ zᵢyᵢ / Σzᵢ
- q 的估计:C 硬币的正面频率,等于 Σ(1-zᵢ)yᵢ / Σ(1-zᵢ)
现在 zᵢ 未知,但 E 步已经算出了 μᵢ,也就是 zᵢ 的后验期望。一个自然的替代操作是:用 μᵢ 代替 zᵢ,做软计数。于是 M 步的参数更新公式直接写为:
π_new = (1/n) Σ μᵢ
p_new = Σ μᵢyᵢ / Σ μᵢ
q_new = Σ (1-μᵢ)yᵢ / Σ (1-μᵢ)
这三个公式朴素得让人怀疑是不是搞错了。但如果从 Q 函数求导的角度看,它们的来路非常清晰。EM 每次迭代实际是在最大化一个替代目标 Q 函数:
Q(θ, θ⁽ᵗ⁾) = Σᵢ { μᵢ·log[π·p^{yᵢ}(1-p)^{1-yᵢ}] + (1-μᵢ)·log[(1-π)·q^{yᵢ}(1-q)^{1-yᵢ}] }
对 Q 关于 π 求偏导并令其为零:
∂Q/∂π = Σ μᵢ/π - Σ(1-μᵢ)/(1-π) = 0
整理后就是 π = Σμᵢ / n。对 p 和 q 分别求导,同样可以得到上面那两个加权频率公式。所以 M 步本质上就是在做带着后验权重的最大似然估计,每一步都有扎实的推导支撑。
3.3 手算一个最小例子
公式再多,不如亲手算一遍。设观测序列很短,只有 3 个样本:y = [1, 0, 1],初始参数设为 θ⁽⁰⁾ = (π=0.6, p=0.7, q=0.3)。我们手动迭代一轮。
先做 E 步。对每个样本计算分子和分母:
| 样本 i | yᵢ | 分子(选 B 的概率) | 分母(选 B 加选 C 的概率) | μᵢ |
|---|---|---|---|---|
| 1 | 1 | 0.6 × 0.7 = 0.42 | 0.42 + 0.4 × 0.3 = 0.54 | 0.7778 |
| 2 | 0 | 0.6 × 0.3 = 0.18 | 0.18 + 0.4 × 0.7 = 0.46 | 0.3913 |
| 3 | 1 | 0.6 × 0.7 = 0.42 | 0.42 + 0.4 × 0.3 = 0.54 | 0.7778 |
然后做 M 步。先算 Σμᵢ = 0.7778 + 0.3913 + 0.7778 = 1.9469:
π_new = 1.9469 / 3 ≈ 0.6490
p_new = (0.7778×1 + 0.3913×0 + 0.7778×1) / 1.9469 ≈ 1.5556 / 1.9469 ≈ 0.7990
q_new = (0.2222×1 + 0.6087×0 + 0.2222×1) / (0.2222 + 0.6087 + 0.2222) ≈ 0.4444 / 1.0531 ≈ 0.4220
一轮迭代下来,p 从 0.7 提升到 0.799,q 从 0.3 提升到 0.422,π 从 0.6 提升到 0.649。方向很好理解:3 个观测里有 2 个正面,模型正在调整参数去解释这些正面观测。唯一需要提醒的是,因为样本量只有 3,信息量太少,算法会把两个硬币的参数往整体正面频率方向拉,这是正常现象。这个手算过程,强烈建议你拿纸笔完整走一遍,胜过干看十遍公式。
4. 用 Python 完整实现 EM 算法
4.1 代码结构设计
手算验证了公式,接下来用代码验证算法。实现 EM 不需要调用任何高级库,numpy 足够。代码分三块:生成模拟数据、E 步、M 步,外面套一层迭代循环。模拟数据的价值在于:我们知道真实参数,可以对照 EM 的恢复效果。
import numpy as np def generate_data(pi_true, p_true, q_true, n=1000, seed=42): np.random.seed(seed) z = np.random.binomial(1, pi_true, size=n) # 每轮到底选了 B 还是 C probs = np.where(z == 1, p_true, q_true) # 所选硬币的正面概率 y = np.random.binomial(1, probs) # 最终观测结果 return y def e_step(y, pi, p, q): # 分子:本轮选 B 且得到结果 y 的概率 num_b = pi * (p ** y) * ((1 - p) ** (1 - y)) # 分母:还要加上本轮选 C 的概率 denom = num_b + (1 - pi) * (q ** y) * ((1 - q) ** (1 - y)) return num_b / denom def m_step(y, mu): pi_new = mu.mean() p_new = np.dot(mu, y) / mu.sum() q_new = np.dot(1 - mu, y) / (1 - mu).sum() return pi_new, p_new, q_new def log_likelihood(y, pi, p, q): term = pi * (p ** y) * ((1 - p) ** (1 - y)) + (1 - pi) * (q ** y) * ((1 - q) ** (1 - y)) return np.sum(np.log(term)) def em_algorithm(y, init, max_iter=500, tol=1e-8): pi, p, q = init ll_prev = -np.inf for i in range(max_iter): mu = e_step(y, pi, p, q) pi, p, q = m_step(y, mu) ll = log_likelihood(y, pi, p, q) if abs(ll - ll_prev) < tol: return pi, p, q, i + 1 ll_prev = ll return pi, p, q, max_iter这段代码只有约 30 行,结构非常透明。E 步对应 e_step 函数,M 步对应 m_step 函数,循环里先算 μ 再更新参数。需要注意的一点是:收敛判断用的是对数似然的绝对变化量,而不是参数变化量。原因是 EM 的目标是最大化观测似然,直接监控目标函数才是最可靠的停止标准。
4.2 用模拟数据跑一遍
设计一个真实参数环境:π = 0.4,p = 0.8,q = 0.3,生成 1000 次观测。从初值 (π=0.2, p=0.6, q=0.4) 出发,跑 EM 算法。我在一次运行中得到的大致轨迹如下:
| 迭代轮次 | π | p | q | 对数似然 |
|---|---|---|---|---|
| 初始值 | 0.200 | 0.600 | 0.400 | -697.34 |
| 第 5 轮 | 0.312 | 0.688 | 0.348 | -695.87 |
| 第 10 轮 | 0.365 | 0.748 | 0.325 | -694.56 |
| 第 20 轮 | 0.394 | 0.786 | 0.308 | -693.98 |
| 收敛后 | 0.401 | 0.798 | 0.302 | -693.72 |
收敛后得到的参数和真实值 (0.4, 0.8, 0.3) 非常接近,说明算法正确恢复出了生成数据的机制。你可以看到对数似然每一轮都在上升(至少不下降),这是 EM 算法的理论保证:每次迭代都让观测似然单调不减。这个性质在优化领域算是非常友好的,你几乎不用担心发散问题。
两个工程细节值得留意。第一,μ 的计算中分子分母都是概率乘积,三硬币模型只有一次伯努利观测,不会下溢;但扩展高斯混合模型时,连续密度值的乘积会非常小,需要改用 log-sum-exp 技巧保证数值稳定。第二,M 步中分母 Σμᵢ 理论上不会为 0,但实践中如果所有 μ 都极端接近 0,除零还是会炸,稳妥的做法是给分母加一个 1e-12 量级的 epsilon。
4.3 初值会带你到不同的山头上
同一批数据,换一个初值,结果可能完全不同。这是 EM 最需要警惕的特性。我做了一个对比实验,全部使用相同的数据,仅改变初始参数:
第一组,初值取 (π=0.5, p=0.5, q=0.5)。E 步算出来的所有 μᵢ 都等于 0.5,M 步得到 π_new = 0.5,p_new = q_new = 观测正面频率(大约 0.48)。下一轮依然如此,算法直接停在了一个对称解上:π = 0.5,p = q ≈ 0.48,对数似然约 -693.7。这个解和真实参数下的似然几乎一样,因为模型无法区分“用 B 的概率和用 C 的概率各占一半,且两枚硬币参数相同”与真实的混合结构。
第二组,初值取 (π=0.8, p=0.4, q=0.2)。这次 EM 收敛到 π ≈ 0.60,p ≈ 0.29,q ≈ 0.79。看起来和真实参数完全不同,但它其实是真实解的镜像:真实是 (0.4, 0.8, 0.3),镜像把 B 和 C 的角色对调,再把 π 取补数,得到 (0.6, 0.3, 0.8)。这两个参数化方式生成的是完全相同的观测分布,对数似然几乎没差别。
这个实验说明两件事:第一,EM 只能保证找到局部最优,不保证全局最优;第二,三硬币模型本身具有参数不可辨识性——p 和 q 交换、π 变为 1-π 后,观测分布不变。所以你跑 EM 时不要只看最终参数,还要同时看对数似然值,必要的话用多组初值并行跑,取似然最高的结果。
5. 实践中的坑与排查技巧
5.1 初值敏感与多峰问题
初值敏感是 EM 最出名的问题。三硬币模型还算温和,到了高斯混合模型,不同初值可能收敛到完全不同质量的聚类结果。我的习惯做法是:多组初值并行跑。比如随机生成 10 组初值,每组跑 50 次迭代,先看对数似然的快速排名,取前几名精细收敛。这样能显著降低落到坏局部最优的概率。
还有一个常用的启发式初始化:先用简单的聚类或分组方法给数据打上粗略标签,再由这些标签估计初始参数。比如三硬币模型里,可以先按结果是正面还是反面把样本粗略分成两组,用组内正面比例作为 p 和 q 的初值。这种初始化虽然不是最优,但往往比纯随机更接近真实解。
5.2 收敛判断与阈值选择
收敛判断有几种常见做法:监控对数似然变化量、监控参数变化量、或者两者同时监控。我推荐监控对数似然,因为它直接指向优化目标。阈值方面,三硬币模型这种低维问题可以设 1e-8,高斯混合模型这类复杂模型设 1e-6 就够了,太严格的阈值会导致大量无意义的迭代。
另一个经验:迭代上限一定要设置。EM 在接近收敛时是线性收敛,速度可能很慢,如果不设上限,遇到病态数据可能跑几千轮。我一般设 500 次上限,同时打印每 20 轮的轨迹,肉眼确认参数是否在缓慢漂移。
5.3 参数不可辨识性
三硬币模型里,交换 B 和 C 的角色、π 替换为 1-π,观测分布完全不变。这不是算法的 bug,是模型结构的固有冗余。如果你的目标只是拟合数据,这个特性无所谓;但如果你想给参数赋予明确解释,比如“硬币 B 的正面概率到底是多少”,就必须引入约束来打破对称性。常见做法是固定 π < 0.5,或者初始化时人为让 p > q,使算法始终收敛到约定的分支。
5.4 E 步数值下溢
虽然三硬币模型不会遇到这个问题,但所有 EM 实现都应该提前养成好习惯。算 μᵢ 时分子分母都是多个概率的乘积,样本量大或维度高时很小的概率乘起来会下溢成 0。解决思路是在 log 空间计算,最后再 exp 回去,或者使用 log-sum-exp 技巧。这个坑在高斯混合模型里几乎必踩,提前了解能省很多排查时间。
6. 下一步:从三硬币走向 GMM 与更广的 EM
6.1 三硬币模型就是最简混合模型
把三硬币模型稍微变一下,就是高斯混合模型:观测从二值的 0/1 变成连续实数,硬币 B 和 C 对应两个高斯分布,π 对应混合系数,p 和 q 对应两个高斯成分的均值和方差。E 步算的 μᵢ 在 GMM 里叫“responsibility”,意思是“第 i 个样本由第 k 个成分生成的责任”。M 步的更新公式同样是加权统计:新均值是样本按 responsibility 加权平均,新方差是加权平方残差。你会发现,GMM 的 EM 和三硬币模型的 EM 几乎是同一套骨架,只是把“数硬币正反面”换成了“算高斯密度”。
6.2 K-Means 是硬 EM
GMM 和 K-Means 的亲戚关系也值得理解一遍。当 GMM 中每个成分的方差趋近于 0,后验概率 μᵢₖ 会逐渐变成 one-hot 向量:每个样本被唯一分配给概率最大的那个成分。在这个极限下,E 步变成了“按距离最近分配”,M 步变成了“重新计算簇中心”,这就是 K-Means。所以你可以把 K-Means 理解成一种“硬 EM”:它不做软分配,直接用最硬的分组更新参数。理解了这层关系,再回头看为什么 K-Means 对初始簇中心敏感、为什么可能收敛到局部最优,一切都是顺理成章的。
6.3 扩展方向
三硬币模型学完之后,继续前进的路线很清晰:先啃一维 GMM 的完整推导和实现,再学多维 GMM,之后可以进入隐马尔可夫模型——HMM 的 E 步不再是简单的贝叶斯公式,而是需要前向-后向算法在整个隐状态序列上做平滑。再往后,潜在狄利克雷分配这类主题模型里,后验分布过于复杂,精确 EM 做不到,于是有变分推断来近似 E 步。这些方向虽然越来越复杂,但核心思想都没有离开你现在学到的这套“猜隐变量、软计数、更新参数”的循环。
我自己带新人学 EM 的时候,从来不让大家一上来就啃 GMM 的公式推导。三硬币模型这个案例,我至少讲了几十遍,每次讲完都有人一拍大腿:原来 Q 函数就是干这个的。如果你正在学 EM,我的建议很直接:把上面的代码亲手敲一遍,把书上的手算例子也过一遍,然后重点试那个初值 (0.5, 0.5, 0.5) 的对称陷阱,看看算法如何停在看似合理却完全错误的解上。这个坑只要亲手踩过一次,你就比对着公式背十遍的人理解得深得多。