수학

MATH / 중급 29번

젠센 부등식: 볼록함수와 기댓값의 순서

확산 모델과 VAE의 손실은 원래 목표가 아니라 그 하한입니다. 왜 하한을 대신 올려도 되는가 — 그 근거는 부등호 하나입니다. 볼록함수를 정의하고 접선 하나로 젠센 부등식을 증명한 뒤, 지난 글에서 미뤄 둔 «교차엔트로피의 초과분은 음수가 될 수 없다»를 갚습니다.

PALDYN Team37 MIN READ

확산 모델이나 VAE의 학습 코드를 열면 손실이 이렇게 생겼습니다.

loss = -(recon_term - kl_term)      # = -ELBO

ELBO는 evidence lower bound, 우리말로 증거 하한입니다. 이름에 「하한」이 박혀 있습니다. 모델이 진짜로 키우고 싶은 것은 데이터의 로그가능도 log⁡p(x)\log p(x) 인데, 그 값은 숨은 변수에 대한 적분이라 계산할 수 없습니다. 그래서 계산할 수 있는 다른 값을 대신 최대화합니다.

이상하게 들리는 대목이 여기입니다. 목표를 대신할 값을 아무거나 골라서는 안 될 텐데, 무엇이 이 바꿔치기를 정당하게 만드는가. 답은 부등호 하나이고, 그 부등호의 이름이 젠센 부등식입니다. 지난 글에서 「H(p,q)−H(p)H(p,q) - H(p) 는 음수가 될 수 없다」고 적어 놓고 증명을 미뤘는데, 그것도 같은 부등호에서 나옵니다. 이 글에서 도구 하나를 세우고 빚 둘을 갚습니다.

볼록함수의 정의

현과 나누는 비율

두 점을 이은 선분을 현(chord)이라고 합니다. 볼록성의 정의는 그 현의 위치에 대한 이야기입니다.

정의. 함수 ff 가 볼록(convex)하다는 것은 정의역의 임의의 두 점 aa, bb 와 0≤λ≤10 \le \lambda \le 1 에 대해 f(λa+(1−λ)b)  ≤  λf(a)+(1−λ)f(b)f(\lambda a + (1-\lambda) b) \;\le\; \lambda f(a) + (1-\lambda) f(b) 가 성립한다는 뜻이다.

왼쪽은 두 점 사이의 어떤 자리에서 잰 함숫값이고, 오른쪽은 같은 자리에서 잰 현의 높이입니다. 부등호는 「그래프가 현보다 아래에 있다」를 적은 것입니다.

볼록함수는 현이 언제나 그래프 위에 있다

λ\lambda 와 1−λ1-\lambda 는 더해서 1인 두 수이므로, 두 점 사이를 나눈 비율로 읽는 것이 가장 정확합니다. λ=1\lambda = 1 이면 aa 자리이고 λ=0\lambda = 0 이면 bb 자리이며, 그 사이는 aa 쪽으로 λ\lambda 만큼 치우친 자리입니다. f(x)=x2f(x) = x^2 과 a=1a = 1, b=4b = 4 에서 λ=1/3\lambda = 1/3 으로 두고 양변을 실제로 재 봅니다.

λa+(1−λ)b=13(1)+23(4)=3⟹f(3)=9\lambda a + (1-\lambda) b = \tfrac{1}{3}(1) + \tfrac{2}{3}(4) = 3 \quad \Longrightarrow \quad f(3) = 9

λf(a)+(1−λ)f(b)=13(1)+23(16)=333=11\lambda f(a) + (1-\lambda) f(b) = \tfrac{1}{3}(1) + \tfrac{2}{3}(16) = \tfrac{33}{3} = 11

9≤119 \le 11 입니다. 왼쪽은 x=3x = 3 에서 포물선의 높이이고 오른쪽은 같은 x=3x = 3 에서 현의 높이이며, 그 차이 2가 두 값이 벌어진 정도입니다. 부등호가 반대이면 오목(concave)하다고 합니다. ff 가 오목한 것과 −f-f 가 볼록한 것은 같은 말이라, 오목함수를 따로 다룰 필요는 없습니다. 부호만 뒤집으면 됩니다.

그래프 위쪽 영역

정의를 도형의 말로 한 번 더 옮겨 적으면 나중에 판정이 쉬워집니다. 그래프와 그 위쪽을 통째로 모은 집합

{(x,t):t≥f(x)}\{(x, t) : t \ge f(x)\}

을 생각합니다. 어떤 집합 안의 두 점을 골라 이은 선분이 언제나 그 집합 안에 들어 있으면 그 집합을 볼록집합(convex set)이라고 하는데, 방금 모은 영역이 볼록집합인 것과 ff 가 볼록한 것은 같은 말입니다.

