수학

MATH / 중급 30번

KL 발산: 거리가 아닌 거리와 방향의 선택

RLHF의 KL 페널티, 증류 손실, ELBO의 간격은 전부 같은 양입니다. 그 양을 정의하고 «두 분포 사이의 거리»라는 설명이 어디서 틀리는지 반례로 확인한 뒤, 방향을 바꾸면 학습이 넓게 덮을지 한 봉우리에 몰릴지가 갈리는 것을 수치로 봅니다.

PALDYN Team31 MIN READ

세 자리에 같은 이름이 붙어 있습니다.

init_kl_coef: 0.05        # RLHF — 참조 모델에서 멀어지지 마라
loss = KLDivLoss()(student_logp, teacher_p)   # 지식 증류
kl_term = ...                                  # VAE의 ELBO 두 항 중 하나

정책을 붙들어 두는 벌점, 교사 모델을 흉내 내는 손실, 그리고 지난 글에서 하한과 참값 사이에 남는다고 했던 간격. 셋의 정체가 같습니다. 흔한 설명은 「두 분포 사이의 거리」인데, 그 설명은 틀렸고 틀린 자리가 실무에서 중요합니다. 방향을 어느 쪽으로 잡느냐에 따라 학습이 전혀 다른 답으로 갑니다.

KL 발산의 정의

초과분에 이름 붙이기

지난 글에서 교차엔트로피가 엔트로피보다 작아질 수 없다는 것을 증명했습니다. 그 초과분에 이름을 붙입니다.

정의. 분포 pp 에 대한 qq 의 쿨백-라이블러 발산(Kullback-Leibler divergence)은 DKL(p ∥ q)=∑ipilog⁡piqiD_{\mathrm{KL}}(p \,\|\, q) = \sum_i p_i \log \frac{p_i}{q_i} 이다. 줄여서 KL 발산이라 부르고, ∥\| 기호는 앞뒤를 바꿔 쓸 수 없다는 표시다.

성질 둘은 이미 증명해 두었습니다. DKL(p ∥ q)≥0D_{\mathrm{KL}}(p \,\|\, q) \ge 0 이라는 것 — 젠센 부등식과 log⁡\log 의 오목성에서 나왔습니다 — 과, 등호가 q=pq = p 일 때만 성립한다는 것입니다.

단위는 로그의 밑을 따라갑니다. 밑이 2면 비트, 자연로그면 내트입니다. 이 글에서는 손으로 세기 좋게 비트를 씁니다.

교차엔트로피와의 분해

로그를 쪼개면 그대로 분해가 나옵니다.

DKL(p ∥ q)=∑ipilog⁡pi−∑ipilog⁡qi=H(p,q)−H(p)D_{\mathrm{KL}}(p \,\|\, q) = \sum_i p_i \log p_i - \sum_i p_i \log q_i = H(p,q) - H(p)

교차엔트로피는 엔트로피와 KL 발산의 합이다

교차엔트로피 = 엔트로피 + KL 발산. 이 한 줄에서 지도학습의 익숙한 사실 하나가 떨어져 나옵니다. 정답 분포 pp 는 데이터가 주는 것이라 학습 중에 바뀌지 않으므로 H(p)H(p) 는 상수이고, 따라서

arg⁡min⁡θH(p,qθ)=arg⁡min⁡θDKL(p ∥ qθ)\arg\min_\theta H(p, q_\theta) = \arg\min_\theta D_{\mathrm{KL}}(p \,\|\, q_\theta)

입니다. 교차엔트로피를 줄이는 것과 KL 발산을 줄이는 것이 같은 일인 이유가 이것입니다. 두 값 자체는 H(p)H(p) 만큼 다르지만 최소가 되는 자리가 같습니다.

한 가지만 구별해 둡니다. 둘이 같은 자리에서 최소가 된다는 것은 어디로 가느냐가 같다는 뜻이지 얼마나 왔는지가 같다는 뜻이 아닙니다. 손실 곡선이 0.8에서 평평해졌다고 해서 모델이 정답에서 0.8만큼 떨어져 있는 것이 아닙니다. 그 0.8 안에는 아무리 학습해도 줄일 수 없는 H(p)H(p) 가 들어 있고, 실제로 줄어든 것은 그 위에 얹힌 KL 부분뿐입니다.

q에 0이 들어가면

실무에서 걸리는 단서를 하나 붙입니다. pi>0p_i > 0 인데 qi=0q_i = 0 이면 그 항이 pilog⁡(pi/0)p_i \log(p_i/0) 이라 값이 무한대입니다. 반대로 pi=0p_i = 0 인 항은 0log⁡0=00 \log 0 = 0 으로 약속해 그냥 사라집니다.

