RLHF 학습 코드를 열어 보면 손실이 이렇게 생겼습니다.
logp = torch.log_softmax(logits, -1).gather(-1, actions[..., None]).squeeze(-1)
loss = -(logp * advantage.detach()).mean()
모델이 실제로 뽑은 토큰의 로그 확률을 꺼내서, 그 답이 받은 점수를 곱하고, 부호를 뒤집어 평균 냅니다. 처음 보면 이상합니다. 점수를 손실 안에 상수처럼 곱해 놓았을 뿐인데 어떻게 이것이 「기대 보상을 올리는 방향」이 되는 걸까요. 게다가 advantage에는 detach()가 붙어 있어 역전파가 그쪽으로는 흐르지도 않습니다.
지난 글에서 우리는 를 세워 놓고 정확히 한 자리에서 막혔습니다. 미분을 합 안으로 넣으면 가 남는데, 가중치 가 확률이 아니라 표본평균으로 추정할 수 없었습니다. 필요한 것은 저 합을 다시 「 로 가중된 평균」 꼴로 되돌려 주는 항등식 하나였습니다.
그 항등식이 위 코드 한 줄의 정체입니다. 이 글은 그 항등식만 다룹니다 — 유도, 성립 조건, 불편성, 그리고 그 대가인 분산입니다. REINFORCE라는 알고리즘으로 조립하는 일은 정책 그래디언트가 맡습니다.
한 줄짜리 항등식
시작은 초등적인 미분 공식 하나입니다. 로그의 미분은 원래 함수를 자기 자신으로 나눈 것입니다.
양변에 를 곱하면 이렇게 됩니다.
이것을 로그 미분 트릭이라고 부릅니다 — 확률의 그래디언트를 「확률 × 로그 확률의 그래디언트」로 바꿔 쓰는 항등식입니다. 별것 아닌 변형처럼 보이지만, 오른쪽에 가 곱해진 인수로 다시 나타났다는 것이 전부입니다. 지난 글에서 필요하다고 적었던 「 안에서 하나를 밖으로 끄집어내기」가 바로 이것입니다.
목적식에 넣어 보겠습니다.
세 걸음뿐입니다. 첫 줄에서 둘째 줄로 갈 때 미분을 적분 안으로 넣었고(이 걸음에 조건이 붙습니다 — 아래에서 봅니다), 둘째에서 셋째로 갈 때 로그 미분 트릭을 썼고, 셋째 줄은 이미 「 로 가중된 평균」 꼴이라 그대로 기댓값으로 읽었습니다.
왼쪽 끝과 오른쪽 끝만 남기면 이 글의 제목이 되는 식입니다.
왼쪽은 계산할 수 없고 오른쪽은 계산할 수 있습니다. 오른쪽은 에서 표본을 뽑아 각 표본에서 를 계산해 평균 내면 되기 때문입니다. 그리고 는 우리가 늘 계산하던 것입니다 — 지도학습에서 로그가능도를 미분하던 그 값입니다.
이 항등식을 쓴 추정량에는 이름이 있습니다. 점수 함수 추정량이라고 부르는데, 통계학에서 를 점수 함수라 불러 온 데서 온 이름입니다. 강화학습 쪽에서는 REINFORCE 추정량이라고도 합니다.
맨 앞 코드로 돌아가면 이제 읽힙니다. logp가 이고, advantage가 자리의 점수이며, 곱해서 평균 내고 backward()를 부르면 자동미분이 를 만들어 냅니다. 부호가 뒤집힌 것은 를 최대화해야 하는데 옵티마이저는 최소화하기 때문이고, detach()가 붙은 것은 점수가 자리의 상수여야 하기 때문입니다 — 위 유도에서 는 와 무관한 함수였습니다.
언제 깨지는가
유도의 둘째 걸음, 그러니까 미분을 적분 안으로 옮긴 자리가 공짜가 아닙니다. 조건은 둘입니다.
첫째, 지지집합이 에 의존하지 않아야 합니다. 지지집합이란 확률이 0이 아닌 값들의 모임입니다. 적분 구간의 끝이 에 따라 움직이면 를 흔들 때 구간 자체가 변하므로, 미분이 피적분함수에만 작용한다고 볼 수 없습니다.
깨지는 예를 하나 보면 확실합니다. 를 구간 위의 균등분포로 두고 를 재 봅니다. 왼쪽은 손으로 바로 나옵니다.
오른쪽은 이므로 이고,
부호까지 반대입니다. 항등식이 성립하지 않는 것이고, 원인은 확률밀도가 에서 뚝 끊기며 그 끊기는 자리가 와 함께 움직인다는 데 있습니다.
둘째, 미분과 적분을 바꿔도 되는 정칙성 조건이 필요합니다. 엄밀하게는 를 와 무관하게 위에서 눌러 주는 적분 가능한 함수가 있어야 합니다. 다행히 우리가 다루는 자리에서는 둘 다 자동으로 만족됩니다. 언어모델의 정책은 소프트맥스에서 나오므로 모든 토큰에 0보다 큰 확률을 주고, 그 지지집합은 어휘 전체라서 가 아무리 변해도 그대로입니다. 유한한 어휘 위의 합은 적분 교환 문제도 일으키지 않습니다.
하나 더 짚어 둘 것이 있습니다. 가 0이면 로그 미분 트릭의 분모가 0이 되는데, 그런 는 애초에 표본으로 뽑히지 않으므로 기댓값에 기여하지 않습니다. 다만 확률이 0에 가까운 표본이 뽑히면 가 매우 커집니다. 이것이 아래에서 볼 분산 문제의 씨앗입니다.
정말 불편 추정량인가
불편 추정량은 추정값의 기댓값이 참값과 같은 추정량입니다. 위 항등식이 성립하면 는 정의상 불편입니다 — 각 항의 기댓값이 이고 평균의 기댓값도 그대로이니까요. 그래도 숫자로 확인하는 편이 확실합니다. 지난 글에서 유한차분이 무너졌던 그 장난감 문제를 그대로 씁니다.
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
소프트맥스의 로그 확률을 미분하면 라서 코드의 np.eye(3)[y] - p가 그 값입니다.
숫자를 지난 글과 나란히 놓으면 격차가 보입니다. 유한차분은 표본 100만 개를 쓰고도 상대오차 219%였습니다. 점수 함수 추정량은 100개로 11%, 만 개로 0.2%입니다. 그리고 유한차분은 파라미터마다 따로 흔들어야 했지만 여기서는 표본 한 묶음으로 모든 파라미터의 편미분을 한꺼번에 얻습니다 — 역전파가 하는 일이 정확히 그것이기 때문입니다.
오차가 으로 줄어드는 것도 표에서 보입니다. 표본을 100배 늘리면 오차가 10분의 1이 됩니다(11.4% → 0.22%가 대략 그렇습니다). 표본오차에서 본 그 법칙 그대로입니다.
대가: 분산이 크다
불편이라는 것은 「평균적으로 맞다」는 말이지 「한 번 재면 맞다」는 말이 아닙니다. 점수 함수 추정량은 불편이지만 분산이 큽니다. 그리고 그 분산이 어디서 오는지 보면, 고칠 자리도 함께 보입니다.
추정량의 한 표본짜리 값은 라는 곱입니다. 두 인수 모두 클 수 있습니다.
- 쪽 — 점수의 절댓값이 크면 곱 전체가 그만큼 커집니다. 그런데 아래에서 보듯 점수에 상수를 더해도 참 그래디언트는 바뀌지 않습니다. 즉 분산은 커지는데 신호는 그대로입니다.
- 쪽 — 확률이 작은 표본이 뽑히면 이 값이 큽니다. 그런 표본은 드물게 뽑히지만 뽑히면 크게 튑니다.
- 길이 쪽 — 언어모델에서 이므로 답이 200토큰이면 200개 항의 합입니다. 항이 늘수록 그 합의 흔들림도 커집니다.
첫 번째가 가장 손대기 쉬운 자리라 확인해 보겠습니다. 세 점수 에 상수 를 똑같이 더하면 어떻게 될까요. 는 만큼 올라가지만 는 그대로입니다 — 모든 답을 똑같이 올려 주는 것은 답들 사이의 우열을 바꾸지 않으니까요.
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라는 이름의 값이 곱해져 있던 것도 이 때문입니다.
다른 길은 왜 막혀 있는가
기댓값의 그래디언트를 얻는 방법이 이것만 있는 것은 아닙니다. 무작위성을 파라미터 밖으로 빼내는 길도 있습니다. 정규분포에서 뽑는 경우라면
로 적을 수 있습니다. 이제 무작위성은 에만 있고 와 무관하므로, 기댓값을 취하는 분포가 에 의존하지 않습니다. 지난 글에서 지도학습이 쉬웠던 그 상황으로 돌아온 것이고, 미분이 기댓값 안으로 그냥 들어갑니다.
이 방식을 재매개변수화 트릭이라 하고, 여기서는 를 실제로 미분해 그 기울기를 씁니다. 그래서 점수 함수 추정량보다 분산이 훨씬 작은 것이 보통입니다 — 점수 하나만 보는 대신 가 어느 쪽으로 기울어 있는지를 직접 읽기 때문입니다.
| 점수 함수 추정량 | 재매개변수화 트릭 | |
|---|---|---|
| 쓰는 항등식 | ||
| 에 요구하는 것 | 값만 주면 된다 | 가 미분 가능해야 한다 |
| 의 종류 | 이산·연속 모두 | 연속만 |
| 분산 | 크다 | 대체로 작다 |
| 대표적 쓰임 | RLHF, 정책 그래디언트 | VAE, 확산모델의 순방향 과정 |
언어모델이 왼쪽 칸에 갇혀 있는 이유는 표의 셋째 줄입니다. 토큰은 이산이라 에 대해 미분 가능한 함수로 적을 수 없습니다 — 「토큰 번호 4128」을 연속적으로 조금 움직인다는 말이 성립하지 않습니다. 그리고 넷째 줄도 있습니다. 보상 모델은 미분 가능하게 만들 수 있지만 단위 테스트 통과 여부나 사람의 판정은 그렇지 않습니다. 의 값만 요구하는 쪽이 훨씬 넓은 채점자를 받아들입니다.
재매개변수화 트릭 자체와 확산모델의 순방향 과정에서 그것이 어떻게 쓰이는지는 중급 59번 · 재매개변수화 트릭과 폐형 forward 과정이 맡습니다.
다시 그 한 줄로
맨 앞의 두 줄짜리 코드는 이제 유도의 마지막 줄을 그대로 옮겨 적은 것으로 읽힙니다.
logp * advantage를 평균 내고 부호를 뒤집어 backward()를 부르면, 자동미분이 를 만들고 점수가 그 앞에 가중치로 붙습니다. 뽑힌 답의 점수가 양수면 그 답의 로그 확률을 올리는 방향으로, 음수면 내리는 방향으로 한 걸음을 냅니다. 「좋았던 것을 더 자주 하게 만든다」는 문장이 식으로는 이렇게 생겼습니다.
그리고 마지막 절에서 본 것 때문에 그 자리에 reward가 아니라 advantage가 들어가 있습니다. 무엇을 빼야 편향 없이 분산만 줄어드는지가 다음 글의 주제입니다.
정리
- 로그 미분 트릭은 라는 항등식이다. 로그의 미분 공식에 를 곱한 것뿐이지만, 확률이 곱해진 인수로 되살아난다는 것이 핵심이다.
- 이것을 목적식에 넣으면 가 되고, 왼쪽과 달리 오른쪽은 표본평균으로 추정할 수 있다. 이 추정량을 점수 함수 추정량(REINFORCE 추정량)이라 부른다.
- 성립 조건은 둘이다 — 지지집합이 에 의존하지 않을 것, 미분과 적분을 바꿔도 될 것. 균등분포에서는 첫 조건이 깨져 부호까지 반대인 값이 나온다. 소프트맥스 정책은 두 조건을 모두 만족한다.
- 항등식이 성립하므로 추정량은 불편이다. 답이 셋뿐인 문제에서 표본 100개로 상대오차 11%, 만 개로 0.2%였다 — 같은 문제에서 유한차분이 100만 개로도 219%였던 것과 대비된다. 오차는 으로 준다.
- 대신 분산이 크다. 점수의 절댓값, 확률이 작은 표본에서 폭발하는 , 그리고 토큰 수만큼 쌓이는 합이 원인이다. 점수에 상수 50을 더하면 참 그래디언트는 그대로인데 같은 정확도에 필요한 표본이 7,400배가 된다.
- 재매개변수화 트릭은 무작위성을 밖으로 빼내 를 직접 미분하는 다른 길이고 분산이 훨씬 작다. 다만 가 연속이고 가 미분 가능해야 한다. 토큰은 이산이므로 언어모델에는 점수 함수 추정량밖에 없다.
읽어주셔서 감사합니다. 😊