확인은 정의를 그대로 읽는 것으로 끝납니다. 영역 안의 두 점 (a,f(a))(a, f(a)) 와 (b,f(b))(b, f(b)) 를 이은 선분의 높이가 λf(a)+(1−λ)f(b)\lambda f(a) + (1-\lambda) f(b) 이고, 그 자리의 그래프 높이가 f(λa+(1−λ)b)f(\lambda a + (1-\lambda) b) 입니다. 선분이 영역 안에 있다는 말은 선분의 높이가 그래프보다 위라는 말이고, 그것이 곧 볼록성의 부등호입니다. 「현이 그래프 위에 있다」와 「위쪽 영역이 볼록집합이다」는 서로 다른 사실이 아니라 같은 문장의 두 표기입니다.

엄격 볼록과 그냥 볼록

λ\lambda 를 0과 1로 놓으면 양변이 같아지니 등호는 늘 가능합니다. 눈여겨볼 자리는 0<λ<10 < \lambda < 1 에서도 등호가 서는가입니다. 서로 다른 두 점과 그 사이의 모든 λ\lambda 에서 부등호가 반드시 강부등호이면 엄격 볼록(strictly convex)하다고 합니다.

직선은 이 구별이 왜 필요한지 보여 줍니다. f(x)=2x+1f(x) = 2x + 1 은 현이 그래프와 정확히 겹치므로 모든 λ\lambda 에서 등호이고, 그래서 볼록이면서 동시에 오목이지만 엄격 볼록은 아닙니다. 반대로 x2x^2 은 위에서 봤듯 λ=1/3\lambda = 1/3 에서 9와 11로 갈리므로 엄격 볼록입니다. 직선은 볼록 쪽에도 오목 쪽에도 들지만 엄격한 쪽에는 어디에도 들지 않는다 — 이 문장이 뒤에서 등호 조건을 말할 때 그대로 쓰입니다.

볼록성의 판정과 조립

이계도함수 판정

정의를 그대로 확인하는 것은 번거롭습니다. 두 번 미분되는 함수라면 훨씬 쉬운 판정법이 있습니다.

판정. ff 가 두 번 미분 가능하면, ff 가 볼록한 것과 모든 점에서 f′′(x)≥0f''(x) \ge 0 인 것은 같은 말이다.

f′′f'' 은 기울기의 변화율이니, f′′≥0f'' \ge 0 은 기울기가 줄지 않는다는 뜻입니다. 왼쪽에서 완만하게 시작해 오른쪽으로 갈수록 가팔라지는 모양이고, 그런 곡선은 현 아래에 놓입니다. AI에서 자주 만나는 함수들을 이 판정으로 갈라 봅니다.

함수 f′′(x)f''(x) 갈래
x2x^2 22 볼록
exe^x ex>0e^x > 0 볼록
−log⁡x-\log x 1/x2>01/x^2 > 0 볼록
log⁡x\log x −1/x2<0-1/x^2 < 0 오목
1/x  (x>0)1/x \;(x>0) 2/x3>02/x^3 > 0 볼록
∥x∥\|x\| 0에서 미분 불가 볼록 (정의로 확인)

마지막 줄이 판정법의 한계를 보여 줍니다. 볼록성은 미분 가능성을 요구하지 않습니다. ∣x∣|x| 는 0에서 꺾이지만 현은 여전히 그래프 위에 있습니다. 판정법은 미분되는 함수에서 쓰는 지름길일 뿐입니다.

세 번째 줄에 밑줄을 그어 둡니다. −log⁡-\log 가 볼록이라는 것, 뒤집어 말해 log⁡\log 가 오목이라는 것이 이 글에서 실제로 쓰는 사실의 거의 전부입니다.

볼록을 보존하는 연산

판정보다 자주 쓰는 것은 조립입니다. 이미 볼록이라고 아는 함수들을 몇 가지 방식으로 이어 붙이면 결과도 볼록이라는 규칙이고, 미분을 한 번도 하지 않고 결론이 납니다.

조립 결과 근거
f+gf + g 볼록 부등호 둘을 변끼리 더한다
cfcf, c≥0c \ge 0 볼록 부등호에 양수를 곱해도 방향이 그대로다
max⁡(f,g)\max(f, g) 볼록 각각의 현 아래에 있으니 둘 중 높은 현 아래에도 있다

세 번째 줄이 ReLU를 설명합니다. ReLU(x)=max⁡(0,x)\mathrm{ReLU}(x) = \max(0, x) 는 상수함수 0과 항등함수 xx 라는 직선 둘의 최댓값이고, 직선은 볼록이므로 ReLU도 볼록입니다. 0에서 꺾여 미분이 안 되는 점 때문에 판정법은 쓸 수 없지만, 조립 규칙은 그 사실을 아예 묻지 않습니다. 힌지 손실 max⁡(0,1−yz)\max(0, 1 - yz) 도 같은 이유로 볼록합니다.

아핀 합성과 헤세 행렬