조건을 한 문장으로 적으면 이렇습니다. 분포가 양수 값을 주는 자리들의 모임을 지지집합(support)이라고 하는데, DKL(p ∥ q)D_{\mathrm{KL}}(p \,\|\, q) 가 유한하려면 qq 의 지지집합이 pp 의 지지집합을 덮어야 합니다. pp 가 조금이라도 확률을 준 자리를 qq 가 하나라도 0으로 버리면 그 순간 발산합니다. 이 비대칭이 뒤에서 방향의 차이를 만드는 뿌리입니다.

이 단서가 실무에서 나타나는 자리가 둘입니다.

  • 빈도에서 추정한 분포. 말뭉치를 세어 만든 분포는 한 번도 안 나온 사건에 정확히 0을 줍니다. 그대로 KL을 재면 무한대가 나오므로, 모든 칸에 작은 ε\varepsilon 을 더하고 다시 정규화하는 스무딩을 먼저 겁니다. 0을 없애는 것이 목적이지 분포를 바꾸는 것이 목적이 아니므로 ε\varepsilon 은 작게 잡습니다.
  • softmax의 언더플로. softmax는 수학적으로는 0을 내지 않지만 로짓 차이가 크면 부동소수점에서 0으로 내려앉고, 그러면 log⁡\log 가 음의 무한대가 됩니다. 확률을 만든 뒤 로그를 씌우지 말고 log_softmax로 한 번에 계산하는 관행이 여기서 나옵니다.

라벨 스무딩은 같은 규칙의 반대쪽입니다. 정답을 원-핫에서 (1−ε)(1-\varepsilon) 과 ε/C\varepsilon/C 로 풀어 놓는 조작인데, pp 쪽의 0을 없애는 일이라 무한대를 막는 것과는 상관이 없습니다. 바뀌는 것은 목표 자체입니다. 원-핫이 목표이면 손실을 끝까지 줄이는 길이 정답 로짓을 무한히 키우는 것밖에 없는데, 목표에 바닥을 깔아 주면 유한한 로짓에서 최소가 서게 됩니다.

거리가 아닌 이유

대칭이 아니다

수학에서 «거리»라는 말에는 지켜야 할 조건이 붙어 있습니다. 그중 둘을 KL 발산이 어깁니다. 첫째는 대칭입니다. p=(0.9, 0.1)p = (0.9,\, 0.1) 과 q=(0.5, 0.5)q = (0.5,\, 0.5) 를 놓고 양쪽으로 재 봅니다.

DKL(p ∥ q)=0.9log⁡20.90.5+0.1log⁡20.10.5=0.763−0.232=0.531D_{\mathrm{KL}}(p \,\|\, q) = 0.9 \log_2 \frac{0.9}{0.5} + 0.1 \log_2 \frac{0.1}{0.5} = 0.763 - 0.232 = 0.531

DKL(q ∥ p)=0.5log⁡20.50.9+0.5log⁡20.50.1=−0.424+1.161=0.737D_{\mathrm{KL}}(q \,\|\, p) = 0.5 \log_2 \frac{0.5}{0.9} + 0.5 \log_2 \frac{0.5}{0.1} = -0.424 + 1.161 = 0.737

같은 두 분포인데 재는 방향에 따라 값이 다르다

같은 두 분포인데 0.531과 0.737입니다. 「A와 B 사이가 0.5인데 B와 A 사이는 0.7」인 자를 거리라고 부를 수는 없습니다. 두 계산을 나란히 보면 어느 항이 값을 끌어올렸는지도 읽힙니다. 아래쪽에서는 0.5log⁡25=1.1610.5 \log_2 5 = 1.161 한 항이 거의 전부인데, 그 자리는 qq 가 0.5를 주는데 pp 는 0.1밖에 안 주는 칸입니다. 앞에 오는 분포가 큰 값을 주는 칸에서 뒤의 분포가 인색하면 벌점이 커집니다.

삼각부등식이 깨진다

둘째는 삼각부등식입니다. 거리라면 「돌아가는 길이 곧장 가는 길보다 짧을 수 없다」가 성립해야 합니다. 세 분포로 반례를 만듭니다.

p=(0.5, 0.5),q=(0.9, 0.1),r=(0.99, 0.01)p = (0.5,\, 0.5), \qquad q = (0.9,\, 0.1), \qquad r = (0.99,\, 0.01)

