Chapter 3: REINFORCE — 샘플로 추정하고, 베이스라인을 빼다

의문

답이 셋뿐인 장난감 문제를 떠올려 보자. "대한민국의 수도는?"에 모델이 “서울입니다”(1점), “부산입니다”(0점), “서울이요”(0.6점) 가운데 하나로만 답한다. 평균 점수를 올리려고 샘플 하나씩 뽑아 걸으면 학습 곡선이 요동친다. 점수마다 똑같이 5를 더해 주면 요동은 더 심해진다. 모든 답에 똑같이 얹은 점수가 어느 답이 나은지를 바꿀 리 없고, 평균을 올리는 방향도 그대로인데 왜 그럴까? 그리고 그 요동은 어떻게 줄이는가?

평균 점수를 올리는 방향은 정책 그래디언트로 적힌다.

∇θJ=Ey∼πθ[ R(y) ∇θlog⁡πθ(y) ]\nabla_\theta \textcolor{#d62728}{J} = \mathbb{E}_{\textcolor{#1c9c60}{y} \sim \textcolor{#1565c0}{\pi_\theta}}\big[\, \textcolor{#d9670b}{R}(\textcolor{#1c9c60}{y})\, \nabla_\theta \log \textcolor{#1565c0}{\pi_\theta}(\textcolor{#1c9c60}{y}) \,\big]
J목적함수: 평균 점수R(y)응답 y의 점수πθ학습 중인 정책 (프롬프트를 받아 답을 고르는 언어모델)\small\begin{array}{ll} \textcolor{#d62728}{J} & \text{목적함수: 평균 점수} \\ \textcolor{#d9670b}{R}(\textcolor{#1c9c60}{y}) & \text{응답 }\textcolor{#1c9c60}{y}\text{의 점수} \\ \textcolor{#1565c0}{\pi_\theta} & \text{학습 중인 정책 (프롬프트를 받아 답을 고르는 언어모델)} \end{array}

정확한 식이지만 언어모델에서는 계산할 수 없다. 기대값을 구하려면 가능한 모든 응답을 더해야 하는데, 어휘가 1만 개인 모델이 20토큰짜리 답만 쓴다고 해도 가짓수가 1000020=108010000^{20} = 10^{80} 이다. 관측 가능한 우주의 원자 수로 흔히 꼽는 값(약 108010^{80})과 맞먹고, 실제 답은 이보다 훨씬 길다. 그러니 할 수 있는 일은 하나뿐이다: 몇 개를 뽑아서 평균을 낸다.

몬테카를로 추정: 기대값 대신 샘플 평균

뽑은 몇 개의 평균이 정말 참 그래디언트를 대신할 수 있을까? 대신할 수 있다면, 한 걸음은 실제로 어떻게 돌아가고, 무엇이 비싼가?

샘플 평균으로 걷기

현재 모델로 응답 N\textcolor{#5f970c}{N}개를 뽑고, 각자의 R⋅∇log⁡π\textcolor{#d9670b}{R} \cdot \nabla \log \textcolor{#1565c0}{\pi}를 평균낸다.

g^=1N∑i=1NR(yi) ∇θlog⁡πθ(yi),yi∼πθ\textcolor{#cc00ff}{\hat g} = \frac{1}{\textcolor{#5f970c}{N}} \sum_{i=1}^{\textcolor{#5f970c}{N}} \textcolor{#d9670b}{R}(\textcolor{#1c9c60}{y_i})\, \nabla_\theta \log \textcolor{#1565c0}{\pi_\theta}(\textcolor{#1c9c60}{y_i}), \qquad \textcolor{#1c9c60}{y_i} \sim \textcolor{#1565c0}{\pi_\theta}
g^그래디언트 ∇θJ의 추정치 (샘플 평균)N뽑은 응답 수 (배치 크기)R(yi)i번째 응답의 점수yi∼πθ지금의 정책이 뽑은 i번째 응답 \small\begin{array}{ll} \textcolor{#cc00ff}{\hat g} & \text{그래디언트 }\nabla_\theta \textcolor{#d62728}{J}\text{의 추정치 (샘플 평균)} \\ \textcolor{#5f970c}{N} & \text{뽑은 응답 수 (배치 크기)} \\ \textcolor{#d9670b}{R}(\textcolor{#1c9c60}{y_i}) & \text{}i\text{번째 응답의 점수} \\ \textcolor{#1c9c60}{y_i} \sim \textcolor{#1565c0}{\pi_\theta} & \text{지금의 정책이 뽑은 }i\text{번째 응답} \end{array}

각 항의 기대값이 ∇J\nabla \textcolor{#d62728}{J}이므로 평균의 기대값도 ∇J\nabla \textcolor{#d62728}{J}다. 이런 추정을 불편(unbiased, 치우침 없음)하다고 한다. 여러 번 뽑아 평균내면 참값으로 수렴한다는 뜻이다. 미니배치 확률적 경사하강법(SGD)이 전체 데이터의 그래디언트 대신 배치 평균을 쓰는 것과 같은 원리다.

실제 학습 루프

REINFORCE 한 걸음은 여섯 단계다.

단계 하는 일 비용
① 생성 현재 모델로 프롬프트마다 응답을 뽑는다 (샘플링 온도 1: 모델이 준 확률 그대로 뽑는다) 모델을 돌려 답을 뽑는 계산 — 가장 느린 단계
② 채점 각 응답에 보상 Ri\textcolor{#d9670b}{R_i}를 매긴다 (검증기, 보상 모델, 사람) 채점기에 따라 다름
③ 로그확률 응답의 토큰 로그확률을 더해 log⁡πθ(yi)\log \textcolor{#1565c0}{\pi_\theta}(\textcolor{#1c9c60}{y_i})를 구한다 순전파 1회
④ 대리 손실(기울기만 빌려 쓰는 손실) L=−1N∑iRilog⁡πθ(yi)\textcolor{#d62728}{\mathcal{L}} = -\frac{1}{\textcolor{#5f970c}{N}}\sum_i \textcolor{#d9670b}{R_i}\log \textcolor{#1565c0}{\pi_\theta}(\textcolor{#1c9c60}{y_i}) 무시할 만함
⑤ 역전파·갱신 ∇L\nabla \textcolor{#d62728}{\mathcal{L}}로 파라미터를 한 걸음 옮긴다 역전파 1회
⑥ 폐기 뽑았던 응답을 버리고 ①로 돌아간다 —

샘플링 온도를 1로 두는 까닭은 위 식의 조건 yi∼πθ\textcolor{#1c9c60}{y_i} \sim \textcolor{#1565c0}{\pi_\theta} 에 있다. 샘플링 온도를 낮춰 뽑으면 샘플이 πθ\textcolor{#1565c0}{\pi_\theta} 가 아닌 더 뾰족한 분포에서 나오고, 그 샘플의 평균은 더 이상 ∇J\nabla \textcolor{#d62728}{J} 를 겨누지 않는다. 코드로 쓰면 SFT(정답을 그대로 따라 쓰게 하는 지도 미세조정)에서 몇 줄 바뀔 뿐이다.

responses = policy.generate(prompts, do_sample=True)          # ① 생성
rewards = score(prompts, responses)                            # ② 채점, shape [N]
logp = policy.logprob(prompts, responses).sum(dim=-1)          # ③ 토큰 로그확률의 합, shape [N]
loss = -(rewards.detach() * logp).mean()                       # ④ 대리 손실
loss.backward(); optimizer.step(); optimizer.zero_grad()       # ⑤

⑥이 중요하다. 로그 미분 트릭은 "지금의 πθ\textcolor{#1565c0}{\pi_\theta}에서 뽑은 샘플"에만 성립하므로, 파라미터를 한 번 갱신하면 방금 뽑은 응답은 남의 응답이 된다. 비싸게 생성한 샘플을 그래디언트 한 걸음에 쓰고 버린다. 이 낭비가 REINFORCE의 큰 약점이다.

왜 이렇게 흔들리는가

첫머리의 장난감 문제로 재 보자. 처음 소프트맥스 정책이 세 답 “서울입니다”(1점), “부산입니다”(0점), “서울이요”(0.6점)에 약 0.14, 0.61, 0.25의 확률을 준다. 이때 "서울입니다"의 로짓(소프트맥스를 거쳐 확률이 되기 전의 값)에 대한 참 그래디언트는 +0.098+0.098이다. 샘플 하나로 추정하면 뽑힌 답에 따라 값이 달라진다. 이 추정치의 표준편차를 계산해 보면:

설정 추정치 평균 표준편차 신호 대 잡음비 잡음보다 신호가 커지려면 필요한 샘플 수
보상 그대로 0.098 0.307 0.32 10개
보상에 +5 0.098 2.024 0.05 430개

(신호 대 잡음비는 평균 ÷ 표준편차다. 추정치가 가리키는 값이 흔들림보다 얼마나 큰지를 잰다. 마지막 열은 (표준편차/평균)2(\text{표준편차}/\text{평균})^2, 즉 N\textcolor{#5f970c}{N}개를 평균내 표준편차가 1/N1/\sqrt{\textcolor{#5f970c}{N}}로 줄었을 때 평균과 같아지는 N\textcolor{#5f970c}{N}이다.)

평균은 두 줄 모두 똑같다. 달라지는 것은 흔들림의 폭, 즉 분산이다. 보상에 5를 더했을 뿐인데 같은 정확도를 얻으려면 샘플이 43배 필요하다. 샘플 하나가 LLM의 응답 생성 한 번이라는 걸 생각하면, 분산은 곧 GPU 비용이다.

샘플 하나의 추정치가 실제로 어떤 값들인지 펼쳐 보면 까닭이 보인다.

보상 그대로 (답 · 뽑힐 확률 → 추정치) 서울입니다 · 0.14 +0.86 부산입니다 · 0.61 0 서울이요 · 0.25 −0.08 평균 +0.098 · 표준편차 0.31 보상에 +5 서울입니다 · 0.14 +5.18 부산입니다 · 0.61 −0.68 서울이요 · 0.25 −0.77 평균 +0.098 · 표준편차 2.02

왜 상수가 분산을 키우는가? 추정치는 R⋅∇log⁡π(y)\textcolor{#d9670b}{R} \cdot \nabla \log \textcolor{#1565c0}{\pi}(\textcolor{#1c9c60}{y})다. ∇log⁡π\nabla \log \textcolor{#1565c0}{\pi}는 어떤 답이 뽑혔느냐에 따라 방향이 크게 바뀌는 벡터이고, R\textcolor{#d9670b}{R}은 그 벡터의 길이를 늘린다. 모든 답의 R\textcolor{#d9670b}{R}이 5 근처라면, 매 샘플이 "뽑힌 답을 5배 세게 올려라"라고 외친다. 좋은 답이 뽑히든 나쁜 답이 뽑히든. 그림 아래쪽에서 "부산입니다"가 뽑힌 샘플이 "서울입니다"의 로짓을 −0.68만큼 내리라고 외치는 것이 그 외침이다. 이 외침들은 평균내면 상쇄되지만(상수 보상의 기대 그래디언트는 0이다 — 확률의 합이 늘 1이므로), 하나하나는 크다.

문제 1 — 몇 명에게 물어야 하나

동아리 회장이 MT를 바다로 갈지 산으로 갈지 정하려고 회원에게 한 명씩 묻는다. 바다를 고른 답은 +1, 산을 고른 답은 −1로 적고 그 평균으로 판단한다. 회원 전체로는 60%가 바다파다. (가) 아무나 한 명에게 물었을 때 그 답의 평균과 표준편차는? (나) 몇 명에게 물어 평균을 내야 평균이 흔들림(표준편차 ÷ √명수)만큼 커지는가? (다) 바다파가 55%라면?

선생님 (평상)
선생님
한 명에게 물으면 답은 +1 아니면 −1이에요. 평균은 얼마일까요?
김민준 (평상)
김민준
0.6 × 1 + 0.4 × (−1) = 0.2요. 제곱의 평균은 늘 1이니까 분산은 1 − 0.04 = 0.96, 표준편차 0.98이에요.
선생님 (평상)
선생님
그럼 몇 명이면 될까요?
김민준 (평상)
김민준
흔들림이 0.98/√n이니까 0.2가 되려면 √n = 4.9, 24명이요. 55%면 평균이 0.1이라서… 절반이니까 48명쯤?
이서연 (의심)
이서연
반으로 줄었으면 두 배가 아니라 네 배 아니야? 명수는 (표준편차/평균)²이니까 제곱으로 들어가.
김민준 (당황)
김민준
아, (0.995/0.1)² = 99명이네요. 평균이 반이 되니 네 배가 필요해요.
이서연 (평상)
이서연
해석학에서 오차가 1/√n로 줄면 자릿수 하나 더 얻는 데 항이 100배 드는 거랑 같네.

정리 (가) 평균 0.2, 표준편차 0.98. (나) (0.98/0.2)2=24(0.98/0.2)^2 = 24명. (다) 평균 0.1, 표준편차 0.995, (0.995/0.1)2=99(0.995/0.1)^2 = 99명. 필요한 수는 (표준편차 ÷ 평균)의 제곱이라, 평균이 반이 되면 네 배를 물어야 한다.

문제 2 — 0/100점으로 바꾸면

채점 기준을 0/1점에서 0/100점으로 바꿨다. 점수를 그대로 로그확률 앞에 곱하고 학습률도 그대로 두면 무엇이 달라지는가? (가) 그래디언트의 기대 방향 (나) 그래디언트의 표준편차 (다) 신호 대 잡음비 (라) 실제 학습 결과

선생님 (평상)
선생님
점수를 100배로 부풀렸어요. 무엇이 바뀌었을까요?
이서연 (평상)
이서연
방향은 그대로예요. 기대 그래디언트가 100배가 될 뿐이니까 같은 방향이고요. 그러니까 아무 문제 없어요.
김민준 (평상)
김민준
저는 반대로 생각했어요. 표준편차가 100배, 분산은 10000배잖아요. 훨씬 나빠졌죠.
선생님 (질문)
선생님
서연 학생, 방향이 같으면 학습률 1e-5로 한 걸음 가는 거랑 1e-3으로 한 걸음 가는 게 같은 결과인가요?
이서연 (생각)
이서연
아니요… 같은 방향이어도 보폭이 100배예요. 사실상 학습률을 100배 올린 거랑 같아요. 값이 한없이 커질 수 있어요.
선생님 (평상)
선생님
그럼 민준 학생, 분산이 10000배면 추정이 100배 나빠진 건가요?
김민준 (평상)
김민준
평균도 100배, 표준편차도 100배니까… 비율은 그대로네요. 신호 대 잡음비는 안 변했어요.
선생님 (미소)
선생님
그래요. 서연 학생은 "방향"만 보고 보폭을 놓쳤고, 민준 학생은 "분산"만 보고 신호도 같이 커진 걸 놓쳤어요. 실제로 바뀐 건 유효 학습률 하나예요.
김민준 (평상)
김민준
아까 MT 투표에서 필요한 인원이 (표준편차/평균)²였잖아요. 둘 다 100배면 필요한 샘플 수도 그대로네요.
이서연 (평상)
이서연
해석학 수업에서 f를 100배 해도 극값의 위치는 그대로지만, 경사하강법 스텝은 달라지는 거랑 같네요.

정리 (가) 같다. (나) 100배. (다) 변하지 않는다. (라) 유효 학습률이 100배가 되어 학습이 불안정해질 수 있다. 보상의 규모에 따라 학습률을 다시 맞춰야 한다 — 보상을 표준편차로 나눠 규모를 맞추는 방법이 쓰이는 이유 중 하나다.