벡터를 받는 함수로 넘어가는 다리가 둘 있습니다. 하나는 합성입니다. AA 가 행렬이고 bb 가 벡터일 때 x↦Ax+bx \mapsto Ax + b 처럼 곱하고 더하기만 하는 사상을 아핀사상(affine map)이라고 하는데, 볼록함수 앞에 아핀사상을 끼워 넣은 g(x)=f(Ax+b)g(x) = f(Ax + b) 도 볼록입니다. 아핀사상은 나누는 비율을 그대로 옮겨 주므로 — A(λx1+(1−λ)x2)+b=λ(Ax1+b)+(1−λ)(Ax2+b)A(\lambda x_1 + (1-\lambda)x_2) + b = \lambda(Ax_1 + b) + (1-\lambda)(Ax_2 + b) — 정의의 왼쪽과 오른쪽이 그대로 따라옵니다.

다른 하나는 이계도함수의 자리를 대신하는 행렬입니다. 여러 변수 함수의 이계 편미분 ∂2f/∂xi∂xj\partial^2 f / \partial x_i \partial x_j 를 (i,j)(i, j) 칸에 모은 행렬을 헤세 행렬(Hessian)이라고 하고, 두 번 미분되는 여러 변수 함수가 볼록한 것과 헤세 행렬이 모든 점에서 반양정치인 것 — 즉 모든 벡터 vv 에 대해 v⊤Hv≥0v^\top H v \ge 0 인 것 — 이 같은 말입니다. 한 변수에서 f′′≥0f'' \ge 0 이던 조건의 그대로의 확장입니다.

logsumexp

이 규칙들을 한 번에 쓰는 자리가 분류 손실입니다. 로짓 벡터 zz 에 대해

LSE(z)=log⁡∑iezi\mathrm{LSE}(z) = \log \sum_i e^{z_i}

를 logsumexp라고 하고, 이 함수는 볼록합니다. 교차엔트로피 손실은 정답이 yy 번일 때 LSE(z)−zy\mathrm{LSE}(z) - z_y 로 적히는데, 뒤의 항이 zz 에 대한 일차식이라 빼도 볼록성이 유지됩니다. 즉 softmax 교차엔트로피는 로짓에 대해 볼록합니다.

그런데 학습이 실제로 움직이는 것은 로짓이 아니라 파라미터입니다. 로지스틱 회귀처럼 로짓이 z=Wx+bz = Wx + b 로 파라미터의 아핀함수이면 위의 아핀 합성 규칙이 걸려 손실은 파라미터에 대해서도 볼록하고, 그래서 국소 최솟값이 곧 전역 최솟값입니다. 반면 로짓 앞에 신경망 몇 층을 태우면 zz 가 파라미터의 아핀함수가 아니게 되어 그 규칙이 끊어집니다. 손실의 볼록성은 마지막 한 층에서만 남아 있는 성질이고, 딥러닝이 국소 최솟값을 걱정하는 이유가 정확히 이 자리입니다.

젠센 부등식

접선 하나로 끝나는 증명

정의의 λ\lambda 와 1−λ1-\lambda 를 확률로 읽으면 「aa 가 나올 확률 λ\lambda, bb 가 나올 확률 1−λ1-\lambda 인 확률변수 XX」이고, 그러면 λa+(1−λ)b\lambda a + (1-\lambda) b 는 E[X]\mathbb{E}[X] 이고 λf(a)+(1−λ)f(b)\lambda f(a) + (1-\lambda) f(b) 는 E[f(X)]\mathbb{E}[f(X)] 입니다. 정의를 확률의 말로 옮겨 적기만 해도 부등식이 됩니다.

젠센 부등식. ff 가 볼록하면 임의의 확률변수 XX 에 대해 f(E[X])  ≤  E[f(X)]f\big(\mathbb{E}[X]\big) \;\le\; \mathbb{E}\big[f(X)\big] 이다. ff 가 오목하면 부등호가 뒤집힌다.

μ=E[X]\mu = \mathbb{E}[X] 라 두고 μ\mu 에서 그은 접선을 봅니다. 볼록함수는 어느 점에서 그은 접선보다 아래로 내려가지 않습니다.

f(x)  ≥  f(μ)+f′(μ)(x−μ)모든 x 에 대해f(x) \;\ge\; f(\mu) + f'(\mu)(x - \mu) \qquad \text{모든 } x \text{ 에 대해}

볼록함수는 접선 위에 있고, 기댓값을 씌우면 젠센 부등식이 나온다

이 자체도 한 줄로 확인됩니다. g(x)=f(x)−f(μ)−f′(μ)(x−μ)g(x) = f(x) - f(\mu) - f'(\mu)(x-\mu) 로 두면 g(μ)=0g(\mu) = 0 이고 g′(x)=f′(x)−f′(μ)g'(x) = f'(x) - f'(\mu) 인데, f′′≥0f'' \ge 0 이라 f′f' 이 증가함수이므로 g′g' 은 x<μx < \mu 에서 음수, x>μx > \mu 에서 양수입니다. 즉 gg 는 μ\mu 에서 최솟값 0을 가지므로 g≥0g \ge 0 입니다.