잰 것 값 (비트)
DKL(p ∥ q)D_{\mathrm{KL}}(p \,\|\, q) 0.737
DKL(q ∥ r)D_{\mathrm{KL}}(q \,\|\, r) 0.208
두 값의 합 0.945
DKL(p ∥ r)D_{\mathrm{KL}}(p \,\|\, r) 2.329

pp 에서 qq 를 거쳐 rr 로 가는 길이 0.945인데 곧장 가는 길이 2.329입니다. 돌아가는 길이 절반도 안 되게 짧습니다. 삼각부등식이 어긋나는 정도가 오차 수준이 아니라 두 배가 넘습니다.

반례가 왜 이렇게 쉽게 나오는지도 보입니다. rr 은 두 번째 칸에 0.01만 남겨 둔 분포라, 그 칸에 0.5를 주는 pp 에서 곧장 재면 0.5log⁡250=2.820.5 \log_2 50 = 2.82 라는 큰 항이 한 번에 생깁니다. 그런데 중간에 qq 를 끼우면 그 낙차가 0.5→0.1→0.010.5 \to 0.1 \to 0.01 로 두 번에 나뉘고, KL은 비의 로그를 재는 양이라 큰 낙차 한 번이 작은 낙차 두 번보다 훨씬 비쌉니다. 앞에 오는 분포의 가중치가 매번 달라지는 것이 그 차이를 만듭니다.

대칭으로 만든 JS 발산

그래서 이름이 «거리»가 아니라 발산(divergence)입니다. 발산은 「0 이상이고, 같을 때만 0이 되는 양」까지만 약속하고 대칭성과 삼각부등식은 약속하지 않는 이름입니다.

방향을 고르고 싶지 않을 때가 있습니다. 두 분포를 대등하게 견주고 싶을 때인데, 그때는 가운뎃점을 하나 만들어 양쪽에서 잽니다.

정의. 혼합분포 m=12(p+q)m = \tfrac{1}{2}(p + q) 에 대해 DJS(p, q)=12DKL(p ∥ m)+12DKL(q ∥ m)D_{\mathrm{JS}}(p,\, q) = \tfrac{1}{2} D_{\mathrm{KL}}(p \,\|\, m) + \tfrac{1}{2} D_{\mathrm{KL}}(q \,\|\, m) 를 젠슨-섀넌 발산(Jensen-Shannon divergence)이라 한다.

앞의 p=(0.9,0.1)p = (0.9, 0.1), q=(0.5,0.5)q = (0.5, 0.5) 로 세어 봅니다. m=(0.7, 0.3)m = (0.7,\, 0.3) 이고

DKL(p ∥ m)=0.168,DKL(q ∥ m)=0.126D_{\mathrm{KL}}(p \,\|\, m) = 0.168, \qquad D_{\mathrm{KL}}(q \,\|\, m) = 0.126

이므로 DJS=0.147D_{\mathrm{JS}} = 0.147 입니다.

대칭·유한·상한

얻는 것이 셋입니다.

  • 대칭입니다. 정의가 pp 와 qq 를 똑같이 다루므로 순서를 바꿔도 같은 값입니다. mm 을 만드는 식에서 pp 와 qq 를 바꾸면 같은 mm 이 나오고, 뒤의 두 항은 서로 자리를 바꿀 뿐입니다.
  • 언제나 유한합니다. mm 은 pp 나 qq 가 양수인 모든 자리에서 양수라, 앞 절의 지지집합 조건이 저절로 지켜집니다. 지지집합이 어긋난 짝에서도 값이 나옵니다.
  • 위로 막혀 있습니다. 값이 log⁡2\log 2 를 넘지 않습니다. 비트로 재면 1, 내트로 재면 0.693입니다.

상한이 어디서 오는지는 mm 의 정의에서 바로 읽힙니다. mi≥12pim_i \ge \tfrac{1}{2} p_i 이므로 pi/mi≤2p_i/m_i \le 2 이고, 로그를 씌우면 각 항이 log⁡2\log 2 이하입니다. 두 항을 반씩 섞어도 그대로 log⁡2\log 2 이하입니다. 등호는 pp 와 qq 의 지지집합이 완전히 갈릴 때 — 한쪽이 양수인 칸에서 다른 쪽이 언제나 0일 때 — 섭니다. 그 짝에서 KL은 무한대인데 JS는 1비트에서 멈춘다는 뜻이고, 이것이 두 값의 성격 차이를 가장 잘 보여 주는 자리입니다.

제곱근이 거리가 되는 자리

DJSD_{\mathrm{JS}} 자체는 아직 거리가 아닙니다. 대칭은 얻었지만 삼각부등식은 여전히 성립하지 않습니다. 그런데 제곱근을 씌운 DJS\sqrt{D_{\mathrm{JS}}} 는 삼각부등식까지 만족해 진짜 거리가 됩니다.

