수학

MATH / 중급 53번

로그 미분 트릭: 샘플만으로 그래디언트를 얻기

RLHF 학습 루프의 손실은 로그 확률에 점수를 곱한 한 줄입니다. 그 한 줄이 어디서 나왔는지를 ∇E[f] = E[f·∇log p]라는 항등식의 유도로 따라가고, 이것이 언제 깨지는지, 왜 불편 추정량인데도 분산이 큰지, 이산 토큰에서는 왜 이 길밖에 없는지를 계산으로 확인합니다.

PALDYN Team21 MIN READ

RLHF 학습 코드를 열어 보면 손실이 이렇게 생겼습니다.

logp = torch.log_softmax(logits, -1).gather(-1, actions[..., None]).squeeze(-1)
loss = -(logp * advantage.detach()).mean()

모델이 실제로 뽑은 토큰의 로그 확률을 꺼내서, 그 답이 받은 점수를 곱하고, 부호를 뒤집어 평균 냅니다. 처음 보면 이상합니다. 점수를 손실 안에 상수처럼 곱해 놓았을 뿐인데 어떻게 이것이 「기대 보상을 올리는 방향」이 되는 걸까요. 게다가 advantage에는 detach()가 붙어 있어 역전파가 그쪽으로는 흐르지도 않습니다.

지난 글에서 우리는 J(θ)=Ey∼πθ[r(y)]J(\theta) = \mathbb{E}_{y\sim\pi_\theta}[r(y)] 를 세워 놓고 정확히 한 자리에서 막혔습니다. 미분을 합 안으로 넣으면 ∑yr(y)∇θπθ(y)\sum_y r(y)\nabla_\theta\pi_\theta(y) 가 남는데, 가중치 ∇θπθ\nabla_\theta\pi_\theta 가 확률이 아니라 표본평균으로 추정할 수 없었습니다. 필요한 것은 저 합을 다시 「πθ\pi_\theta 로 가중된 평균」 꼴로 되돌려 주는 항등식 하나였습니다.

그 항등식이 위 코드 한 줄의 정체입니다. 이 글은 그 항등식만 다룹니다 — 유도, 성립 조건, 불편성, 그리고 그 대가인 분산입니다. REINFORCE라는 알고리즘으로 조립하는 일은 정책 그래디언트가 맡습니다.

한 줄짜리 항등식

시작은 초등적인 미분 공식 하나입니다. 로그의 미분은 원래 함수를 자기 자신으로 나눈 것입니다.

∇θlog⁡pθ(x)=∇θpθ(x)pθ(x)\nabla_\theta \log p_\theta(x) = \frac{\nabla_\theta p_\theta(x)}{p_\theta(x)}

양변에 pθ(x)p_\theta(x) 를 곱하면 이렇게 됩니다.

∇θpθ(x)=pθ(x) ∇θlog⁡pθ(x)\nabla_\theta p_\theta(x) = p_\theta(x)\,\nabla_\theta \log p_\theta(x)

이것을 로그 미분 트릭이라고 부릅니다 — 확률의 그래디언트를 「확률 × 로그 확률의 그래디언트」로 바꿔 쓰는 항등식입니다. 별것 아닌 변형처럼 보이지만, 오른쪽에 pθp_\theta 가 곱해진 인수로 다시 나타났다는 것이 전부입니다. 지난 글에서 필요하다고 적었던 「∇πθ\nabla\pi_\theta 안에서 πθ\pi_\theta 하나를 밖으로 끄집어내기」가 바로 이것입니다.

목적식에 넣어 보겠습니다.

