12장 — 평균장 근사

변분 추론: 인자를 하나씩 바꾸는 규칙

그런데 tanh 조건은 변수가 ±1인 스핀이라서 나온 모양이다. ML에서 근사하려는 변수는 VAE의 잠재변수처럼 대개 연속인 값이고, 사후분포는 데이터가 들어올 때마다 새로 근사해야 한다. 변수가 연속인 값이거나 여러 값을 가질 때, 가장 나은 곱 분포를 이루는 변수별 분포 qᵢ 하나(인자, factor)는 어떤 규칙을 따를까? 그리고 그 규칙을 ML에서 계산할 수 없는 분포, 예를 들어 VAE의 사후분포에도 쓸 수 있을까?

역사: 50년 걸릴 진단을 1초 안에

1990년대에 들어 통계학과 인공지능에서는 변수 사이의 관계를 그래프로 그린 확률 모형이 널리 쓰였다. 정확한 추론 알고리즘은 그래프에서 서로 얽힌 변수의 묶음이 작을 때만 빠르다. 의학 진단용으로 만든 QMR-DT 네트워크가 그 한계를 보여 주었다. 질병 약 600가지와 증상·검사 소견 약 4000가지를 잇고, 환자에게서 나온 소견으로 어떤 병이 있을지를 계산하는 모형이다. 소견 하나가 여러 질병과 이어져 있어 그래프가 촘촘했고, 이 모형에 맞춰 만든 정확한 알고리즘조차 계산 시간이 양성 소견의 수에 대해 지수적으로 늘었다. 1999년 야콜라와 조던은 어려운 임상 사례 묶음에서 이 알고리즘이 사례 하나에 평균 50년쯤 걸릴 것이라고 어림했다. 샘플링으로 근사하는 방법은 그 묶음 가운데 두 사례에서만 쓸 만한 시간 안에 답을 냈다.

두 사람은 계산하기 어려운 확률을 계산할 수 있는 식의 위아래 한계로 바꾸는 변분 근사를 이 모형에 맞춰 만들었다. 소견을 몇 개만 정확히 다루고 나머지는 근사로 처리하자, 정확한 답을 낼 수 있었던 네 사례에서 정확한 방법이 평균 26.9초 걸린 계산을 0.11초에 해냈고, 정확한 방법으로는 풀 수 없던 사례들에도 답을 냈다. 같은 해 조던, 가라마니, 야콜라, 솔은 그래프 모형을 위한 변분 방법을 입문 논문으로 정리했다. 2017년 블라이, 쿠쿠켈비르, 매컬리프는 통계학자를 위한 리뷰에서 이 흐름의 출발점으로 볼츠만 머신에 평균장을 쓴 피터슨과 앤더슨의 1987년 논문을 꼽으며 「특정한 모형, 곧 신경망에 대한 첫 변분 절차라 할 만하다」고 적었다.

인자 하나의 갱신 규칙

규칙을 찾기 전에 스핀 둘에서 인자 하나만 바꾸는 일을 숫자로 해 보자. βJ = 1, βb = 0.5에서 동전 2를 m₂ = 0.5에 묶어 두면, 동전 1을 어떻게 고르든 결합 에너지의 평균은 −J × 0.5 × s₁이 된다. 그러면 동전 1은 편향 b + 0.5J를 받는 홀로 선 스핀이고, F[q]를 가장 낮추는 동전 1은 볼츠만 분포 그대로 m₁ = tanh(0.5 + 0.5) = 0.762다. 상대를 평균에 묶어 두고 내 쪽 분포를 「묶어 둔 상대로 평균 낸 에너지의 볼츠만 분포」로 바꾸는 이 한 걸음을 일반 변수로 적은 것이 아래 규칙이다.

변수를 h = (h₁, …, h_N)이라 하자. 물리에서는 스핀이나 입자의 위치, ML에서는 잠재변수다. ML 문헌은 잠재변수를 흔히 z로 적지만, 이 책에서 z는 로짓 자리의 값이라 h로 적는다. 글자만 바꾼 것이다. 목표 분포 p(h) = e^(−βE(h))/Z에 대해 곱 분포 q(h) = Πqᵢ(hᵢ)를 두자. Π(대문자 파이)는 Σ의 곱셈판으로, 아래에 붙은 범위의 인자를 모두 곱하라는 기호다. 나머지 인자를 고정한 채 F[q]를 인자 qᵢ 하나에 대해 최소화하면 다음 규칙이 나온다.