제곱근이 붙는 것이 이상해 보이지만, 앞의 반례가 그 이유를 짐작하게 합니다. KL이 삼각부등식을 어긴 방식은 「큰 낙차 한 번이 작은 낙차 두 번보다 지나치게 비싸다」였습니다. 제곱근은 큰 값을 더 많이 줄이는 변환이라 그 불균형을 눌러 줍니다. 대칭이 필요한 자리에서 KL 대신 쓰는 값이 이것입니다.

방향의 선택

정방향은 넓게 덮는다

대칭이 아니라는 것은 불편한 성질이 아니라 고를 수 있는 손잡이입니다. 어느 쪽을 앞에 두느냐가 학습의 답을 바꿉니다. 정의를 다시 보면 이유가 보입니다.

DKL(p ∥ q)=∑ipilog⁡piqiD_{\mathrm{KL}}(p \,\|\, q) = \sum_i p_i \log \frac{p_i}{q_i}

각 항에 pip_i 가 곱해져 있으니 pi=0p_i = 0 인 자리는 qq 가 무엇을 하든 값에 안 들어옵니다. 반대로 pip_i 가 큰 자리에서 qiq_i 가 0에 가까우면 로그가 폭발합니다. 두 봉우리 사이에 얕은 골이 있는 분포로 세어 봅시다.

p=(0.495,  0.01,  0.495)p = (0.495,\; 0.01,\; 0.495)

여기에 후보 둘을 댑니다. qAq_A 는 한 봉우리에 몰려 나머지를 아예 버렸고, qBq_B 는 셋에 고르게 폈습니다.

qA=(0.99,  0.01,  0),qB=(13,  13,  13)q_A = (0.99,\; 0.01,\; 0), \qquad q_B = \left(\tfrac{1}{3},\; \tfrac{1}{3},\; \tfrac{1}{3}\right)

정방향은 넓게 덮게 하고 역방향은 한 봉우리에 몰리게 한다

qAq_A — 한 봉우리 qBq_B — 넓게
정방향 DKL(p ∥ q)D_{\mathrm{KL}}(p \,\|\, q) ∞\infty 0.514
역방향 DKL(q ∥ p)D_{\mathrm{KL}}(q \,\|\, p) 0.990 1.306

정방향은 qBq_B 를 고릅니다. qAq_A 가 세 번째 칸을 0으로 버렸는데 pp 는 거기에 0.495를 주고 있으므로 지지집합 조건이 깨져 값이 무한대입니다. qq 는 pp 가 조금이라도 있는 곳을 하나도 빠뜨리면 안 되고, 여러 봉우리를 다 덮으려다 그 사이의 빈 골까지 덮어 버립니다. 이것을 모드 커버링(mode-covering)이라 부릅니다.

역방향은 한 봉우리에 몰린다

역방향은 qAq_A 를 고릅니다. DKL(q ∥ p)D_{\mathrm{KL}}(q \,\|\, p) 에서는 qi=0q_i = 0 인 자리가 0log⁡0=00 \log 0 = 0 으로 통째로 사라지므로, 세 번째 칸을 버린 것이 아무 벌점도 받지 않습니다. qq 는 자기가 확률을 준 곳에서만 pp 와 맞으면 되므로 가장 편한 봉우리 하나에 몰리는 편이 이득입니다. 이것을 모드 시킹(mode-seeking)이라 부릅니다.

qBq_B 가 역방향에서 1.306으로 더 비싼 것도 같은 규칙입니다. qBq_B 는 가운데 칸에 1/31/3 을 주는데 pp 는 거기에 0.01밖에 안 줍니다. 그 한 항이 13log⁡233.3=1.68\tfrac{1}{3}\log_2 33.3 = 1.68 로 값의 대부분을 만듭니다. 정방향에서는 이 칸의 가중치가 pp 의 0.01이라 거의 안 보이던 자리인데, 방향이 바뀌면서 가중치가 1/31/3 로 올라가 지배적인 항이 되었습니다. 같은 칸이 방향에 따라 무시되기도 하고 값을 독차지하기도 합니다.

어느 쪽이 어디에 쓰이는가

이 갈래가 곧 익숙한 두 학습법의 갈래입니다.

