Skip to content

7 · 혼합 모델과 EM

관측 변수와 잠재 변수의 결합 분포를 정의하면, 복잡한 관측 변수의 주변 분포를 상대적으로 더 다루기 쉬운 확장 공간(관측+잠재)의 결합 분포로 표현할 수 있다.

7.1 K 평균 집단화

  • D차원 데이터 {x1,,xN}K개 집단으로 나누는 문제에서, 각 집단의 원형prototype(중심) μk와 이진 표시 변수 rnk{0,1}(원 핫)을 도입한다. 목표는 각 점에서 가장 가까운 중심까지의 거리 제곱합인 뒤틀림 척도distortion measure를 최소화하는 것이다.
J=n=1Nk=1Krnkxnμk2
  • J를 두 단계로 번갈아 최소화한다. E 단계에서 μk를 고정하고 각 점을 가장 가까운 중심에 할당하고(rnk=1 if k=argminjxnμj2), M 단계에서 rnk를 고정하고 중심을 갱신한다.
μk=nrnkxnnrnk

즉 집단 k에 할당된 점들의 평균이다. 이를 K-평균K-means 알고리즘이라 하며, 제곱 유클리드 거리 대신 일반 거리 V(,)를 쓰면 K-메도이드K-medoids가 된다.

그림 7.1 · (a) 데이터와 두 초기 중심(적·청 ×). (b) E 단계: 각 점을 가까운 중심에 할당. (c) M 단계: 중심 재계산. (d)~(i) 수렴까지 EM 반복.

실습 · K-평균 집단화

E 단계(가까운 중심에 할당)와 M 단계(중심을 집단 평균으로 갱신)를 번갈아 반복하며, 뒤틀림 척도 J가 단조 감소해 수렴하는 것을 확인합니다.

python
import numpy as np
rng = np.random.default_rng(0)
X = np.vstack([rng.normal([0,0], 0.6, (80,2)),
               rng.normal([3,3], 0.6, (80,2)),
               rng.normal([0,3], 0.6, (80,2))])
K = 3; mu = X[rng.choice(len(X), K, replace=False)].copy()
prev = None
for it in range(12):
    d = ((X[:,None,:]-mu[None,:,:])**2).sum(2); r = d.argmin(1)     # E: 할당
    J = d[np.arange(len(X)), r].sum()
    for k in range(K):
        if (r==k).any(): mu[k] = X[r==k].mean(0)                    # M: 중심 갱신
    print(f"반복 {it+1}: J={J:.2f}")
    if prev is not None and abs(J-prev) < 1e-9: break
    prev = J

7.2 혼합 가우시안

  • 가우시안 혼합은 가우시안들의 선형 중첩이다.
(9.7)p(x)=k=1KπkN(xμk,Σk)
  • 원 핫 잠재 변수 z(zk{0,1}, kzk=1)를 도입해 p(zk=1)=πk, 즉 p(z)=kπkzk로 두고, 조건부 분포를 p(xz)=kN(xμk,Σk)zk로 두면, z를 주변화해 (9.7)을 얻는다. x가 주어졌을 때 성분 k의 사후 확률(책임responsibility)은
γ(zk)p(zk=1x)=πkN(xμk,Σk)j=1KπjN(xμj,Σj)

그림 7.2 · (a) 결합 분포 p(z)p(xz)의 표본. (b) 주변 분포 p(x)의 표본. (c) 책임값 γ(znk)로 색칠한 점들.

7.2.1 최대 가능도 방법

  • 로그 가능도는 다음과 같다.
(9.14)lnp(Xπ,μ,Σ)=n=1Nln{k=1KπkN(xnμk,Σk)}
  • 최대 가능도는 특이점singularity 문제를 갖는다. 어떤 성분이 한 데이터 포인트로 붕괴하면 σj0이 되어 로그 가능도가 무한대로 발산한다.

그림 7.3·7.4 · (왼쪽) N개 데이터에 대한 가우시안 혼합의 그래프 표현({zn}은 잠재). (오른쪽) 한 성분이 데이터 포인트로 붕괴해 특이점이 생기는 모습.