이제 양변에 기댓값을 씌웁니다. 기댓값은 부등호의 방향을 지키고, 오른쪽은 xx 에 대해 일차식이라 그대로 계산됩니다.

E[f(X)]  ≥  f(μ)+f′(μ)(E[X]−μ)=f(μ)+0=f(E[X])\mathbb{E}[f(X)] \;\ge\; f(\mu) + f'(\mu)\big(\mathbb{E}[X] - \mu\big) = f(\mu) + 0 = f\big(\mathbb{E}[X]\big)

끝입니다. 오른쪽 항이 0이 되는 것이 증명의 전부이고, 그것은 접선을 하필 μ\mu 에서 그었기 때문입니다.

등호 조건도 여기서 읽힙니다. 등호가 서려면 XX 가 실제로 값을 갖는 모든 자리에서 g(x)=0g(x) = 0 이어야 합니다. 즉 그 구간에서 ff 가 접선과 겹쳐 직선이거나, 아니면 XX 가 μ\mu 하나만 갖는 상수여야 합니다. 앞 절에서 갈라 둔 엄격 볼록이 여기서 값을 합니다 — ff 가 엄격 볼록이면 어떤 구간에서도 직선일 수 없으므로 남는 것은 XX 가 상수인 경우 하나뿐입니다.

값이 n개일 때와 연속 분포

정의는 값이 둘일 때의 이야기였는데 위의 증명에는 값이 몇 개인지가 한 번도 들어가지 않았습니다. 쓴 것은 「기댓값이 부등호를 지킨다」와 「기댓값은 일차식을 그대로 통과시킨다」 둘뿐입니다. 그래서 값이 nn 개인 이산 분포든 밀도를 갖는 연속 분포든 같은 세 줄이 그대로 돕니다.

조건을 굳이 적자면 가중치가 확률이기만 하면 된다는 것입니다. 음이 아니고 합이 1 — 이 둘이 전부입니다. 실제로 E[X]\mathbb{E}[X] 가 두 점 사이가 아니라 여러 점의 볼록결합이 되고, 그 점이 정의역 안에 있다는 것만 보장되면 μ\mu 에서 접선을 그을 수 있습니다. 가중치의 합이 1이 아닌 경우에 부등식이 깨지는 것도 같은 자리에서 보입니다. 접선 항의 f(μ)f(\mu) 앞에 합이 곱해져 나오기 때문입니다.

산술평균과 기하평균

XX 가 1과 100을 반반의 확률로 갖는다고 합시다. log⁡\log 는 오목하니 부등호가 뒤집혀 log⁡E[X]≥E[log⁡X]\log \mathbb{E}[X] \ge \mathbb{E}[\log X] 여야 합니다.

E[X]=1+1002=50.5⟹log⁡E[X]=ln⁡50.5=3.922\mathbb{E}[X] = \frac{1 + 100}{2} = 50.5 \quad \Longrightarrow \quad \log \mathbb{E}[X] = \ln 50.5 = 3.922

E[log⁡X]=ln⁡1+ln⁡1002=0+4.6052=2.303\mathbb{E}[\log X] = \frac{\ln 1 + \ln 100}{2} = \frac{0 + 4.605}{2} = 2.303

로그는 오목해서 현이 그래프 아래에 놓인다

3.922≥2.3033.922 \ge 2.303 이고 차이가 1.619나 됩니다. 두 번째 값을 지수로 되돌려 보면 정체가 드러납니다.

exp⁡(E[log⁡X])=exp⁡(ln⁡1+ln⁡1002)=1×100=10\exp\big(\mathbb{E}[\log X]\big) = \exp\left(\frac{\ln 1 + \ln 100}{2}\right) = \sqrt{1 \times 100} = 10

기하평균입니다. 즉 log⁡E[X]≥E[log⁡X]\log \mathbb{E}[X] \ge \mathbb{E}[\log X] 는 「산술평균 50.5가 기하평균 10보다 크거나 같다」와 같은 말이고, 중학교에서 배우는 산술-기하평균 부등식이 젠센 부등식의 한 조각이었던 셈입니다.

젠센 갭의 2차 어림

부등식의 양변 차이를 젠센 갭이라고 부릅니다. 갭이 얼마나 벌어지는지를 어림하는 식이 하나 있습니다. ff 를 μ\mu 근처에서 2차까지 전개해 기댓값을 씌우면 일차항이 사라지고 이차항만 남습니다.

E[f(X)]−f(E[X])  ≈  12f′′(μ) Var(X)\mathbb{E}[f(X)] - f(\mathbb{E}[X]) \;\approx\; \tfrac{1}{2} f''(\mu)\, \mathrm{Var}(X)

읽는 법은 두 가지입니다. 굽은 정도와 퍼진 정도의 곱이라는 것, 그리고 둘 중 하나라도 0이면 갭이 0이라는 것입니다. ff 가 직선이면 f′′=0f'' = 0 이고 XX 가 상수이면 분산이 0이니, 앞에서 얻은 등호 조건 둘이 이 어림식 안에 그대로 들어 있습니다.