정방향 DKL(목표 ∥ 내가 고르는 것)D_{\mathrm{KL}}(\text{목표} \,\|\, \text{내가 고르는 것}) 역방향 DKL(내가 고르는 것 ∥ 목표)D_{\mathrm{KL}}(\text{내가 고르는 것} \,\|\, \text{목표})
성격 모드 커버링 모드 시킹
0을 못 두는 쪽 뒤 앞
쓰이는 자리 최대가능도, 지도학습의 교차엔트로피, 지식 증류 변분추론, RLHF·DPO의 KL 페널티
실패하는 모습 흐릿한 평균, 아무도 안 쓰는 자리까지 덮는다 한 봉우리에 눌러앉아 다양성을 잃는다

최대가능도는 데이터 분포가 앞에 오므로 모델이 데이터의 모든 구석을 덮으려 합니다. 학습이 덜 된 생성 모델이 흐릿한 평균 같은 것을 내놓는 것이 그 그림자입니다. 변분추론과 RLHF의 KL 페널티는 우리가 고르는 분포가 앞에 오므로 참조 분포의 한 봉우리에 안정적으로 붙습니다.

어느 쪽이 좋은지는 정해져 있지 않습니다. 고를 때 물어야 할 것은 「빠뜨리는 것과 없는 것을 지어내는 것 중 어느 쪽이 더 나쁜가」입니다. 검색 후보를 만드는 모델이면 빠뜨리는 쪽이 치명적이라 정방향이고, 사용자에게 답 하나를 내놓는 정책이면 엉뚱한 답을 지어내는 쪽이 치명적이라 역방향입니다.

세 자리에 같은 양이 있다

표로 읽기

처음의 세 조각으로 돌아갑니다. 무엇을 앞에 두었는지만 확인하면 각 자리의 성격이 읽힙니다.

같은 식이 세 자리에서 다른 이름으로 불린다

쓰이는 자리 식 방향 하는 일
RLHF·DPO의 KL 페널티 β DKL(πθ ∥ πref)\beta\, D_{\mathrm{KL}}(\pi_\theta \,\|\, \pi_{\text{ref}}) 역방향 정책이 참조 모델의 봉우리에서 벗어나지 않게 붙든다
지식 증류 DKL(p교사 ∥ q학생)D_{\mathrm{KL}}(p_{\text{교사}} \,\|\, q_{\text{학생}}) 정방향 학생이 교사가 확률을 준 자리를 다 덮게 한다
ELBO의 간격 DKL(q(z) ∥ p(z∣x))D_{\mathrm{KL}}(q(z) \,\|\, p(z \mid x)) 역방향 하한과 참값의 거리가 정확히 이 값이다

가운데 줄이 「증류에서 왜 정답 하나가 아니라 교사의 분포 전체를 흉내 내는가」에 대한 답입니다. 교사가 오답들에 매긴 작은 확률까지 목표에 들어 있고, 정방향이라 학생은 그것들을 빠뜨릴 수 없습니다. 그 작은 확률들이 이른바 «어두운 지식」이고, 원-핫 정답에는 없는 정보입니다.

맨 아래 줄은 지난 글에서 미뤄 둔 답입니다. ELBO가 log⁡p(x)\log p(x) 에 못 미치는 그 간격이 바로 KL 발산이고, KL이 0이 되는 조건은 q(z)=p(z∣x)q(z) = p(z \mid x) 하나뿐이므로 하한이 참값에 닿는 것은 우리가 고른 분포가 참 사후분포와 같아질 때뿐입니다. 자세한 유도는 중급 60번 · 변분 하한과 ELBO가 맡습니다.

같은 양인데 왜 이름이 셋인가

셋이 같은 식이라면 이름이 왜 셋인가 — 여기에 답해 두지 않으면 표가 우연의 목록처럼 읽힙니다. 갈리는 것은 식이 아니라 그 식에서 무엇을 움직이느냐입니다.

  • 증류에서 움직이는 것은 뒤에 있는 학생입니다. 앞의 교사는 고정이라 H(p교사)H(p_{\text{교사}}) 가 상수이고, 그래서 이 항을 최소화하는 것과 교차엔트로피를 최소화하는 것이 같은 일입니다. 「손실」이라고 부르는 이유가 여기 있습니다.
  • RLHF에서 움직이는 것은 앞에 있는 정책입니다. 뒤의 참조 모델이 고정이라 이 항은 정책이 얼마나 멀리 갔는지를 재는 값이 되고, 보상에 더해지는 벌점으로 읽힙니다.
  • ELBO에서는 움직이는 것이 앞의 q(z)q(z) 인데 목적이 이 항을 0으로 만드는 것이 아니라 다른 양을 최대화하는 것이고, 이 항은 그 과정에서 남는 간격입니다. 재고는 있지만 직접 손대지는 않는 값입니다.