7.2.2 가우시안 혼합 분포에 대한 EM

  • (9.14)를 μk, Σk, πk에 대해 미분해 0으로 두면(마지막은 kπk=1 제약에 라그랑주 승수) 다음을 얻는다(Nk=nγ(znk)는 성분 k의 유효 데이터 수).
(9.17)μk=1Nkn=1Nγ(znk)xn(9.19)Σk=1Nkn=1Nγ(znk)(xnμk)(xnμk),πk=NkN
  • 이 식들은 책임값 γ(znk)가 매개변수에 복잡하게 의존하므로 닫힌 형태의 해가 아니다. 대신 E 단계(현재 매개변수로 책임값 계산)와 M 단계(책임값으로 매개변수 재추정)를 번갈아 반복한다. 보통 K-평균으로 초기화한 뒤 EM을 적용한다.

가우시안 혼합 분포에 대한 EM 알고리즘

  1. 평균 μk, 공분산 Σk, 혼합 계수 πk를 초기화하고 로그 가능도를 계산한다.

  2. E 단계 — 책임값을 계산한다.

γ(znk)=πkN(xnμk,Σk)j=1KπjN(xnμj,Σj)
  1. M 단계 — 매개변수를 재추정한다(Nk=nγ(znk)).
μknew=1Nknγ(znk)xn,Σknew=1Nknγ(znk)(xnμknew)(xnμknew),πknew=NkN
  1. 로그 가능도를 계산하고, 수렴하지 않았으면 2단계로 돌아간다.

그림 7.5 · (a) 초기 두 성분. (b) 최초 E 단계 후 책임값으로 색칠. (c) 최초 M 단계 후 평균·공분산 갱신. (d)~(f) 2·5·20단계 후, (f)에서 거의 수렴.

실습 · 가우시안 혼합 EM

E 단계(책임값 γ 계산)와 M 단계(μ,Σ,π 재추정)를 반복하며 로그 가능도가 단조 증가해 수렴하고, 참 평균·혼합 계수를 회복하는 것을 확인합니다.

python
import numpy as np
def gauss(X, mu, S):
    D = X.shape[1]; d = X-mu; Si = np.linalg.inv(S)
    return np.exp(-0.5*np.einsum('ni,ij,nj->n', d, Si, d))/np.sqrt((2*np.pi)**D*np.linalg.det(S))
rng = np.random.default_rng(1)
X = np.vstack([rng.multivariate_normal([0,0], [[0.5,0],[0,0.5]], 120),
               rng.multivariate_normal([3,3], [[0.6,0.3],[0.3,0.6]], 120)])
K = 2; N = len(X)
mu = X[rng.choice(N, K, replace=False)].copy(); S = [np.eye(2) for _ in range(K)]; pi = np.ones(K)/K
prev = -np.inf
for it in range(40):
    R = np.stack([pi[k]*gauss(X, mu[k], S[k]) for k in range(K)], 1)     # E 단계
    ll = np.log(R.sum(1)).sum(); R = R/R.sum(1, keepdims=True)
    Nk = R.sum(0)
    for k in range(K):                                                    # M 단계
        mu[k] = (R[:,k,None]*X).sum(0)/Nk[k]
        d = X-mu[k]; S[k] = (R[:,k,None,None]*np.einsum('ni,nj->nij', d, d)).sum(0)/Nk[k]
    pi = Nk/N
    if it < 4 or it % 6 == 0: print(f"반복 {it+1:2d}: 로그가능도={ll:.2f}")
    if ll-prev < 1e-6: print(f"수렴 (반복 {it+1})"); break
    prev = ll
print("추정 π:", np.round(pi,3).tolist(), " 평균:", np.round(mu,2).tolist())

7.3 EM에 대한 다른 관점

  • 관측 데이터 X, 잠재 변수 Z, 매개변수 θ에 대해 로그 가능도는 lnp(Xθ)=ln{Zp(X,Zθ)}이다. {X,Z}완전한complete 데이터, X만을 불완전한incomplete 데이터라 한다. 잠재 변수에 대한 지식은 사후 분포 p(ZX,θ)로만 주어진다.

일반적인 EM 알고리즘