∇θ Ex∼pθ[f(x)]=∇θ∫pθ(x) f(x) dx=∫∇θpθ(x)⋅f(x) dx=∫pθ(x) f(x) ∇θlog⁡pθ(x) dx=Ex∼pθ[ f(x) ∇θlog⁡pθ(x) ]\begin{aligned} \nabla_\theta\,\mathbb{E}_{x\sim p_\theta}[f(x)] &= \nabla_\theta \int p_\theta(x)\,f(x)\,dx \\ &= \int \nabla_\theta p_\theta(x)\cdot f(x)\,dx \\ &= \int p_\theta(x)\,f(x)\,\nabla_\theta \log p_\theta(x)\,dx \\ &= \mathbb{E}_{x\sim p_\theta}\big[\,f(x)\,\nabla_\theta \log p_\theta(x)\,\big] \end{aligned}

세 걸음뿐입니다. 첫 줄에서 둘째 줄로 갈 때 미분을 적분 안으로 넣었고(이 걸음에 조건이 붙습니다 — 아래에서 봅니다), 둘째에서 셋째로 갈 때 로그 미분 트릭을 썼고, 셋째 줄은 이미 「pθp_\theta 로 가중된 평균」 꼴이라 그대로 기댓값으로 읽었습니다.

로그 미분 트릭이 그래디언트를 기댓값으로 되돌리는 세 걸음

왼쪽 끝과 오른쪽 끝만 남기면 이 글의 제목이 되는 식입니다.

∇θ Ex∼pθ[f(x)]=Ex∼pθ[f(x) ∇θlog⁡pθ(x)]\nabla_\theta\,\mathbb{E}_{x\sim p_\theta}[f(x)] = \mathbb{E}_{x\sim p_\theta}\big[f(x)\,\nabla_\theta \log p_\theta(x)\big]

왼쪽은 계산할 수 없고 오른쪽은 계산할 수 있습니다. 오른쪽은 pθp_\theta 에서 표본을 뽑아 각 표본에서 f(x)∇θlog⁡pθ(x)f(x)\nabla_\theta\log p_\theta(x) 를 계산해 평균 내면 되기 때문입니다. 그리고 ∇θlog⁡pθ(x)\nabla_\theta \log p_\theta(x) 는 우리가 늘 계산하던 것입니다 — 지도학습에서 로그가능도를 미분하던 그 값입니다.

이 항등식을 쓴 추정량에는 이름이 있습니다. 점수 함수 추정량이라고 부르는데, 통계학에서 ∇θlog⁡pθ\nabla_\theta \log p_\theta 를 점수 함수라 불러 온 데서 온 이름입니다. 강화학습 쪽에서는 REINFORCE 추정량이라고도 합니다.

g^=1n∑i=1nf(x(i)) ∇θlog⁡pθ(x(i)),x(i)∼pθ\widehat{g} = \frac{1}{n}\sum_{i=1}^{n} f(x^{(i)})\,\nabla_\theta \log p_\theta(x^{(i)}),\qquad x^{(i)}\sim p_\theta

맨 앞 코드로 돌아가면 이제 읽힙니다. logp가 log⁡πθ(y)\log \pi_\theta(y) 이고, advantage가 ff 자리의 점수이며, 곱해서 평균 내고 backward()를 부르면 자동미분이 ∇θlog⁡πθ\nabla_\theta\log\pi_\theta 를 만들어 냅니다. 부호가 뒤집힌 것은 JJ 를 최대화해야 하는데 옵티마이저는 최소화하기 때문이고, detach()가 붙은 것은 점수가 ff 자리의 상수여야 하기 때문입니다 — 위 유도에서 ff 는 θ\theta 와 무관한 함수였습니다.

언제 깨지는가

유도의 둘째 걸음, 그러니까 미분을 적분 안으로 옮긴 자리가 공짜가 아닙니다. 조건은 둘입니다.

첫째, 지지집합이 θ\theta 에 의존하지 않아야 합니다. 지지집합이란 확률이 0이 아닌 값들의 모임입니다. 적분 구간의 끝이 θ\theta 에 따라 움직이면 θ\theta 를 흔들 때 구간 자체가 변하므로, 미분이 피적분함수에만 작용한다고 볼 수 없습니다.

