7장 — 볼츠만 분포

에너지 기반 모델: 에너지를 배우는 볼츠만 분포

지금까지 에너지는 물리계가 정해 주는 것이었고, 분배함수는 에너지가 정해지면 함께 정해지는 합이었다. 그렇다면 에너지 함수 자체를 신경망으로 배우면 어떻게 될까? 매개변수가 바뀔 때마다 Z도 함께 바뀔 텐데, 그때도 Z를 나눗셈에만 쓰고 잊어도 될까?

역사: 확률보다 에너지를 먼저

볼츠만 분포를 신경망에 들여온 대표적인 예는 1985년의 볼츠만 머신이다. 켜지고 꺼지는 뉴런들의 상태 전체에 에너지를 매기고 확률을 볼츠만 분포로 정한 모델인데, 학습할 때마다 모델에서 샘플을 뽑아야 해서 무척 느렸다. 그 뒤 많은 확률 모델이 같은 벽 앞에 섰다. 확률 모델은 모든 가능한 경우에 대한 합이나 적분으로 정규화를 지켜야 하는데, 그 합을 계산할 수 없는 경우가 많았던 것이다.

2006년 르쿤, 초프라, 해드셀, 란자토, 황은 「에너지 기반 학습 튜토리얼」에서 순서를 뒤집어 보자고 제안했다. 변수들의 조합마다 얼마나 잘 맞는지를 재는 숫자 하나, 곧 에너지를 매기고, 학습은 관측된 조합에 낮은 에너지를, 관측되지 않은 조합에 높은 에너지를 주는 에너지 함수를 찾는 일로 보자는 것이다. 그러면 정규화를 지킬 필요가 없어 계산할 수 없는 적분을 피해 갈 수 있고, 여러 분류기와 생성 모델이 한 틀에 들어온다고 적었다. 이 절은 거꾸로, 그 에너지를 볼츠만 분포로 다시 확률에 이으면 학습이 무엇을 해야 하는지를 본다.

데이터에서 내리고 모델에서 올린다

에너지를 직접 배우는 가장 작은 모델부터 보자. 상태가 A, B, C 셋뿐이고, 모델은 세 상태의 에너지 E_A, E_B, E_C를 매개변수로 직접 가지며 확률을 온도 1의 볼츠만 분포로 정한다. 상태가 셋뿐이면 분배함수를 손으로 더할 수 있어서, 학습이 에너지를 어느 쪽으로 옮기는지 숫자로 볼 수 있다. 데이터 10개 가운데 A가 6개, B가 3개, C가 1개 나왔고, 처음에는 세 에너지가 모두 0이라 모델 확률이 1/3씩이라고 하자.

데이터의 음의 로그우도는 데이터에서 본 평균 에너지에 ln Z를 더한 것, 곧 0.6E_A + 0.3E_B + 0.1E_C + ln(e^(−E_A) + e^(−E_B) + e^(−E_C))다. 이것을 E_A로 미분하면 첫 항에서 0.6(데이터에서 A의 비율)이, ln Z에서 −1/3(모델이 지금 A에 주는 확률)이 나온다. 세 상태 모두 「데이터 비율 − 모델 확률」이다.

상태 데이터 비율 모델 확률 (처음) 기울기 경사하강으로 빼면
A 0.6 0.333 +0.267 에너지가 내려간다
B 0.3 0.333 −0.033 조금 올라간다
C 0.1 0.333 −0.233 올라간다

학습률 1로 한 걸음 가면 모델 확률은 (0.426, 0.316, 0.258)로 데이터 쪽에 다가가고, 걸음을 되풀이하면 모델 확률이 데이터 비율 (0.6, 0.3, 0.1)과 같아질 때 기울기가 0이 되어 멈춘다. 그때 에너지의 차이는 −ln(비율)의 차이, 곧 B가 A보다 ln 2 = 0.693, C가 A보다 ln 6 = 1.792 높다. 데이터가 많은 곳은 내리고, 모델이 데이터보다 많이 내놓는 곳은 올린 결과다.

상태 셋짜리 에너지 기반 모델의 첫 걸음. 위: 데이터 비율 (0.6, 0.3, 0.1)과 처음 모델 확률 (1/3씩). 아래: 기울기 「데이터 비율 − 모델 확률」만큼 에너지를 옮기면 A는 0.267 내려가고 B는 0.033, C는 0.233 올라간다
상태 셋짜리 에너지 기반 모델의 첫 걸음. 위: 데이터 비율 (0.6, 0.3, 0.1)과 처음 모델 확률 (1/3씩). 아래: 기울기 「데이터 비율 − 모델 확률」만큼 에너지를 옮기면 A는 0.267 내려가고 B는 0.033, C는 0.233 올라간다