같은 자에 세 이름이 붙은 것이고, 이름은 그 자를 어디에 대고 있는지가 정합니다. 그래서 새 논문에서 KL 항을 만나면 물어야 할 것도 둘뿐입니다 — 앞에 오는 것이 무엇이고, 그중 무엇이 학습으로 움직이는가.

코드로 확인하기

세 성질을 수로

import numpy as np

def kl(p, q):                                  # 비트 단위, 0 log 0 = 0
    m = p > 0
    return float((p[m] * np.log2(p[m] / q[m])).sum())

p = np.array([0.9, 0.1]); q = np.array([0.5, 0.5])
print(round(kl(p, q), 4), round(kl(q, p), 4))          # 0.531 0.737  대칭이 아니다

# ① 삼각부등식 반례
a = np.array([0.5, 0.5]); b = np.array([0.9, 0.1]); c = np.array([0.99, 0.01])
print(round(kl(a, b) + kl(b, c), 4), round(kl(a, c), 4))   # 0.9454 2.3292

# ② 방향이 답을 뒤집는다
P  = np.array([0.495, 0.01, 0.495])
qA = np.array([0.99, 0.01, 0.0])                # 한 봉우리
qB = np.full(3, 1 / 3)                          # 넓게
print(kl(P, qA), round(kl(P, qB), 4))           # inf 0.5142   정방향은 qB
print(round(kl(qA, P), 4), round(kl(qB, P), 4)) # 0.99 1.306   역방향은 qA

# ③ JS 발산은 대칭이고 유한하다
def js(p, q):
    m = (p + q) / 2
    return 0.5 * kl(p, m) + 0.5 * kl(q, m)
print(round(js(p, q), 4), round(js(q, p), 4))   # 0.1468 0.1468
print(round(js(P, qA), 4))                      # 0.3082  KL은 무한이던 짝

②가 이 글의 핵심입니다. 같은 두 후보에 같은 식을 대는데 인자의 순서만 바꾸면 이기는 쪽이 뒤집힙니다.

분해가 실제로 맞는지

분해를 코드로도 한 번 확인해 둡니다. 세 값을 각각 따로 계산해 맞대는 것이라 오타를 잡기에 좋습니다.

def H(p):        return float(-(p[p > 0] * np.log2(p[p > 0])).sum())
def H_cross(p, q):
    m = p > 0
    return float(-(p[m] * np.log2(q[m])).sum())

p3 = np.array([0.5, 0.25, 0.25]); q3 = np.array([0.25, 0.25, 0.5])
print(round(H(p3), 4), round(H_cross(p3, q3), 4), round(kl(p3, q3), 4))
# 1.5 1.75 0.25        교차엔트로피 - 엔트로피 = KL

print(round(kl(q3, p3), 4))    # 0.25   이 짝은 양쪽이 우연히 같다

마지막 줄이 경계 하나를 짚어 줍니다. 대칭이 아니라는 것은 «언제나 다르다»는 뜻이 아닙니다. 이 짝처럼 두 값이 우연히 맞아떨어지는 경우도 있고, 그 한 예를 보고 「KL은 대칭이더라」로 결론 내리면 안 됩니다. 반례 하나가 비대칭을 증명하듯, 일치하는 예 하나는 대칭을 증명하지 못합니다.

연습 문제

연습 1 — 양쪽으로 재기