깨지는 예를 하나 보면 확실합니다. pθp_\theta 를 구간 [0,θ][0,\theta] 위의 균등분포로 두고 f(x)=xf(x) = x 를 재 봅니다. 왼쪽은 손으로 바로 나옵니다.

E[x]=θ2⟹∇θE[x]=12\mathbb{E}[x] = \frac{\theta}{2} \quad\Longrightarrow\quad \nabla_\theta \mathbb{E}[x] = \frac{1}{2}

오른쪽은 log⁡pθ(x)=−log⁡θ\log p_\theta(x) = -\log\theta 이므로 ∇θlog⁡pθ(x)=−1/θ\nabla_\theta \log p_\theta(x) = -1/\theta 이고,

E[x⋅(−1/θ)]=−1θ⋅θ2=−12\mathbb{E}\big[x\cdot(-1/\theta)\big] = -\frac{1}{\theta}\cdot\frac{\theta}{2} = -\frac{1}{2}

부호까지 반대입니다. 항등식이 성립하지 않는 것이고, 원인은 확률밀도가 x=θx=\theta 에서 뚝 끊기며 그 끊기는 자리가 θ\theta 와 함께 움직인다는 데 있습니다.

지지집합이 θ와 함께 움직이면 항등식이 깨진다

둘째, 미분과 적분을 바꿔도 되는 정칙성 조건이 필요합니다. 엄밀하게는 ∥∇θpθ(x)f(x)∥\|\nabla_\theta p_\theta(x) f(x)\| 를 θ\theta 와 무관하게 위에서 눌러 주는 적분 가능한 함수가 있어야 합니다. 다행히 우리가 다루는 자리에서는 둘 다 자동으로 만족됩니다. 언어모델의 정책은 소프트맥스에서 나오므로 모든 토큰에 0보다 큰 확률을 주고, 그 지지집합은 어휘 전체라서 θ\theta 가 아무리 변해도 그대로입니다. 유한한 어휘 위의 합은 적분 교환 문제도 일으키지 않습니다.

하나 더 짚어 둘 것이 있습니다. pθ(x)p_\theta(x) 가 0이면 로그 미분 트릭의 분모가 0이 되는데, 그런 xx 는 애초에 표본으로 뽑히지 않으므로 기댓값에 기여하지 않습니다. 다만 확률이 0에 가까운 표본이 뽑히면 ∇log⁡pθ\nabla\log p_\theta 가 매우 커집니다. 이것이 아래에서 볼 분산 문제의 씨앗입니다.

정말 불편 추정량인가

불편 추정량은 추정값의 기댓값이 참값과 같은 추정량입니다. 위 항등식이 성립하면 g^\widehat g 는 정의상 불편입니다 — 각 항의 기댓값이 ∇θJ\nabla_\theta J 이고 평균의 기댓값도 그대로이니까요. 그래도 숫자로 확인하는 편이 확실합니다. 지난 글에서 유한차분이 무너졌던 그 장난감 문제를 그대로 씁니다.

import numpy as np

R = np.array([1.0, 0.0, -1.0])          # 세 가지 답의 점수
def pi(th):
    e = np.exp(th - th.max()); return e / e.sum()

th = np.array([0.3, 0.1, -0.2])
p = pi(th)
exact = p * (R - p @ R)                 # 답이 셋뿐이라 손으로 구할 수 있다
print("정확한 ∇J =", np.round(exact, 6))

rng = np.random.default_rng(0)
for n in (100, 10_000, 1_000_000):
    y = rng.choice(3, size=n, p=p)      # 정책에서 답을 뽑고
    G = R[y][:, None] * (np.eye(3)[y] - p)   # r(y) · ∇log π(y) 를 계산해
    g = G.mean(0)                            # 평균 낸다
    err = np.linalg.norm(g - exact) / np.linalg.norm(exact)
    print(f"  샘플 {n:>9,}개  추정 ∇J = {np.round(g, 6)}  상대오차 {err:.4f}")