ln⁡qi(hi)=⟨ln⁡p(h)⟩q−i+상수=−β ⟨E(h)⟩q−i+상수\ln \textcolor{#bcbd22}{q_i}(\textcolor{#1b9e77}{h_i}) = \big\langle \ln \textcolor{#e377c2}{p}(\textcolor{#1b9e77}{h}) \big\rangle_{\textcolor{#bcbd22}{q_{-i}}} + \text{상수} = -\textcolor{#8c564b}{\beta}\,\big\langle \textcolor{#ff7f0e}{E}(\textcolor{#1b9e77}{h}) \big\rangle_{\textcolor{#bcbd22}{q_{-i}}} + \text{상수}
qi(hi)변수 i의 인자 (나머지 인자를 고정했을 때 가장 나은 것)q−ii를 뺀 나머지 인자들의 곱 (이것으로 평균을 낸다)p(h)목표 분포 (ML에서는 사후분포)E(h)에너지 (ML에서는 관측값과 잠재변수의 결합 분포에 −ln을 씌운 것)상수hᵢ와 무관한 값 (qᵢ의 합이 1이 되게 맞추면 정해진다)h=(h1,…,hN)변수 전체β역온도 1/kT (ML에서는 1)\begin{array}{ll} \textcolor{#bcbd22}{q_i}(\textcolor{#1b9e77}{h_i}) & \text{변수 i의 인자 (나머지 인자를 고정했을 때 가장 나은 것)} \\ \textcolor{#bcbd22}{q_{-i}} & \text{i를 뺀 나머지 인자들의 곱 (이것으로 평균을 낸다)} \\ \textcolor{#e377c2}{p}(\textcolor{#1b9e77}{h}) & \text{목표 분포 (ML에서는 사후분포)} \\ \textcolor{#ff7f0e}{E}(\textcolor{#1b9e77}{h}) & \text{에너지 (ML에서는 관측값과 잠재변수의 결합 분포에 −ln을 씌운 것)} \\ \text{상수} & \text{hᵢ와 무관한 값 (qᵢ의 합이 1이 되게 맞추면 정해진다)} \\ \textcolor{#1b9e77}{h} = (\textcolor{#1b9e77}{h_1}, \dots, \textcolor{#1b9e77}{h_N}) & \text{변수 전체} \\ \textcolor{#8c564b}{\beta} & \text{역온도 1/kT (ML에서는 1)} \end{array}

이징 모형에서는 나머지 스핀으로 평균 낸 에너지가 sᵢ의 일차식 −(bᵢ + ΣⱼJᵢⱼmⱼ)sᵢ + (sᵢ와 무관한 값)이므로, 이 규칙이 곧 tanh 조건이다. 인자 하나를 이렇게 바꾸는 것은 그 인자에 대한 정확한 최소화라서 F[q]는 한 번 바꿀 때마다 줄거나 그대로이고, 인자들을 하나씩 돌아가며 바꾸면 F[q]가 멈출 때까지 내려간다.

이 규칙에서 목표 분포 p가 볼츠만 분포일 필요는 없었다. 통계학과 ML에서는 목표 분포가 관측값 x가 주어진 잠재변수 h의 사후분포 p(h | x)이고, 에너지는 온도 1의 −ln p(x, h)이다. 계산할 수 있는 분포의 모임을 정하고 그 안에서 D_KL(q‖p)를 가장 작게, 곧 ELBO = −F[q]를 가장 크게 만드는 q로 사후분포를 대신하는 것을 변분 추론 (근사 분포의 모임 안에서 ELBO를 최대화해 사후분포를 대신하는 방법, variational inference)이라 한다. 모임을 곱 분포로 잡은 것이 평균장 변분 추론이고, 위의 규칙을 인자마다 돌아가며 적용하는 알고리즘을 좌표 상승 변분 추론(coordinate ascent variational inference, CAVI)이라 부른다. 물리의 평균장 근사와 ML의 평균장 변분 추론은 같은 최소화를 서로 다른 이름으로 부르는 것이다.

로그우도는 ELBO와 KL 두 토막으로 나뉜다. 결합된 스핀 둘(βJ = 1, βb = 0.5)에서 ln Z = 2.211은 가장 나은 곱 분포의 ELBO, 곧 −βF[q] = 2.108과 틈 D_KL(q‖p) = 0.103의 합이다. ML에서는 ln Z 자리에 데이터의 로그우도 ln p(x)가 온다. ELBO를 키우는 것은 로그우도가 정해져 있으니 KL을 줄이는 것과 같다.
로그우도는 ELBO와 KL 두 토막으로 나뉜다. 결합된 스핀 둘(βJ = 1, βb = 0.5)에서 ln Z = 2.211은 가장 나은 곱 분포의 ELBO, 곧 −βF[q] = 2.108과 틈 D_KL(q‖p) = 0.103의 합이다. ML에서는 ln Z 자리에 데이터의 로그우도 ln p(x)가 온다. ELBO를 키우는 것은 로그우도가 정해져 있으니 KL을 줄이는 것과 같다.

코드: 상관된 정규분포를 곱 분포로 근사하기

이 장 첫머리의 사후분포(평균 (1, −1), 분산 1, 상관계수 0.9인 정규분포)를 좌표 상승 변분 추론으로 근사한다. 정규분포의 곱 분포에서는 규칙 ln qᵢ = ⟨ln p⟩_(나머지) + 상수가 닫힌 식(closed form: 되풀이 계산 없이 공식 한 줄로 나오는 답)으로 풀려서, 각 인자는 평균이 상대 인자의 평균에 따라 움직이고 분산은 정밀도 행렬(공분산 행렬의 역행렬)의 대각 원소의 역수로 고정된 정규분포가 된다.

import numpy as np

mu = np.array([1.0, -1.0])                   # 사후분포 p(h | x) = N(mu, Sigma): 두 잠재변수의 상관계수 0.9
Sigma = np.array([[1.0, 0.9], [0.9, 1.0]])
Lam = np.linalg.inv(Sigma)                   # 정밀도 행렬 (Sigma의 역행렬)

def kl_gauss(m_q, S_q):                      # D_KL(q ‖ p) = ln p(x) − ELBO(q)
    d = m_q - mu
    return 0.5 * (np.trace(Lam @ S_q) + d @ Lam @ d - 2 + np.log(np.linalg.det(Sigma) / np.linalg.det(S_q)))

# 곱 분포 q(h₁)q(h₂): 하나씩 번갈아 ln qᵢ = ⟨ln p⟩_(나머지) + 상수 로 바꾼다 (좌표 상승, CAVI)
m = np.zeros(2); var = 1 / np.diag(Lam)      # 각 인자의 분산은 첫 갱신에서 곧바로 1/Λᵢᵢ로 정해진다
for sweep in range(1, 31):
    m[0] = mu[0] - Lam[0, 1] / Lam[0, 0] * (m[1] - mu[1])
    m[1] = mu[1] - Lam[1, 0] / Lam[1, 1] * (m[0] - mu[0])
    if sweep in (1, 5, 10, 30):
        print(f"{sweep:2d}번째 훑기: 평균 ({m[0]:.4f}, {m[1]:.4f}), 틈 KL = {kl_gauss(m, np.diag(var)):.4f}")
print(f"곱 분포의 표준편차 {np.sqrt(var[0]):.4f} (참 표준편차 {np.sqrt(Sigma[0, 0]):.4f})")
print(f"상관을 표현할 수 있는 q (완전 공분산)의 틈 = {kl_gauss(mu, Sigma):.4f}")
for rho in (0.5, 0.9, 0.99):
    print(f"상관계수 {rho}: 곱 분포의 표준편차 {np.sqrt(1 - rho**2):.3f}, 남는 틈 {-0.5 * np.log(1 - rho**2):.3f} nat, "
          f"훑기마다 오차가 {rho**2:.4f}배")
#  1번째 훑기: 평균 (1.9000, -0.1900), 틈 KL = 1.2354
#  5번째 훑기: 평균 (1.3874, -0.6513), 틈 KL = 0.9054
# 10번째 훑기: 평균 (1.1351, -0.8784), 틈 KL = 0.8395
# 30번째 훑기: 평균 (1.0020, -0.9982), 틈 KL = 0.8304
# 곱 분포의 표준편차 0.4359 (참 표준편차 1.0000)
# 상관을 표현할 수 있는 q (완전 공분산)의 틈 = 0.0000
# 상관계수 0.5: 곱 분포의 표준편차 0.866, 남는 틈 0.144 nat, 훑기마다 오차가 0.2500배
# 상관계수 0.9: 곱 분포의 표준편차 0.436, 남는 틈 0.830 nat, 훑기마다 오차가 0.8100배
# 상관계수 0.99: 곱 분포의 표준편차 0.141, 남는 틈 1.959 nat, 훑기마다 오차가 0.9801배

평균은 참값 (1, −1)로 찾아가지만 표준편차는 0.436에서 움직이지 않고, 틈은 0.830 nat에서 멈춘다. 상관을 표현할 수 있는 q라면 틈이 0이 되므로, 남은 틈은 모두 곱 분포라는 모임의 한계에서 온다. 상관이 강할수록 곱 분포는 더 좁아지고 틈은 커지며, 한 번 훑을 때마다 평균의 오차가 상관계수의 제곱배로만 줄어 수렴도 느려진다. 상관계수를 ρ라 하면 곱 분포의 표준편차는 √(1 − ρ²), 끝까지 남는 틈은 −½ ln(1 − ρ²) nat이고, 오차는 한 번 훑을 때마다 ρ²배가 된다. 상관계수가 0.99이면 오차를 1000분의 1로 줄이는 데 약 340번을 훑어야 한다.

직접 움직여 보기상관된 정규분포와 좌표 상승새 창에서 열기 ↗

ML에서: 닫힌 식으로 도는 갱신

베이즈 모형에서 사후분포 p(h | x)의 분배함수 p(x)는 대개 계산할 수 없다. 평균장 변분 추론은 잠재변수를 몇 묶음으로 나누고 묶음끼리 독립인 q를 둔 뒤, 인자마다 ln qᵢ = ⟨ln p(x, h)⟩_(나머지) + 상수를 번갈아 적용해 ELBO를 키운다. 토픽 모형이나 베이즈 혼합 모형처럼 인자들이 지수족이 되는 모형에서는 이 갱신이 닫힌 식으로 풀려 샘플링 없이 빠르게 돌아간다.

위젯 1에는 평균장 방정식을 푸는 세 가지 갱신 방식(두 평균을 한꺼번에, 하나씩, 새 값과 옛 값을 반씩 섞어 한꺼번에)의 경로를 자유에너지 지도 위에 그리는 단추가 있다. 아래 문제 6을 풀어 본 뒤 「문제 6 불러오기」로 세 경로를 견주어 보자.

직접 움직여 보기두 스핀의 자유에너지 지도새 창에서 열기 ↗

문제 5. 복도에서 마주친 두 사람

좁은 복도에서 두 사람이 마주 걸어온다. 복도의 자리는 북쪽 벽 쪽과 남쪽 벽 쪽 둘뿐이고, 둘이 같은 쪽에 있으면 부딪힌다. 두 사람은 모두 1초마다 상대가 지금 서 있는 쪽의 반대쪽으로 옮겨 선다. 처음에는 둘 다 북쪽에 있다. (가) 두 사람이 매초 동시에 움직이면 10초 동안 몇 번 마주치는가? (나) 매초 한 사람이 먼저 옮겨 서고, 다른 사람은 그것을 본 뒤에 옮겨 서면 어떻게 되는가?

김민준 M03
김민준

둘 다 상대 반대쪽으로 가니까 한 번 움직이면 끝나는 거 아니에요? 북쪽에서 둘 다 남쪽으로… 어, 둘 다 남쪽이네요.

이서연 S01
이서연

다음 초에는 둘 다 상대가 남쪽에 있는 걸 보고 북쪽으로 가. 북, 남, 북, 남… 매초 같은 쪽에 서니까 10초면 열 번 다 마주쳐.

선생님 T02
선생님

규칙은 둘 다 옳게 지켰는데 왜 안 풀리죠?

이서연 S07
이서연

둘 다 상대가 1초 전에 서 있던 자리에 맞춰 움직이니까요. 내가 움직이는 동안 상대도 움직인다는 걸 둘 다 몰라요.

김민준 M07
김민준

(나)는 앞사람이 남쪽으로 옮기고, 뒷사람은 남쪽에 선 앞사람을 보고 북쪽에 그대로 있어요. 1초 만에 지나가요.

선생님 T13
선생님

한 사람이 바뀐 자리를 보고 움직이기만 해도 풀려요. 이 순서의 차이가 아래 문제에서 그대로 다시 나와요.

문제 6. 한꺼번에 갱신하면

βJ = 2, 편향이 없는 스핀 두 개의 평균장 방정식을 (0.9, −0.9)에서 출발해 푼다. (가) 두 평균을 한꺼번에 m ← tanh(βJ × 상대의 m)로 바꾸는 일을 되풀이하면 어떻게 되는가? (나) 하나씩 바꾸면 어떻게 되는가? (다) 새 값과 옛 값을 반씩 섞어 한꺼번에 바꾸면 어떻게 되는가? (풀어 본 뒤 위젯 1의 「문제 6 불러오기」로 확인해 보자.)

김민준 M05
김민준

스핀이 많아지면 반복문이 느리니까 행렬 곱 한 번으로 한꺼번에 바꾸는 게 당연하죠. 돌려 보면… (−0.947, 0.947), (0.956, −0.956), (−0.957, 0.957)… 어? 부호가 계속 번갈아 바뀌면서 안 멈춰요.

이서연 S04
이서연

−βF[q]를 봐. 처음 −1.223에서 −1.628까지만 올라가고 거기서 머물러. 두 스핀이 늘 반대 방향이라 결합 에너지를 오히려 손해 보고 있어.

선생님 T02
선생님

하나씩 바꾸면요?

김민준 M11
김민준

먼저 m₁ = tanh(2 × (−0.9)) = −0.947로 바꾸고, 바뀐 값을 넣어 m₂ = tanh(2 × (−0.947)) = −0.956이에요. 한 번 훑었을 뿐인데 둘이 같은 방향이 되고, 몇 번 더 훑으면 (−0.9575, −0.9575)에 멈춰요. −βF[q]가 2.039고요.

이서연 S07
이서연

하나씩 바꾸는 건 그 변수에 대해 F[q]를 정확히 최소화하는 거라서 F[q]가 줄기만 해. 한꺼번에 바꾸면 각자 상대의 옛 값에 맞추느라 서로 엇갈리고, 그런 보장이 없어.

김민준 M12
김민준

아까 복도에서 마주친 두 사람이랑 똑같네요. 동시에 같은 쪽으로 비켜서 또 마주치고, 한 사람만 먼저 움직이면 바로 지나가고요.

선생님 T12
선생님

(다)는 흔히 쓰는 처방인데, 결과를 봐요.

김민준 M06
김민준

반씩 섞으니까 진동은 멈췄는데, 40번 뒤에 (0, 0)에 가 있어요. −βF[q] = 2 ln 2 = 1.386이라 하나씩 바꾼 2.039보다 낮아요.

이서연 S08
이서연

0도 자기 일관 방정식의 해이긴 하니까 멈출 수는 있지. 그런데 βJ = 2에서는 F[q]가 가장 낮은 곳이 아니라 봉우리와 골짜기 사이에 걸린 안장점이야. 방정식을 만족한다고 가장 나은 곱 분포인 건 아니구나.

선생님 T14
선생님

격자에서는 서로 직접 묶이지 않은 스핀끼리는 한꺼번에 바꿔도 돼요. 바둑판의 검은 칸끼리는 이웃이 아니니 검은 칸을 모두 바꾸고, 그다음 흰 칸을 모두 바꾸는 식이죠. 한꺼번에 바꾸는 속도를 얻으면서 하나씩 바꾸는 보장도 지키는 방법이에요.