다만 이것은 어림이고, 어디까지 믿을 수 있는지가 실전에서는 더 중요합니다. 평균 50.5를 고정한 채 두 점을 벌려 가며 log⁡\log 의 갭과 어림을 나란히 재 봅니다.

두 값 분산 실제 갭 2차 어림
49, 52 2.25 0.000441 0.000441
40, 61 110.25 0.0221 0.0216
20, 81 930.25 0.2269 0.1824
1, 100 2450.25 1.6194 0.4804

분산이 커지면 젠센 갭이 2차 어림을 앞지른다

첫 줄에서는 소수 여섯째 자리까지 맞고 마지막 줄에서는 세 배 넘게 어긋납니다. 경계가 어디인지는 어림의 출처를 보면 읽힙니다. 2차까지만 쓴 전개라, XX 가 μ\mu 에서 멀리 흩어져 3차 이상 항이 무시되지 않는 순간 무너집니다. 위 표에서는 분산이 100 언저리까지가 그 경계이고, 그 위로는 어림이 실제 갭을 밑으로 빗나갑니다 — −log⁡-\log 의 3차 도함수가 갭을 더 키우는 쪽으로 붙기 때문입니다. 갭을 어림으로 대신 세우려 할 때는 분산이 평균의 제곱에 비해 작은 구간인지 먼저 재 보아야 합니다.

교차엔트로피 초과분의 부호

젠센으로 가는 길

지난 글에서 H(p,q)≥H(p)H(p, q) \ge H(p) 를 세 예로 확인만 하고 증명을 미뤘습니다. 이제 갚습니다. 초과분을 한 덩어리로 묶는 것부터 시작합니다.

H(p,q)−H(p)=−∑ipilog⁡qi+∑ipilog⁡pi=∑ipilog⁡piqiH(p,q) - H(p) = -\sum_i p_i \log q_i + \sum_i p_i \log p_i = \sum_i p_i \log \frac{p_i}{q_i}

부호를 밖으로 빼면 기댓값 꼴이 보입니다.

H(p,q)−H(p)=−∑ipilog⁡qipi=− Ei∼p ⁣[log⁡qipi]H(p,q) - H(p) = -\sum_i p_i \log \frac{q_i}{p_i} = -\,\mathbb{E}_{i \sim p}\!\left[\log \frac{q_i}{p_i}\right]

대괄호 안이 log⁡(어떤 확률변수)\log(\text{어떤 확률변수}) 이고 log⁡\log 는 오목하니, 젠센을 그대로 적용합니다.

Ei∼p ⁣[log⁡qipi]  ≤  log⁡ Ei∼p ⁣[qipi]=log⁡∑ipiqipi=log⁡∑iqi=log⁡1=0\mathbb{E}_{i \sim p}\!\left[\log \frac{q_i}{p_i}\right] \;\le\; \log\, \mathbb{E}_{i \sim p}\!\left[\frac{q_i}{p_i}\right] = \log \sum_i p_i \frac{q_i}{p_i} = \log \sum_i q_i = \log 1 = 0

pip_i 가 약분되어 qq 의 총합만 남고, 그것은 1입니다. 앞에 마이너스가 붙어 있었으니 부등호가 뒤집혀

H(p,q)−H(p)  ≥  0H(p,q) - H(p) \;\ge\; 0

입니다. 등호 조건도 함께 옵니다. log⁡\log 는 엄격 오목이므로 등호는 확률변수 qi/piq_i/p_i 가 상수일 때만 서는데, pp 와 qq 가 둘 다 합이 1이라 그 상수는 1일 수밖에 없습니다. 즉 q=pq = p 입니다.

접선 부등식으로 가는 길

같은 결론을 젠센을 부르지 않고 얻는 길이 있습니다. 출발은 부등식 하나입니다.

log⁡t  ≤  t−1(t>0)\log t \;\le\; t - 1 \qquad (t > 0)

h(t)=t−1−log⁡th(t) = t - 1 - \log t 로 두면 h(1)=0h(1) = 0 이고 h′(t)=1−1/th'(t) = 1 - 1/t 라 t=1t = 1 에서 부호가 바뀌므로, hh 는 t=1t = 1 에서 최솟값 0을 갖습니다. 이 부등식을 t=qi/pit = q_i / p_i 에 넣고 pip_i 를 곱해 더하면 끝입니다.

∑ipilog⁡qipi  ≤  ∑ipi(qipi−1)=∑iqi−∑ipi=1−1=0\sum_i p_i \log \frac{q_i}{p_i} \;\le\; \sum_i p_i \left(\frac{q_i}{p_i} - 1\right) = \sum_i q_i - \sum_i p_i = 1 - 1 = 0