# 정확한 ∇J = [ 0.345432 -0.054769 -0.290663]
#   샘플       100개  추정 ∇J = [ 0.333507 -0.013503 -0.320004]  상대오차 0.1144
#   샘플    10,000개  추정 ∇J = [ 0.34625  -0.05506  -0.291189]  상대오차 0.0022
#   샘플 1,000,000개  추정 ∇J = [ 0.345403 -0.054558 -0.290845]  상대오차 0.0006

소프트맥스의 로그 확률을 미분하면 ∇θklog⁡π(y)=1[y=k]−πk\nabla_{\theta_k}\log\pi(y) = \mathbb{1}[y=k] - \pi_k 라서 코드의 np.eye(3)[y] - p가 그 값입니다.

숫자를 지난 글과 나란히 놓으면 격차가 보입니다. 유한차분은 표본 100만 개를 쓰고도 상대오차 219%였습니다. 점수 함수 추정량은 100개로 11%, 만 개로 0.2%입니다. 그리고 유한차분은 파라미터마다 따로 흔들어야 했지만 여기서는 표본 한 묶음으로 모든 파라미터의 편미분을 한꺼번에 얻습니다 — 역전파가 하는 일이 정확히 그것이기 때문입니다.

오차가 1/n1/\sqrt n 으로 줄어드는 것도 표에서 보입니다. 표본을 100배 늘리면 오차가 10분의 1이 됩니다(11.4% → 0.22%가 대략 그렇습니다). 표본오차에서 본 그 법칙 그대로입니다.

대가: 분산이 크다

불편이라는 것은 「평균적으로 맞다」는 말이지 「한 번 재면 맞다」는 말이 아닙니다. 점수 함수 추정량은 불편이지만 분산이 큽니다. 그리고 그 분산이 어디서 오는지 보면, 고칠 자리도 함께 보입니다.

추정량의 한 표본짜리 값은 f(x)∇θlog⁡pθ(x)f(x)\nabla_\theta\log p_\theta(x) 라는 곱입니다. 두 인수 모두 클 수 있습니다.

  • ff 쪽 — 점수의 절댓값이 크면 곱 전체가 그만큼 커집니다. 그런데 아래에서 보듯 점수에 상수를 더해도 참 그래디언트는 바뀌지 않습니다. 즉 분산은 커지는데 신호는 그대로입니다.
  • ∇log⁡pθ\nabla\log p_\theta 쪽 — 확률이 작은 표본이 뽑히면 이 값이 큽니다. 그런 표본은 드물게 뽑히지만 뽑히면 크게 튑니다.
  • 길이 쪽 — 언어모델에서 log⁡πθ(y)=∑tlog⁡πθ(yt∣⋅)\log\pi_\theta(y) = \sum_t \log\pi_\theta(y_t\mid\cdot) 이므로 답이 200토큰이면 200개 항의 합입니다. 항이 늘수록 그 합의 흔들림도 커집니다.

첫 번째가 가장 손대기 쉬운 자리라 확인해 보겠습니다. 세 점수 (1,0,−1)(1, 0, -1) 에 상수 cc 를 똑같이 더하면 어떻게 될까요. JJ 는 cc 만큼 올라가지만 ∇θJ\nabla_\theta J 는 그대로입니다 — 모든 답을 똑같이 올려 주는 것은 답들 사이의 우열을 바꾸지 않으니까요.

for c in (0.0, 5.0, 50.0):
    G = (R + c)[:, None] * (np.eye(3) - p)   # 각 답에서의 추정량 값
    mean = p @ G
    var = (p @ (G ** 2) - mean ** 2).sum()
    need = var / (0.1 * np.linalg.norm(mean)) ** 2
    print(f"c={c:>5}  ∇J={np.round(mean, 6)}  총분산={var:9.4f}  10%오차에 필요한 표본≈{need:,.0f}")

