f-발산의 하한: 밀도 비를 켤레 뒤로 숨기기
이 장을 열며 던진 물음으로 돌아가자. 손에 있는 것이 두 분포의 샘플뿐일 때, 두 분포가 얼마나 다른지를 어떻게 잴까? 식대로 계산하려고 하면 매번 두 밀도의 비가 걸림돌이 된다.
결과가 둘뿐인 분포의 KL
가장 작은 예부터 보자. 결과가 둘뿐인 두 분포 p = (0.8, 0.2)와 q = (0.5, 0.5)의 KL은 0.8 ln(0.8/0.5) + 0.2 ln(0.2/0.5) = 0.8 × 0.4700 + 0.2 × (−0.9163) = 0.1927이다. 결과가 둘뿐이라 손으로 끝까지 셀 수 있어서 고른 예다. 계산의 모든 항에 두 확률의 비 0.8/0.5와 0.2/0.5가 들어 있다. 확률은 모르고 샘플만 가진 사람이 모르는 것이 바로 이 비다.
발산을 샘플 평균으로 바꾸기
KL만 그런 것이 아니다. 데이터 분포 p와 모델 분포 q가 얼마나 다른지 재는 발산은 대부분 다음 모양으로 쓸 수 있다. f는 f(1) = 0인 볼록 함수이고, 이렇게 만든 발산을 f-발산(f-divergence)이라 한다. f(u) = u ln u를 고르면 KL 발산이 된다.
이 식을 샘플로 계산할 수 없는 이유는 f 안에 밀도의 비 pᵢ/qᵢ가 들어 있기 때문이다. 여기서 르장드르 변환이 등장한다. f는 볼록이므로 두 번 변환하면 제자리, 곧 f(u) = sup_z (zu − f*(z))로 쓸 수 있다. 이것을 각 i에 따로 적용하되 z를 i마다 다르게 고를 수 있게 하면, qᵢ가 분모를 지워 다음 부등식이 나온다. 이 한 줄이 왜 성립하는지는 아래 문제 10에서 결과마다 따라가 본다.
오른쪽에는 밀도가 없다. 데이터 샘플에서 z의 평균을, 모델 샘플에서 f*(z)의 평균을 내면 된다. 여기서 샘플마다 점수를 매기는 함수 z(x)를 비평 함수 (두 분포를 가르는 점수 함수, critic)라 부른다. 어떤 z를 골라도 펜헬–영 부등식 때문에 발산의 하한이 되고, 등호는 z(x)가 정확히 f′(p/q)일 때 성립한다. KL이면 f*(z) = e^(z − 1)이므로 하한은 ⟨z⟩_p − ⟨e^(z − 1)⟩_q가 되고, 최적의 비평 함수는 1 + ln(p/q), 곧 밀도 비의 로그다. 이 하한은 Nguyen, Wainwright, Jordan(2010)이 발산 추정에 쓴 것이어서 세 사람 이름의 앞글자를 따 NWJ 하한이라 부른다.
ML에서: f-GAN
Nowozin, Cseke, Tomioka(2016)의 f-GAN은 이 부등식을 그대로 학습 규칙으로 썼다. 비평 함수는 오른쪽을 키워 발산을 최대한 정확히 재려 하고, 생성기는 모델 분포 q를 움직여 그 값을 줄이려 한다. 최초의 GAN이 판별기와 생성기의 게임이었던 것도 f를 특정하게 고른 경우라는 것이 이 논문이 보인 핵심 가운데 하나다. 판별기의 로짓을 비평 함수라 보면, 판별기가 배우는 것은 결국 밀도 비의 로그다.
문제 10. 분모를 지우는 한 수
결과가 둘뿐인 두 분포 p = (0.9, 0.1), q = (0.6, 0.4)가 있다. (가) KL의 하한 ⟨z⟩_p − ⟨e^(z − 1)⟩_q에서 z를 결과마다 따로 고를 수 있을 때, 이 값을 가장 크게 하는 z₁, z₂와 그때의 값을 구하라. (나) 그 값이 KL과 같음을 확인하고, 왜 q가 분모에서 사라지는지 설명하라.

아무거나 넣어 볼게요. z = (1, 1)이면 0.9 + 0.1 − (0.6 + 0.4) × e⁰ = 0. 하한이 0이면 쓸모가 없는데요?

그 식을 결과마다 따로 적어 봐요. 결과 i의 몫은 무엇이죠?

pᵢzᵢ − qᵢe^(zᵢ − 1)이에요. 결과마다 따로 떨어져 있으니까 결과마다 최대로 만들면 되겠네요. zᵢ로 미분해 0으로 두면 zᵢ = 1 + ln(pᵢ/qᵢ).

그럼 z = (1.4055, −0.3863)이고 값은 0.2263이에요. KL을 직접 계산하면 0.9 ln 1.5 + 0.1 ln 0.25 = 0.3649 − 0.1386 = 0.2263. 딱 같아요!

결과 i의 몫에서 qᵢ를 앞으로 묶어 내면 qᵢ × (zᵢ × (pᵢ/qᵢ) − e^(zᵢ − 1))이에요. 괄호 안은 어디서 본 모양이죠?

u = pᵢ/qᵢ로 두면 zu − f*(z)예요. f(u) = u ln u의 켤레가 e^(z − 1)이니까, 괄호 안의 최댓값은 f(u)고요. 그걸 qᵢ를 곱해 더하면 Σ qᵢf(pᵢ/qᵢ), 곧 f-발산 식 그대로네요. 바깥에 곱해진 qᵢ가 pᵢ/qᵢ의 분모를 지운 거였어요.

그러면 샘플로 할 때는 Σ pᵢzᵢ를 데이터 샘플의 z 평균으로, Σ qᵢe^(zᵢ − 1)을 모델 샘플의 평균으로 바꾸면 되니까 확률을 몰라도 되고요.

그런데 신경망이 모든 결과에서 가장 좋은 z를 고르지 못하면요?

결과마다 괄호 안이 f(u)보다 작아지니까 합도 KL보다 작아요. 그래서 등호가 아니라 하한이에요.

조별 과제 점수가 조원마다 받은 점수의 합이면, 한 명이라도 대충 하면 딱 그만큼 합이 깎이는 거랑 같네요.