결합 분포 p(X,Zθ)에서 가능도 p(Xθ)θ에 대해 최대화한다.

  1. θold를 초기화한다.

  2. E 단계 — 사후 분포 p(ZX,θold)를 계산한다.

  3. M 단계 — 완전 데이터 로그 가능도의 기댓값을 최대화한다.

θnew=argmaxθ Q(θ,θold),Q(θ,θold)=Zp(ZX,θold)lnp(X,Zθ)
  1. 수렴하지 않으면 θoldθnew로 두고 2단계로 돌아간다.

7.3.1 베르누이 분포들의 혼합

  • 이진 변수들에 대한 베르누이 혼합(잠재 클래스 분석latent class analysis)을 생각하자. 각 성분은 p(xμk)=i=1Dμkixi(1μki)1xi이고 혼합은 p(xμ,π)=kπkp(xμk)이다. 책임값과 M 단계는
γ(znk)=πkp(xnμk)j=1Kπjp(xnμj),μk=xk=1Nknγ(znk)xn,πk=NkN

가우시안 혼합과 달리 μki(0,1)이라 특이점이 생기지 않는다.

실습 · 베르누이 혼합 EM (이진 데이터 군집화)

두 이진 원형 패턴에서 생성한 잡음 섞인 데이터를 베르누이 혼합 EM으로 군집화하면, 각 성분의 μk가 원래 패턴을 회복하는 것을 확인합니다.

python
import numpy as np
rng = np.random.default_rng(2)
proto = np.array([[1,1,1,0,0,0], [0,0,0,1,1,1]], float)     # 두 원형 패턴
X = np.vstack([(rng.uniform(size=(60,6)) < 0.85*proto[0]+0.05).astype(float),
               (rng.uniform(size=(60,6)) < 0.85*proto[1]+0.05).astype(float)])
K = 2; N, D = X.shape
mu = rng.uniform(0.25, 0.75, (K, D)); pi = np.ones(K)/K; prev = -np.inf
def bern(X, mu):                                             # p(x|μ_k), (N,K)
    return np.prod(mu**X[:,None,:]*(1-mu)**(1-X[:,None,:]), 2)
for it in range(60):
    P = pi*bern(X, mu); ll = np.log(P.sum(1)).sum(); R = P/P.sum(1, keepdims=True)   # E
    Nk = R.sum(0); mu = (R.T@X)/Nk[:,None]; pi = Nk/N                                # M
    if ll-prev < 1e-7: break
    prev = ll
print(f"수렴 (반복 {it+1}), 로그가능도={ll:.2f}")
print("성분1 μ:", np.round(mu[0], 2).tolist())
print("성분2 μ:", np.round(mu[1], 2).tolist())

7.3.2 베이지안 선형 회귀에 대한 EM

  • 베이지안 선형 회귀에서 가중치 w를 잠재 변수로 보면 초매개변수 α, β를 EM으로 추정할 수 있다. w의 사후 분포에 대한 완전 데이터 로그 가능도의 기댓값을 α에 대해 최대화하면
α=ME[ww]=MmNmN+Tr(SN)

mN·SN은 사후 분포의 평균·공분산이며, β도 비슷하게 재추정한다.

7.4 일반적 EM 알고리즘

  • 잠재 변수 분포 q(Z)를 도입하면 로그 가능도가 다음처럼 분해된다.
(9.70)lnp(Xθ)=L(q,θ)+KL(qp)L(q,θ)=Zq(Z)lnp(X,Zθ)q(Z),KL(qp)=Zq(Z)lnp(ZX,θ)q(Z)
  • KL(qp)0이므로 L(q,θ)는 로그 가능도의 하한이다. EM은 이 하한을 두 단계로 올린다. E 단계에서 θold를 고정하고 q(Z)=p(ZX,θold)로 두면 KL=0이 되어 하한이 로그 가능도와 같아진다. M 단계에서 q를 고정하고 L(q,θ)θ에 대해 최대화하면, 하한이 오르고 로그 가능도도 최소한 그만큼 오른다(이때 q가 새 사후 분포와 달라져 KL>0).

그림 7.6~7.8 · (왼쪽) lnp(X)=L(q,θ)+KL(qp) 분해. (가운데) E 단계: q를 사후 분포로 두어 KL=0. (오른쪽) M 단계: q 고정, θ 최대화로 하한 상승.

