거리 아닌 거리, 직선 아닌 직선

발산: 가까이서 보면 계량이 된다

KL 발산은 방향에 따라 값이 다르다. 그런데 실제로 KL을 재는 자리는 대개 아주 가까운 두 분포 사이다. 모델을 학습시키는 일은 매개변수를 조금씩 바꾸는 걸음의 연속이라, 걸음마다 바로 앞의 분포와 아주 조금 달라진 분포를 견주게 된다. 그렇다면 두 분포가 아주 가까워지면 어떻게 될까? 비대칭은 끝까지 남을까?

정규분포 두 개로 재 보자. N(0,1)N(0, 1)과 N(1,22)N(1, 2^2)처럼 멀리 떨어진 두 분포는 한 방향이 0.443, 반대 방향이 1.307로 세 배 가까이 다르다. 이번에는 두 번째 분포를 N(0.1, 1.052)N(0.1,\ 1.05^2)로 바짝 붙인다. 두 방향의 값은 0.00684와 0.00746으로 거의 같아진다. 게다가 둘 다 0.0075라는 한 값 가까이에 모인다.

p = N(0, 1)은 그대로 두고 q = N(0, σ²)의 표준편차 σ만 바꾸며 잰 KL 발산의 두 방향(실선과 긴 점선)과 2차 근사(빨간 점선). σ = 1, 곧 q = p 근처에서는 세 곡선이 한데 붙고, 멀어질수록 두 방향이 갈라진다. 두 방향의 곡선은 σ를 1/σ로 바꾸면 서로 겹친다
p = N(0, 1)은 그대로 두고 q = N(0, σ²)의 표준편차 σ만 바꾸며 잰 KL 발산의 두 방향(실선과 긴 점선)과 2차 근사(빨간 점선). σ = 1, 곧 q = p 근처에서는 세 곡선이 한데 붙고, 멀어질수록 두 방향이 갈라진다. 두 방향의 곡선은 σ를 1/σ로 바꾸면 서로 겹친다

가까이에서는 비대칭이 사라지는 것이다. 0.0075가 어디서 왔는지는 D(p∥q)D(p \| q)를 q=pq = p 근방에서 테일러 전개하면 보인다.