앞 소절과 같은 결론에 같은 등호 조건(qi/piq_i / p_i 가 모두 1)입니다. 이 길에는 기댓값을 씌우는 단계도, 오목성을 부르는 단계도 없습니다. 항마다 부등식을 걸고 더했을 뿐입니다.

두 길이 같은 접선인 것

두 증명이 닮은 것은 우연이 아닙니다. 젠센 증명에서 쓴 도구는 μ\mu 에서 그은 접선이었고, 지금 경우 X=qi/piX = q_i/p_i 의 기댓값은 ∑iqi=1\sum_i q_i = 1 이므로 μ=1\mu = 1 입니다. f=log⁡f = \log 의 t=1t = 1 에서의 접선을 그대로 적으면

f(μ)+f′(μ)(t−μ)=log⁡1+1⋅(t−1)=t−1f(\mu) + f'(\mu)(t - \mu) = \log 1 + 1 \cdot (t - 1) = t - 1

이고, 이것이 방금 쓴 접선 부등식의 오른쪽입니다. 두 증명은 같은 접선을 쓰고 있었고, 다른 것은 접선을 언제 꺼내느냐뿐입니다. 젠센은 접선을 먼저 일반 정리로 만들어 두고 나중에 불러 쓰고, 접선 부등식은 필요한 그 점에서 바로 꺼내 씁니다.

그래서 「모델 분포를 정답 분포에 맞추는 것이 손실을 줄이는 유일한 길」이라는 주장이 이제 증명된 사실이 되었습니다. 이 초과분에 이름이 있고, 그 이름과 성질은 다음 글이 맡습니다.

하한을 대신 올리는 논법

로그를 기댓값 안으로

처음의 ELBO로 돌아갑니다. 숨은 변수 zz 가 있는 모델에서 데이터의 로그가능도는 적분입니다.

log⁡p(x)=log⁡∫p(x,z) dz\log p(x) = \log \int p(x, z)\, dz

로그가 적분 바깥에 있는 것이 문제입니다. 안쪽 적분을 닫힌 꼴로 못 구하고, 표본으로 어림하자니 로그 안쪽이라 평균으로 대신할 수도 없습니다. 여기서 우리가 고른 아무 분포 q(z)q(z) 를 곱하고 나누는 수를 씁니다.

log⁡p(x)=log⁡∫q(z)p(x,z)q(z) dz=log⁡ Ez∼q ⁣[p(x,z)q(z)]\log p(x) = \log \int q(z) \frac{p(x,z)}{q(z)}\, dz = \log\, \mathbb{E}_{z \sim q}\!\left[\frac{p(x,z)}{q(z)}\right]

이제 «로그가 씌워진 기댓값»이고, 젠센이 이것을 뒤집어 줍니다.

log⁡p(x)  ≥  Ez∼q ⁣[log⁡p(x,z)q(z)]  =  ELBO\log p(x) \;\ge\; \mathbb{E}_{z \sim q}\!\left[\log \frac{p(x,z)}{q(z)}\right] \;=\; \text{ELBO}

로그가 기댓값 안으로 들어갔습니다. 안쪽은 이제 표본을 뽑아 평균 내면 되는 값이라 계산이 됩니다.

하한을 밀어 올리면 그 위의 참값도 따라 올라간다

바꿔치기가 정당한 이유가 이 부등호입니다. ELBO는 언제나 log⁡p(x)\log p(x) 아래에 있으므로, ELBO를 올리면 그 위에 있는 참값을 아래에서 떠받쳐 올립니다.

갭이 곧 젠센 갭

다만 「올린 만큼 참값이 올라간다」는 보장은 없습니다. 둘 사이의 간격이 줄어드는 쪽으로 힘이 갈 수도 있기 때문입니다. 그 간격이 무엇인지는 이미 세워 둔 말로 부를 수 있습니다 — 하한이 얼마나 팽팽한지가 곧 이 부등식의 젠센 갭입니다.

그리고 젠센 갭은 확률변수가 얼마나 퍼져 있는지가 정한다고 했습니다. 여기서 그 확률변수는 p(x,z)/q(z)p(x,z)/q(z) 이고, 퍼진 정도를 정하는 것은 qq 입니다. 즉 qq 를 고르는 자유가 갭을 줄이는 손잡이입니다. qq 를 잘 골라 비가 어느 zz 에서나 비슷한 값이 되게 만들면 갭이 작아지고, qq 가 엉뚱한 자리에 질량을 두면 비가 크게 흔들려 갭이 벌어집니다. 그 간격이 정확히 어떤 양이고 언제 0이 되는지는 중급 60번 · 변분 하한과 ELBO가 다룹니다.

표본을 여러 개 뽑으면

갭을 줄이는 다른 손잡이가 하나 더 있습니다. 표본을 KK 개 뽑아 안쪽에서 먼저 평균 낸 뒤 로그를 씌우는 것입니다.