# c=  0.0  ∇J=[ 0.345432 -0.054769 -0.290663]  총분산=   0.2200  10%오차에 필요한 표본≈106
# c=  5.0  ∇J=[ 0.345432 -0.054769 -0.290663]  총분산=  16.5922  10%오차에 필요한 표본≈8,023
# c= 50.0  ∇J=[ 0.345432 -0.054769 -0.290663]  총분산=1634.2694  10%오차에 필요한 표본≈790,237

∇J 열은 세 줄이 완전히 같습니다. 그런데 필요한 표본 수는 106개에서 79만 개로 7,400배가 되었습니다. 점수를 어디에 맞춰 재느냐는 정보를 하나도 바꾸지 않는데 학습 비용은 네 자릿수로 달라집니다.

점수에 상수를 더하면 그래디언트는 그대로인데 분산만 커진다

이 관찰이 다음 글의 출발점입니다. 참 그래디언트를 건드리지 않으면서 분산만 줄일 수 있는 자유도가 여기 있다는 뜻이니까요. 실제 RLHF 코드에서 reward가 아니라 advantage라는 이름의 값이 곱해져 있던 것도 이 때문입니다.

다른 길은 왜 막혀 있는가

기댓값의 그래디언트를 얻는 방법이 이것만 있는 것은 아닙니다. 무작위성을 파라미터 밖으로 빼내는 길도 있습니다. 정규분포에서 뽑는 경우라면

x=μθ+σθ ϵ,ϵ∼N(0,1)x = \mu_\theta + \sigma_\theta\,\epsilon,\qquad \epsilon\sim\mathcal{N}(0,1)

로 적을 수 있습니다. 이제 무작위성은 ϵ\epsilon 에만 있고 θ\theta 와 무관하므로, 기댓값을 취하는 분포가 θ\theta 에 의존하지 않습니다. 지난 글에서 지도학습이 쉬웠던 그 상황으로 돌아온 것이고, 미분이 기댓값 안으로 그냥 들어갑니다.