비트 단위로 계산합니다. log⁡23=1.585\log_2 3 = 1.585, log⁡25=2.322\log_2 5 = 2.322 를 써도 됩니다.

  1. p=(0.5, 0.25, 0.25)p = (0.5,\, 0.25,\, 0.25), q=(0.25, 0.25, 0.5)q = (0.25,\, 0.25,\, 0.5) 에서 DKL(p ∥ q)D_{\mathrm{KL}}(p \,\|\, q) 와 DKL(q ∥ p)D_{\mathrm{KL}}(q \,\|\, p) 를 각각 구하세요.
    앞쪽은 0.5log⁡22+0.25log⁡21+0.25log⁡212=0.5+0−0.25=0.250.5\log_2 2 + 0.25\log_2 1 + 0.25\log_2 \tfrac{1}{2} = 0.5 + 0 - 0.25 = 0.25 입니다. 뒤쪽은 0.25log⁡212+0.25log⁡21+0.5log⁡22=−0.25+0+0.5=0.250.25\log_2 \tfrac{1}{2} + 0.25\log_2 1 + 0.5\log_2 2 = -0.25 + 0 + 0.5 = 0.25 로 같은 값입니다. 두 분포가 첫 칸과 셋째 칸을 맞바꾼 꼴이라 항들이 그대로 자리만 바뀝니다.
  2. p=(0.6, 0.3, 0.1)p = (0.6,\, 0.3,\, 0.1), q=(0.2, 0.3, 0.5)q = (0.2,\, 0.3,\, 0.5) 에서 양쪽을 구하고 어느 쪽이 큰지 적으세요.
    DKL(p ∥ q)=0.6log⁡23+0+0.1log⁡20.2=0.951−0.232=0.719D_{\mathrm{KL}}(p \,\|\, q) = 0.6\log_2 3 + 0 + 0.1\log_2 0.2 = 0.951 - 0.232 = 0.719 이고, DKL(q ∥ p)=0.2log⁡213+0+0.5log⁡25=−0.317+1.161=0.844D_{\mathrm{KL}}(q \,\|\, p) = 0.2\log_2 \tfrac{1}{3} + 0 + 0.5\log_2 5 = -0.317 + 1.161 = 0.844 입니다. 뒤쪽이 큽니다. 가운데 칸은 두 분포가 같은 값을 주므로 양쪽 모두 0으로 빠집니다.
  3. p=(0.5, 0.5, 0)p = (0.5,\, 0.5,\, 0), q=(13, 13, 13)q = (\tfrac{1}{3},\, \tfrac{1}{3},\, \tfrac{1}{3}) 에서 양쪽을 구하세요.
    DKL(p ∥ q)=0.5log⁡21.5+0.5log⁡21.5=log⁡21.5=0.585D_{\mathrm{KL}}(p \,\|\, q) = 0.5\log_2 1.5 + 0.5\log_2 1.5 = \log_2 1.5 = 0.585 입니다. pp 의 셋째 항은 0log⁡0=00\log 0 = 0 으로 사라집니다. 반대로 DKL(q ∥ p)D_{\mathrm{KL}}(q \,\|\, p) 는 셋째 칸에서 q3=1/3>0q_3 = 1/3 > 0 인데 p3=0p_3 = 0 이라 무한대입니다.

연습 2 — 분해와 반례

  1. 1번의 pp 와 qq 로 H(p)H(p) 와 H(p,q)H(p, q) 를 각각 구하고, 그 차가 1번의 앞쪽 답과 같은지 확인하세요.
    H(p)=−(0.5log⁡20.5+0.25log⁡20.25+0.25log⁡20.25)=0.5+0.5+0.5=1.5H(p) = -(0.5\log_2 0.5 + 0.25\log_2 0.25 + 0.25\log_2 0.25) = 0.5 + 0.5 + 0.5 = 1.5 이고, H(p,q)=−(0.5log⁡20.25+0.25log⁡20.25+0.25log⁡20.5)=1+0.5+0.25=1.75H(p,q) = -(0.5\log_2 0.25 + 0.25\log_2 0.25 + 0.25\log_2 0.5) = 1 + 0.5 + 0.25 = 1.75 입니다. 차는 0.25로 1번의 앞쪽 답과 같습니다.
  2. 본문의 (0.5,0.5)(0.5, 0.5), (0.9,0.1)(0.9, 0.1), (0.99,0.01)(0.99, 0.01) 말고 다른 세 분포로 삼각부등식이 깨지는 짝을 하나 만들고, 세 값을 적으세요.
    한 예로 p=(0.5,0.5)p = (0.5, 0.5), q=(0.8,0.2)q = (0.8, 0.2), r=(0.98,0.02)r = (0.98, 0.02) 를 쓰면 DKL(p∥q)=0.322D_{\mathrm{KL}}(p\|q) = 0.322, DKL(q∥r)=0.430D_{\mathrm{KL}}(q\|r) = 0.430, 합이 0.752인데 DKL(p∥r)=1.837D_{\mathrm{KL}}(p\|r) = 1.837 입니다. 만드는 요령은 둘째 칸의 확률을 세 분포에서 계단처럼 크게 떨어뜨리는 것입니다.

연습 3 — JS 발산의 대칭

  1. DJS(p,q)=DJS(q,p)D_{\mathrm{JS}}(p, q) = D_{\mathrm{JS}}(q, p) 임을 정의에서 보이세요.
    m=12(p+q)m = \tfrac{1}{2}(p+q) 인데 pp 와 qq 를 바꾸면 12(q+p)\tfrac{1}{2}(q+p) 로 같은 분포입니다. 따라서 DJS(q,p)=12DKL(q∥m)+12DKL(p∥m)D_{\mathrm{JS}}(q, p) = \tfrac{1}{2}D_{\mathrm{KL}}(q\|m) + \tfrac{1}{2}D_{\mathrm{KL}}(p\|m) 이고, 이는 원래 식의 두 항을 순서만 바꿔 더한 것이라 값이 같습니다. 대칭이 정의에서 바로 나온다는 점이 핵심이고, KL 자체의 성질은 하나도 쓰이지 않습니다.