LK=E ⁣[log⁡1K∑k=1Kp(x,zk)q(zk)]\mathcal{L}_K = \mathbb{E}\!\left[\log \frac{1}{K}\sum_{k=1}^{K} \frac{p(x, z_k)}{q(z_k)}\right]

K=1K = 1 이면 원래의 ELBO이고, KK 를 키우면 하한이 올라가면서도 여전히 log⁡p(x)\log p(x) 아래에 머뭅니다. 이유는 앞 소절 그대로입니다. 안쪽에서 KK 개를 먼저 평균 내면 로그가 보는 확률변수의 분산이 1/K1/K 배로 줄고, 젠센 갭은 그 분산을 따라 줄기 때문입니다. 갭을 줄이는 길이 둘이라는 뜻이기도 합니다 — qq 를 더 잘 고르거나, 같은 qq 로 표본을 더 뽑거나.

같은 논법이 서는 다른 자리

이 모양을 한 번 익혀 두면 여러 곳에서 다시 보입니다. 중요도 샘플링으로 어떤 기댓값의 로그를 어림할 때, 표본 평균을 먼저 내고 로그를 씌운 값은 참값보다 작게 나오는 쪽으로 치우칩니다. 로그가 오목하기 때문이고, 치우친 크기가 곧 젠센 갭입니다. 대조 손실에서 분모의 합을 배치 안의 몇 개로만 어림하는 것도 같은 자리입니다 — 로그 안쪽의 합을 표본으로 바꾸는 순간 부등호가 하나 생기고, 배치를 키우는 것이 그 갭을 줄이는 손잡이가 됩니다.

지금 필요한 것은 발상의 뿌리 하나입니다. 구할 수 없는 값은 볼록성으로 밀어 내려 하한을 만들고, 그 하한을 대신 최대화한다.

코드로 확인하기

정의와 등호 조건

import numpy as np

# ① 정의를 직접 확인 — f(x) = -log x 는 볼록
f = lambda t: -np.log(t)
a, b, lam = 1.0, 100.0, 0.5
print(round(f(lam * a + (1 - lam) * b), 4),      # -3.922  그래프
      round(lam * f(a) + (1 - lam) * f(b), 4))   # -2.3026 현 — 현이 위에 있다

# ② log 는 오목 — 부등호가 뒤집힌다
x = np.array([1.0, 100.0]); w = np.array([0.5, 0.5])
print(round(np.log(x @ w), 4), round(w @ np.log(x), 4))    # 3.922 2.3026
print(round(np.exp(w @ np.log(x)), 4))                     # 10.0  기하평균

# ③ 등호는 X가 상수일 때만
x2 = np.array([50.5, 50.5])
print(round(np.log(x2 @ w), 4), round(w @ np.log(x2), 4))  # 3.922 3.922

# ④ 젠센이 보장한 H(p,q) - H(p) >= 0
p = np.array([0.5, 0.25, 0.125, 0.125])
for q in [p, np.full(4, 0.25), np.array([0.125, 0.125, 0.25, 0.5])]:
    print(round(float((p * np.log2(p / q)).sum()), 4))     # 0.0  0.25  0.875

④가 지난 글의 세 예와 같은 수입니다. 교차엔트로피 1.75·2.0·2.625에서 엔트로피 1.75를 뺀 값이고, 하나도 음수가 아닙니다.

역수의 평균과 평균의 역수

부등식이 실무의 수를 어떻게 바꾸는지 보이는 자리가 있습니다. 응답 지연을 다섯 번 재고, 그것을 초당 처리량으로 바꿔 평균 내는 경우입니다.

lat = np.array([20., 25., 30., 35., 200.])   # 밀리초
print(round(1000 / lat.mean(), 2))           # 16.13  전체 시간으로 계산한 처리량
print(round((1000 / lat).mean(), 2))         # 31.38  요청별 처리량의 평균

같은 다섯 번의 측정인데 두 배 가까이 갈립니다. 1/x1/x 가 x>0x > 0 에서 볼록하므로 젠센이 E[1/X]≥1/E[X]\mathbb{E}[1/X] \ge 1/\mathbb{E}[X] 를 보장하고, 두 번째 줄이 언제나 큰 쪽입니다. 200밀리초짜리 한 번이 첫 줄에서는 평균을 끌어올려 처리량을 끌어내리는데, 둘째 줄에서는 5라는 작은 값 하나로만 들어가 네 번의 큰 값에 묻힙니다. 요청별 처리량을 평균 낸 값을 시스템의 처리량이라고 부르면 안 되는 이유가 이 부등호입니다.

배치별 perplexity

언어 모델에서도 같은 일이 일어납니다. 손실의 지수를 취한 값을 perplexity라고 하는데, 배치마다 계산해 평균 내면 전체로 한 번에 계산한 값보다 큽니다.

nll = np.array([2.0, 2.5, 4.0])              # 배치 셋의 평균 손실
print(np.round(np.exp(nll), 2))              # [ 7.39 12.18 54.6 ]
print(round(float(np.exp(nll).mean()), 2))   # 24.72  배치별 ppl 의 평균
print(round(float(np.exp(nll.mean())), 2))   # 17.00  전체 손실에서 한 번에

