KL 발산은 방향에 따라 값이 다르다. 그런데 실제로 KL을 재는 자리는 대개 아주 가까운 두 분포 사이다. 모델을 학습시키는 일은 매개변수를 조금씩 바꾸는 걸음의 연속이라, 걸음마다 바로 앞의 분포와 아주 조금 달라진 분포를 견주게 된다. 그렇다면 두 분포가 아주 가까워지면 어떻게 될까? 비대칭은 끝까지 남을까?
정규분포 두 개로 재 보자. N(0,1)과 N(1,22)처럼 멀리 떨어진 두 분포는 한 방향이 0.443, 반대 방향이 1.307로 세 배 가까이 다르다. 이번에는 두 번째 분포를 N(0.1,1.052)로 바짝 붙인다. 두 방향의 값은 0.00684와 0.00746으로 거의 같아진다. 게다가 둘 다 0.0075라는 한 값 가까이에 모인다.
p = N(0, 1)은 그대로 두고 q = N(0, σ²)의 표준편차 σ만 바꾸며 잰 KL 발산의 두 방향(실선과 긴 점선)과 2차 근사(빨간 점선). σ = 1, 곧 q = p 근처에서는 세 곡선이 한데 붙고, 멀어질수록 두 방향이 갈라진다. 두 방향의 곡선은 σ를 1/σ로 바꾸면 서로 겹친다
가까이에서는 비대칭이 사라지는 것이다. 0.0075가 어디서 왔는지는 D(p∥q)를 q=p 근방에서 테일러 전개하면 보인다.
1차 항이 없는 것은 q=p에서 발산이 최솟값 0을 갖기 때문이다. 2차 항은 변위 δ를 −δ로 바꿔도 그대로다. 그래서 두 점이 충분히 가까우면 D(p∥q)≈D(q∥p)가 된다.
2차 항의 계수 gij는 리만 계량이다. 계량은 점마다 놓인 자라고 생각하면 된다. 그 점 근처에서 작은 걸음의 길이와 각도를 재는 눈금이다. KL 발산이 만드는 자는 피셔 정보행렬(매개변수를 조금 바꿀 때 분포가 얼마나 달라지는지를 재는 행렬)이다. 앞의 0.0075는 이 자로 잰 걸음 길이의 제곱의 절반이다. 지구 표면이 좁은 영역에서는 평면처럼 보이듯, 비대칭인 발산도 한 점 바로 근처에서는 대칭인 자로 보인다.
비대칭은 3차 항에서 처음 나타난다. 위의 0.00684와 0.00746이 조금 다른 것도 이 3차 항 때문이다. 이 3차 항의 계수를 큐빅 텐서 C(첨자가 셋인 텐서)라 부르는데, 짝을 이루는 두 접속(이웃한 점의 화살표, 곧 벡터를 서로 비교하는 규칙) ∇와 ∇∗가 얼마나 벌어지는지가 바로 이 양으로 정해져서, 유클리드 거리의 제곱처럼 대칭인 발산에서는 C=0이고 두 접속이 하나로 합쳐진다. 여기서 텐서는 좌표를 바꿀 때 정해진 규칙대로만 성분이 바뀌는 양이다(PyTorch의 tensor와는 이름만 같다).
KL 발산에서 쓴 성질만 추려 내면 "발산"이라는 이름에 필요한 조건이 된다.
발산 (비대칭 거리 / Divergence, D(p∥q)) — 두 점 사이의 “방향 있는 떨어짐”. 세 가지 조건을 만족한다: (1) D(p∥q)≥0, (2) D(p∥q)=0⟺p=q, (3) q=p 근방에서 D(p∥p+δ)=21gijδiδj+O(δ3)으로, 0이 아닌 δ에서 늘 양수인 이차식(양의 정부호 이차형식)을 만든다. 대칭성과 삼각부등식은 요구하지 않는다. 거리보다 약한 조건이지만, 리만 계량을 만들어 내기에는 충분하다.
따라 계산해 보기 1 — 두 정규분포 사이의 KL과 피셔 계량
정규분포 N(μ,σ2)(평균 μ, 표준편차 σ) 전체는 좌표 (μ,σ)를 갖는 2차원 매니폴드다. 두 정규분포 사이의 KL 발산은 적분을 직접 계산하면 닫힌 식으로 나온다.
이 되어, σ=1에서 21(0.01+2×0.0025)=0.0075를 준다. 두 방향의 정확한 값 0.00684와 0.00746이 모두 이 값 가까이에 있고, 둘의 차이 0.0006은 δ의 세제곱 크기다.
위의 계산을 이번에는 세 가지 결과를 갖는 분포들의 공간(확률 삼각형) 위에서 해 본다. 아래 위젯에는 두 분포 p와 q가 놓여 있다. 두 점을 끌어 보면서 D(p∥q)와 D(q∥p)가 얼마나 다른지, q를 p에 가까이 가져가면 둘이 피셔 계량의 이차식 21gijδiδj로 함께 수렴하는지 확인해 보자. p 둘레의 세 곡선은 "발산이 같은 값인 점들"이다. 그 발산 값을 키우면 세 곡선이 어떻게 갈라지는지도 보자. 아래 문제 5의 두 분포도 단추로 불러올 수 있다.
ML에서: 신뢰 영역을 KL로 잰다
강화학습의 신뢰 영역 정책 최적화(TRPO, Schulman 등, 2015)는 정책을 한 번에 너무 많이 바꾸지 않도록, 곧 한 번에 움직여도 믿을 수 있는 범위(신뢰 영역) 안에 머물도록, 새 정책과 옛 정책 사이의 평균 KL 발산이 작은 값을 넘지 않게 제약한다. 실제로 이 제약을 풀 때는 KL 발산을 위의 2차 근사 21gijδiδj로 바꾸어 피셔 정보행렬로 계산한다. 여기서 움직이는 것은 정책 신경망의 매개변수이고, 가까운 곳에서는 발산이 곧 계량이라는 사실이 알고리즘에 그대로 쓰인다.
제약이 왜 하필 KL인지는 목적함수 쪽에서도 보인다. TRPO는 옛 정책으로 모은 데이터로 새 정책을 평가하므로, 각 표본에 두 정책 확률의 비 πnew/πold(π는 정책을 뜻하는 글자)를 곱한다(중요도 샘플링). 이 비가 표본마다 얼마나 들쭉날쭉한지(분산)는 χ2(카이제곱) 발산이라는 또 다른 발산과 같다. 두 정책이 가까우면 χ2 발산도 같은 자로 잰 이차식 gijδiδj, 곧 KL의 두 배가 된다. 추정을 믿을 수 있는 범위와 KL 제약이 같은 자로 재진다는 뜻이다. 다만 이 연결은 두 정책이 가까울 때만 성립한다.
문제 4 — 가까운 두 정규분포와 피셔 계량
계산p=N(0,1), q=N(0.05,0.952)에 대해 정확한 D(p∥q)를 구하고, 2차 근사 21gijδiδj와 비교하라. 좌표를 (μ,σ) 대신 (μ,v=σ2)로 잡아도 같은 근삿값이 나오는지 확인하라.
같이 풀기 세션 4
김민준
정확한 값은 따라 계산해 보기 식에 넣으면 0.00411이에요. 근사는 δμ=0.05, δσ=−0.05니까 21(0.052+0.052)=0.0025요. 꽤 어긋나네요.
이서연
너 계량을 단위행렬로 놨네. μ랑 σ는 그냥 좌표 이름표일 뿐이고, 눈금은 피셔 계량이 정해. g=diag(1/σ2,2/σ2)라서 σ 방향에 2가 붙어. 그러면 21(0.0025+2×0.0025)=0.00375야.
김민준
아, 좌표 격자 간격이 곧 거리는 아니라고 그렇게 들었는데… 좌표만 보고 유클리드로 계산해 버렸네요.
선생님
민준 학생은 과제에서 이런 실수 해 본 적 없어요?
김민준
있어요. 입력 특성 두 개를 표준화 안 하고 거리를 쟀다가 단위가 큰 특성이 결과를 다 먹어 버려서, 조교가 “특성마다 눈금이 다르잖아요” 하고 코멘트를 달았어요.
이서연
그럼 분산 좌표도 해 볼게. δv=0.952−1=−0.0975니까, 같은 계량에 넣으면 21(0.0025+2×0.09752)≈0.0108. 어, 세 배 가까이 나오네?
선생님
계량을 (μ,σ) 좌표에서 구해 놓고 (μ,v) 좌표의 변위를 넣었죠. 좌표를 바꾸면 계량 성분도 같이 바뀌어요.
이서연
v=σ2이면 dv=2σdσ니까, σ22dσ2=σ22⋅4σ2dv2=2v2dv2. 그러면 gvv=2v21이고, 21(0.0025+0.09752/2)≈0.00363. 이제 0.00375랑 거의 같아요. 남은 차이는 3차 항이고요.
선생님
그래요. 근삿값이 좌표에 따라 흔들리면 안 된다는 것, 그게 계량이 텐서여야 하는 이유예요.
문제 5 — 드문 행동을 건드리면
계산 행동이 셋인 정책(한 상황에서 세 행동을 고를 확률)이 지금 p=(0.6,0.3,0.1)이다. 신뢰 영역 방법처럼 한 번의 갱신에서 D(p∥q)≤0.01을 지키려 한다. 두 갱신 후보가 있다. (가) 첫째 행동의 확률 0.05를 둘째 행동으로 옮긴 q=(0.55,0.35,0.10), (나) 0.05를 셋째 행동으로 옮긴 q=(0.55,0.30,0.15). 각각 D(p∥q)와 2차 근사 21gijδiδj=21∑i(qi−pi)2/pi(결과가 셋인 분포의 피셔 계량으로 잰 값)를 구하고, 제약을 지키는지 판정하라. 위 도구의 「문제 5 (가)」「문제 5 (나)」 단추로 두 분포를 불러와 풀이와 견주어 볼 수 있다.
같이 풀기 세션 5
김민준
둘 다 0.05씩 옮겼으니까 바뀐 크기가 같잖아요. 유클리드로 재면 둘 다 21(0.052+0.052)=0.0025라서 둘 다 통과예요.
선생님
위 도구에서 두 단추를 차례로 눌러 볼래요?
김민준
(가)는 D(p∥q)=0.0060인데 (나)는 0.0117이에요. (나)는 제약을 넘네요!
이서연
2차 근사로 보면 (가)는 21(0.052/0.6+0.052/0.3)=0.00625, (나)는 21(0.052/0.6+0.052/0.1)≈0.0146이야. 확률이 0.1인 칸에서 0.05가 바뀌면 분모가 작아서 크게 쳐.
선생님
같은 0.05라도 원래 확률이 작던 행동에서 옮기면 분포가 더 많이 바뀐 거예요. 0.3이 0.35가 되면 1.17배지만, 0.1이 0.15가 되면 1.5배죠.
김민준
앞 문제에서 계량을 단위행렬로 놨던 실수랑 똑같네요. 확률 공간도 칸마다 눈금이 달라요. 그래서 신뢰 영역을 유클리드 거리가 아니라 KL로 재는 거고요.
이서연
(나)에서는 근사 0.0146이 정확한 값 0.0117보다 꽤 크네. 0.1에 견주면 0.05는 작은 걸음이 아니라서 3차 항이 커진 거겠지.