실습 · ELBO 분해 lnp(X)=L(q,θ)+KL(qp)

가우시안 혼합의 한 스냅샷에서, 임의의 q에 대해 L+KL이 로그 가능도와 같음을, 그리고 q를 사후 분포로 두면 KL=0이 되어 L이 로그 가능도와 일치함을 확인합니다.

python
import numpy as np
def gauss(X, mu, S):
    D = X.shape[1]; d = X-mu; Si = np.linalg.inv(S)
    return np.exp(-0.5*np.einsum('ni,ij,nj->n', d, Si, d))/np.sqrt((2*np.pi)**D*np.linalg.det(S))
rng = np.random.default_rng(3)
X = np.vstack([rng.multivariate_normal([0,0], np.eye(2)*0.4, 50),
               rng.multivariate_normal([2.5,2.5], np.eye(2)*0.4, 50)])
K = 2; mu = np.array([[0.5,0.0],[2.0,2.0]]); S = [np.eye(2)*0.5]*2; pi = np.array([0.5,0.5])
px = np.stack([pi[k]*gauss(X, mu[k], S[k]) for k in range(K)], 1)
lnpX = np.log(px.sum(1)).sum()
post = px/px.sum(1, keepdims=True)                          # 사후 분포
def elbo(q):
    L = KL = 0.0
    for k in range(K):
        m = q[:,k] > 1e-12
        L  += np.sum(q[m,k]*(np.log(px[m,k]) - np.log(q[m,k])))
        KL += np.sum(q[m,k]*(np.log(q[m,k]) - np.log(post[m,k])))
    return L, KL
Lb, KLb = elbo(np.full_like(post, 0.5))                     # 임의의 q
Lg, KLg = elbo(post)                                        # q = 사후 분포
print(f"ln p(X)           = {lnpX:.3f}")
print(f"임의 q : L={Lb:.3f}, KL={KLb:.3f}, L+KL={Lb+KLb:.3f}")
print(f"q=사후 : L={Lg:.3f}, KL={KLg:.3f}, L+KL={Lg+KLg:.3f}  (KL≈0 → L=ln p(X))")

연습문제

문제 7.1

잠재 변수의 주변 분포가 (9.10) p(z)=kπkzk, 조건부 분포가 (9.11) p(xz)=kN(xμk,Σk)zk일 때, zp(z)p(xz)가 (9.7)의 가우시안 혼합임을 보여라.

풀이

z는 원 핫이므로 정확히 한 성분만 1이고, 가능한 상태는 z=ek(k번째만 1), k=1,,KK가지다. 각 상태에서

p(z=ek)=j=1Kπjzj=πk,p(xz=ek)=j=1KN(xμj,Σj)zj=N(xμk,Σk)

(zj=0인 항은 πj0=1, N0=1이라 사라진다). 모든 상태에 대해 합하면

zp(z)p(xz)=k=1Kp(z=ek)p(xz=ek)=k=1KπkN(xμk,Σk)

이는 정확히 (9.7)의 가우시안 혼합이다.

문제 7.2

혼합 밀도 p(x)=k=1Kπkp(xk)에서 x=(xa,xb)로 나눌 때, 조건부 밀도 p(xbxa)가 그 자체로 혼합 분포임을 보이고 혼합 계수와 성분 밀도를 구하라.

풀이

조건부 분포의 정의에서

p(xbxa)=p(xa,xb)p(xa)=k=1Kπkp(xa,xbk)j=1Kπjp(xaj)

성분 내에서 p(xa,xbk)=p(xbxa,k)p(xak)로 분해하면

p(xbxa)=k=1Kπkp(xak)j=1Kπjp(xaj)λkp(xbxa,k)

따라서 p(xbxa)는 성분 밀도 p(xbxa,k)와 혼합 계수

λk=πkp(xak)j=1Kπjp(xaj)

를 갖는 혼합 분포다. λk0이고 kλk=1이므로 유효한 혼합 계수이며, 이는 xa를 관측한 뒤의 성분 사후 확률에 해당한다.

PDF