exe^x 가 볼록하니 E[eX]≥eE[X]\mathbb{E}[e^X] \ge e^{\mathbb{E}[X]} 이고, 24.72가 17.00보다 큰 것이 그 부등호입니다. 손실이 튄 배치 하나가 지수를 지나며 크게 부풀어 평균을 끌어올립니다. 보고할 값은 언제나 전체 손실에서 지수를 한 번 취한 쪽입니다.

치우치는 방향은 볼록·오목이 정한다

세 예의 방향을 나란히 놓으면 규칙 하나로 정리됩니다.

ff 갈래 부등호
log⁡\log 오목 E[log⁡X]≤log⁡E[X]\mathbb{E}[\log X] \le \log \mathbb{E}[X]
1/x1/x 볼록 E[1/X]≥1/E[X]\mathbb{E}[1/X] \ge 1/\mathbb{E}[X]
exe^x 볼록 E[eX]≥eE[X]\mathbb{E}[e^X] \ge e^{\mathbb{E}[X]}

「평균을 먼저 내고 함수를 씌울 것인가, 함수를 먼저 씌우고 평균을 낼 것인가」에서 어느 쪽이 큰지는 ff 의 볼록·오목 하나만 보면 정해집니다. 얼마나 갈리는지는 분산이 정하고, 그 크기를 어림하는 식은 앞에서 세운 12f′′(μ)Var(X)\tfrac{1}{2} f''(\mu)\mathrm{Var}(X) 입니다.

정리

  • 볼록함수는 현이 그래프 위에 있는 함수다. f(λa+(1−λ)b)≤λf(a)+(1−λ)f(b)f(\lambda a + (1-\lambda)b) \le \lambda f(a) + (1-\lambda) f(b) 가 정의이고, 그래프 위쪽 영역이 볼록집합이라는 말과 같다.
  • 미분되는 함수라면 f′′≥0f'' \ge 0 으로, 여러 변수 함수라면 헤세 행렬이 반양정치인지로 판정한다. 다만 볼록성은 미분 가능성을 요구하지 않는다 — ∣x∣|x| 와 ReLU가 그 예다.
  • 판정보다 자주 쓰는 것은 조립이다. 합·양수배·최댓값·아핀 합성이 볼록을 보존하고, 그래서 logsumexp에서 softmax 교차엔트로피가 로짓에 대해 볼록하다는 것이 따라 나온다. 그 앞에 신경망을 태우면 그 성질이 사라진다.
  • 젠센 부등식은 f(E[X])≤E[f(X)]f(\mathbb{E}[X]) \le \mathbb{E}[f(X)] 이고, 오목함수에서는 뒤집힌다. 증명은 μ=E[X]\mu = \mathbb{E}[X] 에서 그은 접선 하나로 끝나고, 값이 몇 개인지·연속인지는 증명에 들어가지 않는다 — 가중치가 확률이기만 하면 된다.
  • 등호는 ff 가 그 구간에서 직선이거나 XX 가 상수일 때만 선다. ff 가 엄격 볼록이면 남는 것은 뒤쪽뿐이다.
  • 젠센 갭은 12f′′(μ)Var(X)\tfrac{1}{2} f''(\mu)\mathrm{Var}(X) 로 어림된다. 굽은 정도와 퍼진 정도의 곱이며, 분산이 평균의 제곱에 비해 커지면 이 어림이 실제 갭을 크게 밑돈다.
  • H(p,q)≥H(p)H(p,q) \ge H(p) 는 log⁡\log 의 오목성 + 젠센으로 얻고, 접선 부등식 log⁡t≤t−1\log t \le t-1 로도 같은 결론을 얻는다. 두 길이 쓰는 접선이 같은 직선이다.
  • 구할 수 없는 값은 하한으로 바꿔 최대화한다. 하한이 팽팽한 정도가 곧 젠센 갭이라, qq 를 잘 고르거나 표본을 KK 개 뽑아 안쪽에서 먼저 평균 내면 갭이 줄어든다.

loss = -(recon_term - kl_term)으로 돌아가 봅시다. 저 한 줄이 원래 목표가 아니라 그 하한이라는 사실은 이제 결함이 아니라 설계로 읽힙니다. 계산할 수 없는 로그를 기댓값 안으로 밀어 넣은 대가로 부등호가 하나 생겼고, 볼록성이 그 부등호의 방향을 붙들어 주고 있습니다.

그리고 이 글에서 두 번 나온 양이 하나 있습니다. H(p,q)−H(p)H(p,q) - H(p) 라는 초과분과 ELBO가 참값에 못 미치는 간격 — 둘은 같은 모양의 식입니다. 다음 글이 그 양에 이름을 붙이고, 왜 그것을 «거리»라고 부르면 안 되는지를 봅니다.


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

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