D(p ∥ p+δ)=12 gij δiδj+O(δ3) \textcolor{#7a5c1e}{D}(\textcolor{#2e8b3a}{p} \,\|\, \textcolor{#2e8b3a}{p} + \textcolor{#964079}{\delta}) = \frac{1}{2}\,\textcolor{#d62728}{g}_{ij}\,\textcolor{#964079}{\delta}^i\textcolor{#964079}{\delta}^j + O(\textcolor{#964079}{\delta}^3)
D발산p기준 분포δi좌표 θi 로 잰 작은 변위gij발산이 만들어 내는 리만 계량 (KL 이면 피셔 정보행렬)O(δ3)δ 의 세제곱 크기로 작아지는 나머지 \begin{array}{ll} \textcolor{#7a5c1e}{D} & \text{발산} \\ \textcolor{#2e8b3a}{p} & \text{기준 분포} \\ \textcolor{#964079}{\delta}^i & \text{좌표 } \textcolor{#1b9e77}{\theta}^i \text{ 로 잰 작은 변위} \\ \textcolor{#d62728}{g}_{ij} & \text{발산이 만들어 내는 리만 계량 (KL 이면 피셔 정보행렬)} \\ O(\textcolor{#964079}{\delta}^3) & \textcolor{#964079}{\delta} \text{ 의 세제곱 크기로 작아지는 나머지} \end{array}

1차 항이 없는 것은 q=p\textcolor{#2e8b3a}{q} = \textcolor{#2e8b3a}{p}에서 발산이 최솟값 0을 갖기 때문이다. 2차 항은 변위 δ\textcolor{#964079}{\delta}를 −δ-\textcolor{#964079}{\delta}로 바꿔도 그대로다. 그래서 두 점이 충분히 가까우면 D(p∥q)≈D(q∥p)\textcolor{#7a5c1e}{D}(\textcolor{#2e8b3a}{p} \| \textcolor{#2e8b3a}{q}) \approx \textcolor{#7a5c1e}{D}(\textcolor{#2e8b3a}{q} \| \textcolor{#2e8b3a}{p})가 된다.

2차 항의 계수 gij\textcolor{#d62728}{g}_{ij}는 리만 계량이다. 계량은 점마다 놓인 자라고 생각하면 된다. 그 점 근처에서 작은 걸음의 길이와 각도를 재는 눈금이다. KL 발산이 만드는 자는 피셔 정보행렬(매개변수를 조금 바꿀 때 분포가 얼마나 달라지는지를 재는 행렬)이다. 앞의 0.0075는 이 자로 잰 걸음 길이의 제곱의 절반이다. 지구 표면이 좁은 영역에서는 평면처럼 보이듯, 비대칭인 발산도 한 점 바로 근처에서는 대칭인 자로 보인다.

비대칭은 3차 항에서 처음 나타난다. 위의 0.00684와 0.00746이 조금 다른 것도 이 3차 항 때문이다. 이 3차 항의 계수를 큐빅 텐서 C\textcolor{#8f8f00}{C}(첨자가 셋인 텐서)라 부르는데, 짝을 이루는 두 접속(이웃한 점의 화살표, 곧 벡터를 서로 비교하는 규칙) ∇\textcolor{#9467bd}{\nabla}와 ∇∗\textcolor{#d4449f}{\nabla^*}가 얼마나 벌어지는지가 바로 이 양으로 정해져서, 유클리드 거리의 제곱처럼 대칭인 발산에서는 C=0\textcolor{#8f8f00}{C} = 0이고 두 접속이 하나로 합쳐진다. 여기서 텐서는 좌표를 바꿀 때 정해진 규칙대로만 성분이 바뀌는 양이다(PyTorch의 tensor와는 이름만 같다).

KL 발산에서 쓴 성질만 추려 내면 "발산"이라는 이름에 필요한 조건이 된다.

따라 계산해 보기 1 — 두 정규분포 사이의 KL과 피셔 계량

정규분포 N(μ,σ2)N(\textcolor{#1b9e77}{\mu}, \textcolor{#1b9e77}{\sigma}^2)(평균 μ\textcolor{#1b9e77}{\mu}, 표준편차 σ\textcolor{#1b9e77}{\sigma}) 전체는 좌표 (μ,σ)(\textcolor{#1b9e77}{\mu}, \textcolor{#1b9e77}{\sigma})를 갖는 2차원 매니폴드다. 두 정규분포 사이의 KL 발산은 적분을 직접 계산하면 닫힌 식으로 나온다.

D(N(μ1,σ12) ∥ N(μ2,σ22))=log⁡σ2σ1+σ12+(μ1−μ2)22σ22−12 \textcolor{#7a5c1e}{D}\big(N(\textcolor{#1b9e77}{\mu_1}, \textcolor{#1b9e77}{\sigma_1}^2) \,\|\, N(\textcolor{#1b9e77}{\mu_2}, \textcolor{#1b9e77}{\sigma_2}^2)\big) = \log\frac{\textcolor{#1b9e77}{\sigma_2}}{\textcolor{#1b9e77}{\sigma_1}} + \frac{\textcolor{#1b9e77}{\sigma_1}^2 + (\textcolor{#1b9e77}{\mu_1} - \textcolor{#1b9e77}{\mu_2})^2}{2\textcolor{#1b9e77}{\sigma_2}^2} - \frac{1}{2}
DKL 발산μ,σ정규분포 매니폴드의 좌표 (평균, 표준편차) \begin{array}{ll} \textcolor{#7a5c1e}{D} & \text{KL 발산} \\ \textcolor{#1b9e77}{\mu}, \textcolor{#1b9e77}{\sigma} & \text{정규분포 매니폴드의 좌표 (평균, 표준편차)} \end{array}

멀리 떨어진 두 점. p=N(0,1)\textcolor{#2e8b3a}{p} = N(0, 1), q=N(1,22)\textcolor{#2e8b3a}{q} = N(1, 2^2)을 넣으면 D(p∥q)=log⁡2+1+18−12≈0.443\textcolor{#7a5c1e}{D}(\textcolor{#2e8b3a}{p} \| \textcolor{#2e8b3a}{q}) = \log 2 + \frac{1 + 1}{8} - \frac12 \approx 0.443이고, 순서를 바꾸면 D(q∥p)=log⁡12+4+12−12≈1.307\textcolor{#7a5c1e}{D}(\textcolor{#2e8b3a}{q} \| \textcolor{#2e8b3a}{p}) = \log\frac12 + \frac{4 + 1}{2} - \frac12 \approx 1.307이다.

가까운 두 점. q=N(0.1, 1.052)\textcolor{#2e8b3a}{q} = N(0.1,\ 1.05^2)이면 변위는 δμ=0.1\delta\textcolor{#1b9e77}{\mu} = 0.1, δσ=0.05\delta\textcolor{#1b9e77}{\sigma} = 0.05이고, 같은 식으로 D(p∥q)≈0.00684\textcolor{#7a5c1e}{D}(\textcolor{#2e8b3a}{p} \| \textcolor{#2e8b3a}{q}) \approx 0.00684, D(q∥p)≈0.00746\textcolor{#7a5c1e}{D}(\textcolor{#2e8b3a}{q} \| \textcolor{#2e8b3a}{p}) \approx 0.00746을 얻는다.

2차 근사도 계산해 보자. 위 식을 δμ,δσ\delta\textcolor{#1b9e77}{\mu}, \delta\textcolor{#1b9e77}{\sigma}에 대해 2차까지 전개하면

D≈12((δμ)2σ2+2 (δσ)2σ2),g=(1/σ2002/σ2) \textcolor{#7a5c1e}{D} \approx \frac{1}{2}\left(\frac{(\delta\textcolor{#1b9e77}{\mu})^2}{\textcolor{#1b9e77}{\sigma}^2} + \frac{2\,(\delta\textcolor{#1b9e77}{\sigma})^2}{\textcolor{#1b9e77}{\sigma}^2}\right), \qquad \textcolor{#d62728}{g} = \begin{pmatrix} 1/\textcolor{#1b9e77}{\sigma}^2 & 0 \\ 0 & 2/\textcolor{#1b9e77}{\sigma}^2 \end{pmatrix}
g(μ,σ) 좌표에서의 피셔 계량δμ,δσ두 분포의 좌표 차이 \begin{array}{ll} \textcolor{#d62728}{g} & (\textcolor{#1b9e77}{\mu}, \textcolor{#1b9e77}{\sigma}) \text{ 좌표에서의 피셔 계량} \\ \delta\textcolor{#1b9e77}{\mu}, \delta\textcolor{#1b9e77}{\sigma} & \text{두 분포의 좌표 차이} \end{array}

이 되어, σ=1\textcolor{#1b9e77}{\sigma} = 1에서 12(0.01+2×0.0025)=0.0075\frac12(0.01 + 2 \times 0.0025) = 0.0075를 준다. 두 방향의 정확한 값 0.00684와 0.00746이 모두 이 값 가까이에 있고, 둘의 차이 0.0006은 δ\textcolor{#964079}{\delta}의 세제곱 크기다.

위의 계산을 이번에는 세 가지 결과를 갖는 분포들의 공간(확률 삼각형) 위에서 해 본다. 아래 위젯에는 두 분포 p\textcolor{#2e8b3a}{p}와 q\textcolor{#2e8b3a}{q}가 놓여 있다. 두 점을 끌어 보면서 D(p∥q)\textcolor{#7a5c1e}{D}(\textcolor{#2e8b3a}{p} \| \textcolor{#2e8b3a}{q})와 D(q∥p)\textcolor{#7a5c1e}{D}(\textcolor{#2e8b3a}{q} \| \textcolor{#2e8b3a}{p})가 얼마나 다른지, q\textcolor{#2e8b3a}{q}를 p\textcolor{#2e8b3a}{p}에 가까이 가져가면 둘이 피셔 계량의 이차식 12gijδiδj\frac12 \textcolor{#d62728}{g}_{ij}\textcolor{#964079}{\delta}^i\textcolor{#964079}{\delta}^j로 함께 수렴하는지 확인해 보자. p\textcolor{#2e8b3a}{p} 둘레의 세 곡선은 "발산이 같은 값인 점들"이다. 그 발산 값을 키우면 세 곡선이 어떻게 갈라지는지도 보자. 아래 문제 5의 두 분포도 단추로 불러올 수 있다.

ML에서: 신뢰 영역을 KL로 잰다

강화학습의 신뢰 영역 정책 최적화(TRPO, Schulman 등, 2015)는 정책을 한 번에 너무 많이 바꾸지 않도록, 곧 한 번에 움직여도 믿을 수 있는 범위(신뢰 영역) 안에 머물도록, 새 정책과 옛 정책 사이의 평균 KL 발산이 작은 값을 넘지 않게 제약한다. 실제로 이 제약을 풀 때는 KL 발산을 위의 2차 근사 12gijδiδj\frac12 \textcolor{#d62728}{g}_{ij}\textcolor{#964079}{\delta}^i\textcolor{#964079}{\delta}^j로 바꾸어 피셔 정보행렬로 계산한다. 여기서 움직이는 것은 정책 신경망의 매개변수이고, 가까운 곳에서는 발산이 곧 계량이라는 사실이 알고리즘에 그대로 쓰인다.

제약이 왜 하필 KL인지는 목적함수 쪽에서도 보인다. TRPO는 옛 정책으로 모은 데이터로 새 정책을 평가하므로, 각 표본에 두 정책 확률의 비 πnew/πold\pi_{\text{new}}/\pi_{\text{old}}(π\pi는 정책을 뜻하는 글자)를 곱한다(중요도 샘플링). 이 비가 표본마다 얼마나 들쭉날쭉한지(분산)는 χ2\chi^2(카이제곱) 발산이라는 또 다른 발산과 같다. 두 정책이 가까우면 χ2\chi^2 발산도 같은 자로 잰 이차식 gijδiδj\textcolor{#d62728}{g}_{ij}\textcolor{#964079}{\delta}^i\textcolor{#964079}{\delta}^j, 곧 KL의 두 배가 된다. 추정을 믿을 수 있는 범위와 KL 제약이 같은 자로 재진다는 뜻이다. 다만 이 연결은 두 정책이 가까울 때만 성립한다.

문제 4 — 가까운 두 정규분포와 피셔 계량

계산 p=N(0,1)\textcolor{#2e8b3a}{p} = N(0, 1), q=N(0.05, 0.952)\textcolor{#2e8b3a}{q} = N(0.05,\ 0.95^2)에 대해 정확한 D(p∥q)\textcolor{#7a5c1e}{D}(\textcolor{#2e8b3a}{p} \| \textcolor{#2e8b3a}{q})를 구하고, 2차 근사 12gijδiδj\frac12 \textcolor{#d62728}{g}_{ij}\textcolor{#964079}{\delta}^i\textcolor{#964079}{\delta}^j와 비교하라. 좌표를 (μ,σ)(\textcolor{#1b9e77}{\mu}, \textcolor{#1b9e77}{\sigma}) 대신 (μ,v=σ2)(\textcolor{#1b9e77}{\mu}, \textcolor{#1b9e77}{v} = \textcolor{#1b9e77}{\sigma}^2)로 잡아도 같은 근삿값이 나오는지 확인하라.

같이 풀기 세션 4

김민준 M01
김민준

정확한 값은 따라 계산해 보기 식에 넣으면 0.00411이에요. 근사는 δμ=0.05\delta\textcolor{#1b9e77}{\mu} = 0.05, δσ=−0.05\delta\textcolor{#1b9e77}{\sigma} = -0.05니까 12(0.052+0.052)=0.0025\frac12(0.05^2 + 0.05^2) = 0.0025요. 꽤 어긋나네요.

이서연 S01
이서연

너 계량을 단위행렬로 놨네. μ\textcolor{#1b9e77}{\mu}랑 σ\textcolor{#1b9e77}{\sigma}는 그냥 좌표 이름표일 뿐이고, 눈금은 피셔 계량이 정해. g=diag(1/σ2, 2/σ2)\textcolor{#d62728}{g} = \text{diag}(1/\textcolor{#1b9e77}{\sigma}^2,\ 2/\textcolor{#1b9e77}{\sigma}^2)라서 σ\textcolor{#1b9e77}{\sigma} 방향에 2가 붙어. 그러면 12(0.0025+2×0.0025)=0.00375\frac12(0.0025 + 2 \times 0.0025) = 0.00375야.

김민준 M04
김민준

아, 좌표 격자 간격이 곧 거리는 아니라고 그렇게 들었는데… 좌표만 보고 유클리드로 계산해 버렸네요.

선생님 T01
선생님

민준 학생은 과제에서 이런 실수 해 본 적 없어요?

김민준 M01
김민준

있어요. 입력 특성 두 개를 표준화 안 하고 거리를 쟀다가 단위가 큰 특성이 결과를 다 먹어 버려서, 조교가 “특성마다 눈금이 다르잖아요” 하고 코멘트를 달았어요.

이서연 S01
이서연

그럼 분산 좌표도 해 볼게. δv=0.952−1=−0.0975\textcolor{#964079}{\delta} \textcolor{#1b9e77}{v} = 0.95^2 - 1 = -0.0975니까, 같은 계량에 넣으면 12(0.0025+2×0.09752)≈0.0108\frac12(0.0025 + 2 \times 0.0975^2) \approx 0.0108. 어, 세 배 가까이 나오네?

선생님 T01
선생님

계량을 (μ,σ)(\textcolor{#1b9e77}{\mu}, \textcolor{#1b9e77}{\sigma}) 좌표에서 구해 놓고 (μ,v)(\textcolor{#1b9e77}{\mu}, \textcolor{#1b9e77}{v}) 좌표의 변위를 넣었죠. 좌표를 바꾸면 계량 성분도 같이 바뀌어요.

이서연 S07
이서연

v=σ2\textcolor{#1b9e77}{v} = \textcolor{#1b9e77}{\sigma}^2이면 dv=2σ dσd\textcolor{#1b9e77}{v} = 2\textcolor{#1b9e77}{\sigma}\,d\textcolor{#1b9e77}{\sigma}니까, 2σ2dσ2=2σ2⋅dv24σ2=dv22v2\frac{2}{\textcolor{#1b9e77}{\sigma}^2}d\textcolor{#1b9e77}{\sigma}^2 = \frac{2}{\textcolor{#1b9e77}{\sigma}^2}\cdot\frac{d\textcolor{#1b9e77}{v}^2}{4\textcolor{#1b9e77}{\sigma}^2} = \frac{d\textcolor{#1b9e77}{v}^2}{2\textcolor{#1b9e77}{v}^2}. 그러면 gvv=12v2\textcolor{#d62728}{g}_{vv} = \frac{1}{2\textcolor{#1b9e77}{v}^2}이고, 12(0.0025+0.09752/2)≈0.00363\frac12(0.0025 + 0.0975^2/2) \approx 0.00363. 이제 0.00375랑 거의 같아요. 남은 차이는 3차 항이고요.

선생님 T13
선생님

그래요. 근삿값이 좌표에 따라 흔들리면 안 된다는 것, 그게 계량이 텐서여야 하는 이유예요.

문제 5 — 드문 행동을 건드리면

계산 행동이 셋인 정책(한 상황에서 세 행동을 고를 확률)이 지금 p=(0.6, 0.3, 0.1)\textcolor{#2e8b3a}{p} = (0.6,\ 0.3,\ 0.1)이다. 신뢰 영역 방법처럼 한 번의 갱신에서 D(p∥q)≤0.01\textcolor{#7a5c1e}{D}(\textcolor{#2e8b3a}{p} \| \textcolor{#2e8b3a}{q}) \le 0.01을 지키려 한다. 두 갱신 후보가 있다. (가) 첫째 행동의 확률 0.05를 둘째 행동으로 옮긴 q=(0.55, 0.35, 0.10)\textcolor{#2e8b3a}{q} = (0.55,\ 0.35,\ 0.10), (나) 0.05를 셋째 행동으로 옮긴 q=(0.55, 0.30, 0.15)\textcolor{#2e8b3a}{q} = (0.55,\ 0.30,\ 0.15). 각각 D(p∥q)\textcolor{#7a5c1e}{D}(\textcolor{#2e8b3a}{p} \| \textcolor{#2e8b3a}{q})와 2차 근사 12gijδiδj=12∑i(qi−pi)2/pi\frac12 \textcolor{#d62728}{g}_{ij}\textcolor{#964079}{\delta}^i\textcolor{#964079}{\delta}^j = \frac12\sum_i (\textcolor{#2e8b3a}{q}_i - \textcolor{#2e8b3a}{p}_i)^2/\textcolor{#2e8b3a}{p}_i(결과가 셋인 분포의 피셔 계량으로 잰 값)를 구하고, 제약을 지키는지 판정하라. 위 도구의 「문제 5 (가)」「문제 5 (나)」 단추로 두 분포를 불러와 풀이와 견주어 볼 수 있다.

같이 풀기 세션 5

김민준 M01
김민준

둘 다 0.05씩 옮겼으니까 바뀐 크기가 같잖아요. 유클리드로 재면 둘 다 12(0.052+0.052)=0.0025\frac12(0.05^2 + 0.05^2) = 0.0025라서 둘 다 통과예요.

선생님 T01
선생님

위 도구에서 두 단추를 차례로 눌러 볼래요?

김민준 M06
김민준

(가)는 D(p∥q)=0.0060\textcolor{#7a5c1e}{D}(\textcolor{#2e8b3a}{p} \| \textcolor{#2e8b3a}{q}) = 0.0060인데 (나)는 0.0117이에요. (나)는 제약을 넘네요!

이서연 S01
이서연

2차 근사로 보면 (가)는 12(0.052/0.6+0.052/0.3)=0.00625\frac12(0.05^2/0.6 + 0.05^2/0.3) = 0.00625, (나)는 12(0.052/0.6+0.052/0.1)≈0.0146\frac12(0.05^2/0.6 + 0.05^2/0.1) \approx 0.0146이야. 확률이 0.1인 칸에서 0.05가 바뀌면 분모가 작아서 크게 쳐.

선생님 T01
선생님

같은 0.05라도 원래 확률이 작던 행동에서 옮기면 분포가 더 많이 바뀐 거예요. 0.3이 0.35가 되면 1.17배지만, 0.1이 0.15가 되면 1.5배죠.

김민준 M07
김민준

앞 문제에서 계량을 단위행렬로 놨던 실수랑 똑같네요. 확률 공간도 칸마다 눈금이 달라요. 그래서 신뢰 영역을 유클리드 거리가 아니라 KL로 재는 거고요.

이서연 S01
이서연

(나)에서는 근사 0.0146이 정확한 값 0.0117보다 꽤 크네. 0.1에 견주면 0.05는 작은 걸음이 아니라서 3차 항이 커진 거겠지.