∇θ Eϵ[f(μθ+σθϵ)]=Eϵ[f′(x) ∇θ(μθ+σθϵ)]\nabla_\theta\,\mathbb{E}_{\epsilon}\big[f(\mu_\theta + \sigma_\theta\epsilon)\big] = \mathbb{E}_{\epsilon}\big[f'(x)\,\nabla_\theta(\mu_\theta + \sigma_\theta\epsilon)\big]

이 방식을 재매개변수화 트릭이라 하고, 여기서는 ff 를 실제로 미분해 그 기울기를 씁니다. 그래서 점수 함수 추정량보다 분산이 훨씬 작은 것이 보통입니다 — 점수 하나만 보는 대신 ff 가 어느 쪽으로 기울어 있는지를 직접 읽기 때문입니다.

점수 함수 추정량과 재매개변수화 트릭이 갈리는 자리

점수 함수 추정량 재매개변수화 트릭
쓰는 항등식 ∇p=p ∇log⁡p\nabla p = p\,\nabla\log p x=gθ(ϵ)x = g_\theta(\epsilon)
ff 에 요구하는 것 값만 주면 된다 ff 가 미분 가능해야 한다
xx 의 종류 이산·연속 모두 연속만
분산 크다 대체로 작다
대표적 쓰임 RLHF, 정책 그래디언트 VAE, 확산모델의 순방향 과정

언어모델이 왼쪽 칸에 갇혀 있는 이유는 표의 셋째 줄입니다. 토큰은 이산이라 θ\theta 에 대해 미분 가능한 함수로 적을 수 없습니다 — 「토큰 번호 4128」을 연속적으로 조금 움직인다는 말이 성립하지 않습니다. 그리고 넷째 줄도 있습니다. 보상 모델은 미분 가능하게 만들 수 있지만 단위 테스트 통과 여부나 사람의 판정은 그렇지 않습니다. ff 의 값만 요구하는 쪽이 훨씬 넓은 채점자를 받아들입니다.

재매개변수화 트릭 자체와 확산모델의 순방향 과정에서 그것이 어떻게 쓰이는지는 중급 59번 · 재매개변수화 트릭과 폐형 forward 과정이 맡습니다.

다시 그 한 줄로

맨 앞의 두 줄짜리 코드는 이제 유도의 마지막 줄을 그대로 옮겨 적은 것으로 읽힙니다.

∇θJ(θ)=Ey∼πθ[ r(y) ∇θlog⁡πθ(y∣x) ]\nabla_\theta J(\theta) = \mathbb{E}_{y\sim\pi_\theta}\big[\,r(y)\,\nabla_\theta \log \pi_\theta(y\mid x)\,\big]

logp * advantage를 평균 내고 부호를 뒤집어 backward()를 부르면, 자동미분이 ∇θlog⁡πθ\nabla_\theta\log\pi_\theta 를 만들고 점수가 그 앞에 가중치로 붙습니다. 뽑힌 답의 점수가 양수면 그 답의 로그 확률을 올리는 방향으로, 음수면 내리는 방향으로 한 걸음을 냅니다. 「좋았던 것을 더 자주 하게 만든다」는 문장이 식으로는 이렇게 생겼습니다.

그리고 마지막 절에서 본 것 때문에 그 자리에 reward가 아니라 advantage가 들어가 있습니다. 무엇을 빼야 편향 없이 분산만 줄어드는지가 다음 글의 주제입니다.

정리

  • 로그 미분 트릭은 ∇θpθ(x)=pθ(x)∇θlog⁡pθ(x)\nabla_\theta p_\theta(x) = p_\theta(x)\nabla_\theta\log p_\theta(x) 라는 항등식이다. 로그의 미분 공식에 pθp_\theta 를 곱한 것뿐이지만, 확률이 곱해진 인수로 되살아난다는 것이 핵심이다.
  • 이것을 목적식에 넣으면 ∇θEpθ[f]=Epθ[f ∇θlog⁡pθ]\nabla_\theta\mathbb{E}_{p_\theta}[f] = \mathbb{E}_{p_\theta}[f\,\nabla_\theta\log p_\theta] 가 되고, 왼쪽과 달리 오른쪽은 표본평균으로 추정할 수 있다. 이 추정량을 점수 함수 추정량(REINFORCE 추정량)이라 부른다.
  • 성립 조건은 둘이다 — 지지집합이 θ\theta 에 의존하지 않을 것, 미분과 적분을 바꿔도 될 것. [0,θ][0,\theta] 균등분포에서는 첫 조건이 깨져 부호까지 반대인 값이 나온다. 소프트맥스 정책은 두 조건을 모두 만족한다.
  • 항등식이 성립하므로 추정량은 불편이다. 답이 셋뿐인 문제에서 표본 100개로 상대오차 11%, 만 개로 0.2%였다 — 같은 문제에서 유한차분이 100만 개로도 219%였던 것과 대비된다. 오차는 1/n1/\sqrt n 으로 준다.
  • 대신 분산이 크다. 점수의 절댓값, 확률이 작은 표본에서 폭발하는 ∇log⁡p\nabla\log p, 그리고 토큰 수만큼 쌓이는 합이 원인이다. 점수에 상수 50을 더하면 참 그래디언트는 그대로인데 같은 정확도에 필요한 표본이 7,400배가 된다.
  • 재매개변수화 트릭은 무작위성을 θ\theta 밖으로 빼내 ff 를 직접 미분하는 다른 길이고 분산이 훨씬 작다. 다만 xx 가 연속이고 ff 가 미분 가능해야 한다. 토큰은 이산이므로 언어모델에는 점수 함수 추정량밖에 없다.

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

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