정리

  • KL 발산은 ∑ipilog⁡(pi/qi)\sum_i p_i \log(p_i/q_i) 이고, 곧 H(p,q)−H(p)H(p,q) - H(p) 다. 「qq 를 믿고 적을 때 치르는 초과 비용」이다.
  • 교차엔트로피 = 엔트로피 + KL 발산. pp 가 고정이면 둘을 최소화하는 자리가 같아서 지도학습에서는 구별할 필요가 없다. 다만 값 자체는 H(p)H(p) 만큼 다르므로 손실의 절댓값을 «정답에서 떨어진 거리»로 읽으면 안 된다.
  • DKL≥0D_{\mathrm{KL}} \ge 0 이고 등호는 q=pq = p 일 때만이다. 젠센 부등식에서 나온 결과다.
  • DKL(p∥q)D_{\mathrm{KL}}(p\|q) 가 유한하려면 qq 의 지지집합이 pp 의 지지집합을 덮어야 한다. 빈도 추정에 ε\varepsilon 스무딩을 걸고 log_softmax를 쓰는 관행이 이 조건에서 나온다.
  • 거리가 아니다. 대칭이 아니고 삼각부등식도 어긴다 — 돌아가는 길이 곧장 가는 길의 절반도 안 되는 반례가 있다. 다만 우연히 양쪽 값이 같아지는 짝도 있어, 한 예로 대칭을 결론 내면 안 된다.
  • 정방향은 모드 커버링, 역방향은 모드 시킹이다. 고를 때 물을 것은 빠뜨리는 것과 지어내는 것 중 어느 쪽이 더 나쁜가 하나다.
  • JS 발산은 가운뎃점 mm 을 만들어 양쪽에서 잰 값이라 대칭이고 log⁡2\log 2 이하로 유한하며, 제곱근은 진짜 거리다.

init_kl_coef: 0.05로 돌아갑니다. 이제 저 줄은 「참조 모델에서 멀어지지 마라」보다 정확하게 읽힙니다 — 역방향으로 재고 있으므로, 정책은 참조 모델의 한 봉우리에 눌러앉는 쪽으로 눌립니다. 정책이 다양성을 잃는다는 흔한 관찰이 계수의 크기 문제만이 아니라 방향의 성질이기도 하다는 뜻입니다.

5단원의 마지막 한 편이 남았습니다. 지금까지는 분포 하나 또는 둘을 다뤘는데, 다음 글은 두 확률변수가 얼마나 얽혀 있는가를 재는 양으로 갑니다. 그 양도 결국 KL 발산 하나로 적힙니다.


읽어주셔서 감사합니다. 😊

LATEST

수학의 최신 글

수학2026.09.07

양자화 오차: 격자 사상, 오차 분산, 이상치 채널

실수를 2^b개 격자에 사상할 때 오차의 분산이 왜 Δ²/12인지 유도하고, 그것이 비트당 6.02dB라는 SNR로 번역되는 과정을 실측과 대조했습니다. 이상치 하나가 나머지 값의 유효 비트를 어떻게 먹는지, 그리고 int4에서 성능이 무너지는 지점을 오차 예산으로 미리 계산하는 법까지.

중급18 MIN
수학2026.09.07

수치적으로 안정한 계산 패턴 모음

최댓값 빼기, 로그 공간, log1p·expm1, 분산의 두 공식, 정규화의 ε, fp32 누산, 역행렬 대신 solve — 프레임워크가 몰래 해 주는 일곱 가지를 하나씩 꺼내 각각 어떤 고장을 막는지 직접 재 봤습니다. 수식을 그대로 옮긴 코드가 왜 라이브러리보다 나쁜지에 대한 목록입니다.

중급22 MIN
수학2026.09.07

부동소수점은 어디서 새는가: 반올림, 상쇄, 더하는 순서

0.1 + 0.2가 0.3이 아닌 이유부터 시작해 머신 엡실론을 유도하고, 같은 16비트인데 fp16과 bf16이 서로 다른 지점에서 터지는 이유, 비슷한 수를 뺄 때 유효자리가 사라지는 파괴적 상쇄, 그리고 1,000만 개를 순서만 바꿔 더했을 때 오차가 백만 배 갈리는 실험까지 직접 재 봤습니다.

중급23 MIN