이제 상태가 셋이 아니라 이미지 한 장 한 장이라고 하자. 입력 x마다 신경망이 에너지를 매기고 확률을 온도 1의 볼츠만 분포로 정하는 모델을 에너지 기반 모델 (에너지 함수를 배워 볼츠만 분포로 확률을 정하는 모델, energy-based model)이라 한다.

pθ(x)=e−Eθ(x)Z(θ),Z(θ)=∫e−Eθ(x) dx\textcolor{#bcbd22}{p_\theta}(\textcolor{#1b9e77}{x}) = \frac{e^{-\textcolor{#ff7f0e}{E_\theta}(\textcolor{#1b9e77}{x})}}{\textcolor{#667733}{Z}(\textcolor{#1b9e77}{\theta})}, \qquad \textcolor{#667733}{Z}(\textcolor{#1b9e77}{\theta}) = \int e^{-\textcolor{#ff7f0e}{E_\theta}(\textcolor{#1b9e77}{x})}\, d\textcolor{#1b9e77}{x}
pθ모델 분포Eθ(x)신경망이 매긴 입력 x의 에너지x입력 (이미지 등)θ신경망의 매개변수Z(θ)분배함수 (θ가 바뀌면 함께 바뀐다)\begin{array}{ll} \textcolor{#bcbd22}{p_\theta} & \text{모델 분포} \\ \textcolor{#ff7f0e}{E_\theta}(\textcolor{#1b9e77}{x}) & \text{신경망이 매긴 입력 x의 에너지} \\ \textcolor{#1b9e77}{x} & \text{입력 (이미지 등)} \\ \textcolor{#1b9e77}{\theta} & \text{신경망의 매개변수} \\ \textcolor{#667733}{Z}(\textcolor{#1b9e77}{\theta}) & \text{분배함수 (θ가 바뀌면 함께 바뀐다)} \end{array}

데이터의 음의 로그우도는 ⟨E_θ⟩_데이터 + ln Z(θ)이고, 이것을 매개변수로 미분하면 두 항이 나온다. 둘째 항은 ln Z를 미분한 것인데, 볼츠만 인자의 합의 로그를 미분하면 각 상태의 미분을 볼츠만 확률로 평균한 값이 나오므로 모델 분포에 대한 평균이 된다.

∂∂θ⟨−ln⁡pθ⟩데이터=⟨∂Eθ∂θ⟩데이터−⟨∂Eθ∂θ⟩pθ\frac{\partial}{\partial \textcolor{#1b9e77}{\theta}} \big\langle -\ln \textcolor{#bcbd22}{p_\theta} \big\rangle_{\text{데이터}} = \Big\langle \frac{\partial \textcolor{#ff7f0e}{E_\theta}}{\partial \textcolor{#1b9e77}{\theta}} \Big\rangle_{\text{데이터}} - \Big\langle \frac{\partial \textcolor{#ff7f0e}{E_\theta}}{\partial \textcolor{#1b9e77}{\theta}} \Big\rangle_{\textcolor{#bcbd22}{p_\theta}}
⟨⋅⟩데이터학습 데이터에 대한 평균⟨⋅⟩pθ모델 분포에서 뽑은 샘플에 대한 평균Eθ에너지 함수θ매개변수\begin{array}{ll} \langle \cdot \rangle_{\text{데이터}} & \text{학습 데이터에 대한 평균} \\ \langle \cdot \rangle_{\textcolor{#bcbd22}{p_\theta}} & \text{모델 분포에서 뽑은 샘플에 대한 평균} \\ \textcolor{#ff7f0e}{E_\theta} & \text{에너지 함수} \\ \textcolor{#1b9e77}{\theta} & \text{매개변수} \end{array}

경사하강으로 이 기울기를 빼 주면 데이터가 있는 곳의 에너지는 내려가고 모델이 지금 샘플을 내놓는 곳의 에너지는 올라가며, 두 평균이 같아지면 학습이 멈춘다. 상태 셋짜리 모델의 「데이터 비율 − 모델 확률」이 일반 모델에서는 이 두 평균의 차이가 된 것이다. 상태가 셋이면 모델 쪽 평균을 정확히 더할 수 있지만, 이미지 전체에 에너지를 매기는 모델은 Z도 모델 쪽 평균도 계산할 수 없다. 그래서 모델에서 샘플을 뽑아 평균을 어림하는데, Du와 Mordatch(2019)처럼 에너지의 기울기를 따라 움직이는 랑주뱅 동역학이나 HMC 같은 표본 추출기가 그 역할을 맡는다.

ML에서: 분류기는 이미 에너지 기반 모델이다

Grathwohl 외(2020)의 논문 제목은 「Your Classifier is Secretly an Energy Based Model and You Should Treat it Like One」이다. 분류기의 로짓으로 입력과 라벨 쌍의 에너지를 E(x, y) = −z_y(x)로 정하면 분류기가 내놓는 p(y|x)는 x를 고정한 볼츠만 분포이고, 라벨을 합쳐 없애면 p(x)가 Σ_y e^(z_y(x)), 곧 분류기가 매번 계산하고 버리던 분배함수 Z(x)에 비례한다. 이 장을 마친 독자는 이 제목을 「softmax 분류기는 이미 (x, y) 위의 볼츠만 분포인데 조건부 분포만 쓰고 있었고, 정규화 상수로 버린 Z(x)에 입력 x의 분포가 들어 있다」로 읽게 된다.

문제 13. 정규화 상수를 버린 손실

실수 x 위의 에너지 기반 모델 E_θ(x) = θx²/2 (θ > 0)를 데이터 {−2, −1, 1, 2}로 학습한다. (가) 민준은 모델에서 샘플을 뽑는 계산을 아끼려고 기울기의 모델 쪽 평균 항을 빼고, 데이터의 평균 에너지만 경사하강으로 줄였다. θ = 1에서 학습률 0.1로 10걸음 가면 어떻게 되는가? (나) 음의 로그우도를 θ의 함수로 쓰고 최적의 θ를 구하라.

김민준 M11
김민준

모델 샘플을 뽑는 게 제일 비싸니까 그 항은 빼 볼게요. 데이터의 에너지만 낮춰도 방향은 맞겠죠. 데이터의 x²의 평균이 2.5라서 평균 에너지는 1.25θ이고, θ로 미분하면 항상 1.25예요. θ = 1에서 학습률 0.1로 10걸음 가면 θ = −0.25, 손실은 −0.31로 계속 내려가요. 잘 되는데요?

이서연 S06
이서연

θ가 음수면 e^(−θx²/2) = e^(0.125x²)이라서 x가 커질수록 볼츠만 인자가 커져. 적분하면 무한대라서 확률분포가 아니야.

김민준 M04
김민준

어… θ가 0을 지날 때부터 이미 이상했겠네요. θ = 0이면 모든 x의 에너지가 0이라 평평한데, 실수 전체에 평평하면 정규화가 안 되니까요.

선생님 T14
선생님

에너지를 낮추는 것 자체는 맞는 방향이에요. 그런데 데이터가 있는 곳만이 아니라 모든 곳의 에너지를 함께 낮추면 확률은 하나도 안 올라가요. 그걸 붙잡아 주는 게 Z예요. 이 모델의 Z는요?

이서연 S08
이서연

∫e^(−θx²/2)dx = √(2π/θ)예요. θ가 들어 있어요. 모델 쪽 평균 항은 바로 이 ln Z를 미분한 거라, 빼면 Z가 θ에 따라 커지는 걸 못 보는 거였네요. 음의 로그우도는 1.25θ + ½ ln(2π) − ½ ln θ이고, 미분하면 1.25 − 1/(2θ) = 0에서 θ = 0.4예요. 모델의 분산 1/θ = 2.5가 데이터의 x² 평균과 같아지는 곳이에요.

선생님 T02
선생님

방금 미분한 식을 다르게 읽어 볼까요? 1/(2θ)는 무엇이죠?

이서연 S07
이서연

모델 분포에서 x²/2의 평균이요. 분산이 1/θ니까요. 그러니까 기울기는 ⟨x²/2⟩_데이터 − ⟨x²/2⟩_모델이고, 둘이 같아질 때 멈춰요.

선생님 T13
선생님

그게 에너지 기반 모델의 학습 규칙이에요. 데이터에서는 에너지를 내리고, 모델이 뽑는 샘플에서는 올려요. 뒤의 항을 빼면 모든 곳의 에너지가 함께 내려가다 평평하게 무너지죠. 르쿤과 동료들의 튜토리얼(2006)도 정답의 에너지만 내리는 손실은 다른 답의 에너지를 끌어올리지 않아서, 에너지가 어디서나 같은 값으로 무너진 풀이에 빠질 수 있다고 경고했어요.

김민준 M10
김민준

그럼 분류기 학습도 이거예요? cross-entropy를 로짓으로 미분하면 p − y잖아요. 에너지 −z로 미분하면 y − p고, 경사하강으로 빼 주면 정답 클래스의 에너지는 내려가고 모든 클래스의 에너지가 모델 확률 p만큼씩 올라가요. 제가 외우던 p − y가 「데이터 − 모델」이었네요.

이서연 S12
이서연

너 그거 외우기만 하고 이유는 몰랐지?

김민준 M04
김민준

과제 채점은 코드만 보잖아.