Chapter 5: PPO와 RLHF — 멀리 가지 않게 묶다

의문

비평가(critic: 앞으로 받을 점수를 예측하는 보조 모델)를 두면 토큰마다 어드밴티지 A^t\hat A_t(그 토큰이 그 시점의 평균적인 선택보다 얼마나 나았나)를 추정할 수 있고, 분산도 준다. 그런데 언어모델에 이 방법을 그대로 쓰면 또 하나의 벽에 부딪힌다: 샘플이 너무 비싸다.

REINFORCE도 Actor-Critic도 “지금의 정책(프롬프트를 받아 답을 고르는 언어모델)으로 응답을 뽑고, 그래디언트를 한 번 계산하고, 버린다.” 응답 하나를 생성하려면 토큰 수만큼 모델을 순차적으로 돌려야 한다. RLHF(사람의 선호로 학습한 보상 모델의 점수를 강화학습으로 올리는 절차) 학습에서 실제로 걸리는 시간의 대부분은 역전파가 아니라 생성에 쓰인다. 오픈소스 학습 틀 OpenRLHF를 만든 팀은, 긴 풀이를 쓰는 요즘 모델에서는 생성 단계가 전체 학습 시간의 90% 넘게 차지할 때가 많다고 적었다. 그렇게 비싸게 만든 배치를 한 번의 업데이트에 쓰고 버리는 것은 아깝다.

그래서 묻는다:

한 번 뽑은 배치로 여러 번 업데이트할 수는 없는가?
첫 업데이트가 끝나는 순간 정책은 바뀐다. 배치는 이전 정책이 만든 것이 된다.
정책 그래디언트의 기대값은 지금 정책이 뽑은 샘플 위에서 취한 것이라, 다른 분포에서 뽑은 샘플로는 식이 틀어진다.
이것을 어떻게 보정하고, 얼마나 멀리까지 믿을 수 있는가?

그리고 대화형 모델에는 두 번째 질문이 겹친다: 수학 채점기가 없는 곳에서 점수 RR은 누가 매기는가? 이 둘에 대한 답이 PPO와 RLHF이고, 2022년 ChatGPT를 만든 조합이다.

임포턴스 샘플링: 다른 분포의 샘플로 기대값 구하기

한 배치를 여러 번 쓰려면 먼저 이 물음에 답해야 한다. 다른 분포가 뽑은 샘플로, 내가 알고 싶은 분포의 평균을 구할 수 있는가? 언어모델에서 바로 따지면 응답의 가짓수가 너무 많다. 답을 미리 알고 있는 가장 작은 경우로 옮겨 보자. 주사위다. 공정한 주사위의 평균이 3.5라는 것을 이미 알기 때문에, 어떤 방법이 맞는지 틀리는지를 숫자로 바로 확인할 수 있다.

가중치로 되돌리는 기대값

장난감 예제 — 주사위. 공정한 주사위(각 면 1/6)의 평균 3.5를 알고 싶은데, 손에는 1이 잘 나오는 찌그러진 주사위만 있다. 찌그러진 주사위로 여섯 번 던져 [1,1,2,1,3,4][1, 1, 2, 1, 3, 4]가 나왔다. 그냥 평균을 내면 (1+1+2+1+3+4)/6=2.0(1+1+2+1+3+4)/6 = 2.0이다. 1이 자주 나오는 주사위라 낮게 치우친다.

어떻게 고칠까? 찌그러진 주사위는 1을 0.40의 확률로 낸다. 공정한 주사위의 1/6보다 2.4배 자주 나오니, 1이 나올 때마다 한 번이 아니라 1/2.4≈0.421/2.4 \approx 0.42번으로만 센다. 거꾸로 6은 0.05의 확률이라 공정한 주사위보다 3.33배 드물게 나온다. 한 번 나오면 3.33번으로 센다. 눈마다 이렇게 고쳐 셀 무게를 적으면 표의 셋째 줄이 된다.

눈 1 2 3 4 5 6
공정한 주사위 1/6 1/6 1/6 1/6 1/6 1/6
찌그러진 주사위 0.40 0.20 0.15 0.10 0.10 0.05
고쳐 셀 무게 (공정 ÷ 찌그러짐) 0.42 0.83 1.11 1.67 1.67 3.33

여섯 번의 눈에 무게를 곱해 더하면 0.42×3+0.83×2+1.11×3+1.67×4≈12.920.42 \times 3 + 0.83 \times 2 + 1.11 \times 3 + 1.67 \times 4 \approx 12.92다. 이것을 무엇으로 나누느냐에 따라 추정이 둘 나온다.

왜 둘이 다를까? 무게는 길게 보면 평균이 1이 되도록 만든 수다. 눈마다 「찌그러진 주사위가 그 눈을 낼 확률 × 무게」를 더하면 1/6을 여섯 번 더한 1이 된다. 그래서 아주 많이 던지면 무게의 합이 던진 횟수와 거의 같아지고, 두 추정도 같아진다. 그런데 이번 여섯 번에는 무게가 큰 5와 6이 한 번도 나오지 않아 무게의 합이 6이 아니라 4.86밖에 안 된다. 횟수로 나누면 나오지 않은 눈의 몫을 0으로 세고, 무게의 합으로 나누면 실제로 센 무게끼리만 평균을 낸다. 뒤의 것을 정규화한 추정이라 부른다.

방금 한 일을 기호로 쓰면 이렇다. 평균을 알고 싶은 분포(공정한 주사위)를 πθ\pi_\theta, 샘플을 뽑은 분포(찌그러진 주사위)를 πold\pi_\text{old}, 평균을 내려는 값(눈)을 f(y)f(y)라 하자.

Ey∼πθ[f(y)]=∑yπold(y) πθ(y)πold(y) f(y)=Ey∼πold ⁣[πθ(y)πold(y) f(y)]\mathbb{E}_{\textcolor{#1c9c60}{y} \sim \textcolor{#1565c0}{\pi_\theta}}[\textcolor{#206049}{f}(\textcolor{#1c9c60}{y})] = \sum_{\textcolor{#1c9c60}{y}} \textcolor{#6f6f78}{\pi_\text{old}}(\textcolor{#1c9c60}{y})\,\frac{\textcolor{#1565c0}{\pi_\theta}(\textcolor{#1c9c60}{y})}{\textcolor{#6f6f78}{\pi_\text{old}}(\textcolor{#1c9c60}{y})}\,\textcolor{#206049}{f}(\textcolor{#1c9c60}{y}) = \mathbb{E}_{\textcolor{#1c9c60}{y} \sim \textcolor{#6f6f78}{\pi_\text{old}}}\!\Big[\frac{\textcolor{#1565c0}{\pi_\theta}(\textcolor{#1c9c60}{y})}{\textcolor{#6f6f78}{\pi_\text{old}}(\textcolor{#1c9c60}{y})}\, \textcolor{#206049}{f}(\textcolor{#1c9c60}{y})\Big]
πθ지금 학습 중인 정책πold데이터를 만든 옛 정책f(y)기대값을 구하려는 아무 함수 \small\begin{array}{ll} \textcolor{#1565c0}{\pi_\theta} & \text{지금 학습 중인 정책} \\ \textcolor{#6f6f78}{\pi_\text{old}} & \text{데이터를 만든 옛 정책} \\ \textcolor{#206049}{f}(\textcolor{#1c9c60}{y}) & \text{기대값을 구하려는 아무 함수} \end{array}

w=πθ(y)/πold(y)\textcolor{#c2185b}{w} = \textcolor{#1565c0}{\pi_\theta}(\textcolor{#1c9c60}{y})/\textcolor{#6f6f78}{\pi_\text{old}}(\textcolor{#1c9c60}{y})가 임포턴스 가중치(샘플마다 곱하는 수 — 신경망 파라미터가 아니다)다. 표의 셋째 줄이 바로 이 값이었다. 직관: “πold\textcolor{#6f6f78}{\pi_\text{old}}가 이 샘플을 너무 자주 뽑았으면 가중치를 낮추고, 너무 드물게 뽑았으면 가중치를 높여서, πθ\textcolor{#1565c0}{\pi_\theta}의 기대값에 맞춘다.”

그런데 문제가 보인다. 6이 한 번도 안 나왔다. πold(6)=0.05\textcolor{#6f6f78}{\pi_\text{old}}(6) = 0.05라 잘 안 뽑힌다. 만약 나왔다면 그 한 번은 가중치 3.33으로 세어진다. 드문 샘플은 드물게 나오는 대신, 나오면 크게 세어진다. 샘플 수가 적은데 가중치가 크게 들쭉날쭉하면 추정값이 실험마다 크게 달라진다. 이것을 분산이 터진다(분산 폭발)고 한다.

위젯에서 πold를 "극단"으로 바꿔보라. 6이 나올 확률이 0.5%라 가중치가 33이 된다. 보정 없는 평균은 늘 틀린 곳에 몰려 있고(치우침), 임포턴스 가중치는 평균은 맞추지만 넓게 퍼진다(분산). 가중치를 잘라내면 퍼짐은 줄어드는 대신 다시 치우친다. 맨 아래 줄의 이 절충이 곧 PPO의 선택으로 이어진다. ‘던진 눈’ 칸에 눈을 적으면 그 눈으로 낸 추정값들이 나온다. 단추는 본문의 여섯 번과 아래 문제 2의 일곱 번을 넣는다.

ML에서: 토큰 단위 비율

언어모델로 돌아오자. 한 배치로 두 번째 업데이트를 할 때, 데이터를 뽑은 것은 업데이트 전의 정책 πold\textcolor{#6f6f78}{\pi_\text{old}}이고 기대값이 필요한 것은 지금 정책 πθ\textcolor{#1565c0}{\pi_\theta}다. 주사위에서 찌그러진 주사위와 공정한 주사위의 자리다. 데이터를 뽑은 정책이 지금 학습하는 정책과 같으면 온폴리시(on-policy), 다르면 오프폴리시(off-policy)라 한다. 오프폴리시가 되면 데이터가 나온 분포와 기대값을 취할 분포가 어긋나는데, 이 어긋남을 분포 이동(distribution shift)이라 부른다. (데이터를 학습 도중 새로 만드느냐, 학습 전에 만들어 둔 것만 쓰느냐는 또 다른 축이다. 앞의 것을 온라인, 뒤의 것을 오프라인이라 한다.)

언어모델에서는 훨씬 심하다. 주사위는 면이 6개지만, 어휘가 약 3만 2천 개인 모델(라마 2의 어휘가 이 크기다)로 500토큰짜리 응답을 쓰면 응답의 가짓수는 3200050032000^{500} 수준이다. 게다가 응답 전체의 확률은 토큰 확률의 곱이므로, 응답 단위 가중치도 토큰 비율의 곱이 된다.

πθ(y)πold(y)=∏t=1Tπθ(yt∣st)πold(yt∣st)\frac{\textcolor{#1565c0}{\pi_\theta}(\textcolor{#1c9c60}{y})}{\textcolor{#6f6f78}{\pi_\text{old}}(\textcolor{#1c9c60}{y})} = \prod_{t=1}^{\textcolor{#915a08}{T}} \frac{\textcolor{#1565c0}{\pi_\theta}(\textcolor{#1c9c60}{y_t} \mid \textcolor{#0093b8}{s_t})}{\textcolor{#6f6f78}{\pi_\text{old}}(\textcolor{#1c9c60}{y_t} \mid \textcolor{#0093b8}{s_t})}
πθ(yt∣st)상태 st에서 토큰 yt를 고를 확률πold옛 정책T응답 길이\small\begin{array}{ll} \textcolor{#1565c0}{\pi_\theta}(\textcolor{#1c9c60}{y_t} \mid \textcolor{#0093b8}{s_t}) & \text{상태 }\textcolor{#0093b8}{s_t}\text{에서 토큰 }\textcolor{#1c9c60}{y_t}\text{를 고를 확률} \\ \textcolor{#6f6f78}{\pi_\text{old}} & \text{옛 정책} \\ \textcolor{#915a08}{T} & \text{응답 길이} \end{array}

토큰마다 비율이 1.01 — 겨우 1% 차이 — 이어도 500토큰이면 1.01500≈1451.01^{500} \approx 145다. 반대 방향이면 거의 0이다. 한두 개 응답이 배치 전체의 그래디언트를 지배하게 된다.

그래서 실제 알고리즘은 응답 단위 가중치를 쓰지 않고 토큰 단위 비율만 쓴다.

rt(θ)=πθ(yt∣st)πold(yt∣st),L(θ)=Eπold[ rt(θ) A^t ]\textcolor{#c2185b}{r_t(\theta)} = \frac{\textcolor{#1565c0}{\pi_\theta}(\textcolor{#1c9c60}{y_t} \mid \textcolor{#0093b8}{s_t})}{\textcolor{#6f6f78}{\pi_\text{old}}(\textcolor{#1c9c60}{y_t} \mid \textcolor{#0093b8}{s_t})}, \qquad \textcolor{#d62728}{L}(\theta) = \mathbb{E}_{\textcolor{#6f6f78}{\pi_\text{old}}}\big[\, \textcolor{#c2185b}{r_t(\theta)}\, \textcolor{#8e44ad}{\hat A_t} \,\big]
rt(θ)토큰 확률 비율 (보상 rt와 다르다)L(θ)대리 목적함수 (기울기만 빌려 쓰는 목적함수)A^t어드밴티지 (가까운 미래는 실제 결과, 먼 미래는 비평가 예측으로 섞어 구한 값, GAE)πold샘플을 만든 옛 정책 \small\begin{array}{ll} \textcolor{#c2185b}{r_t(\theta)} & \text{토큰 확률 비율 (보상 }\textcolor{#d9670b}{r_t}\text{와 다르다)} \\ \textcolor{#d62728}{L}(\theta) & \text{대리 목적함수 (기울기만 빌려 쓰는 목적함수)} \\ \textcolor{#8e44ad}{\hat A_t} & \text{어드밴티지 (가까운 미래는 실제 결과, 먼 미래는 비평가 예측으로 섞어 구한 값, GAE)} \\ \textcolor{#6f6f78}{\pi_\text{old}} & \text{샘플을 만든 옛 정책} \end{array}

(기호 주의: 이 교재에서 비율은 항상 θ\theta를 붙여 rt(θ)\textcolor{#c2185b}{r_t(\theta)}로, 토큰 보상은 θ\theta 없이 rt\textcolor{#d9670b}{r_t}로 쓴다. 두 기호 모두 관례라 피하기 어렵다.)

이것은 근사다. 앞선 토큰들이 바뀌어 "어떤 상태에 도달하는가"가 달라지는 효과를 무시했다. 이 근사는 πθ\textcolor{#1565c0}{\pi_\theta}가 πold\textcolor{#6f6f78}{\pi_\text{old}}에 가까울 때만 괜찮다. 응답이 길수록 앞선 토큰들의 변화가 쌓이므로 이 근사는 더 쉽게 무너진다.

덧붙여, θ=θold\theta = \theta_\text{old}인 지점에서 ∇θ rt=∇θlog⁡πθ(yt∣st)\nabla_\theta\, \textcolor{#c2185b}{r_t} = \nabla_\theta \log \textcolor{#1565c0}{\pi_\theta}(\textcolor{#1c9c60}{y_t} \mid \textcolor{#0093b8}{s_t}) 이므로 이 대리 목적함수의 기울기는 그 지점에서 어드밴티지를 빈칸에 넣은 Actor-Critic 그래디언트 E[∑tA^t∇θlog⁡πθ]\mathbb{E}[\sum_t \textcolor{#8e44ad}{\hat A_t} \nabla_\theta \log\textcolor{#1565c0}{\pi_\theta}] 와 정확히 같다.

문제 1 — 학교 앞 설문으로 동네 평균 내기

동네 주민이 하루에 커피를 몇 잔 마시는지 평균을 알고 싶어, 학교 앞에서 100명에게 물었다. 대학생이 80명(평균 2잔), 직장인이 20명(평균 1잔)이었다. 이 동네 주민은 대학생 20%, 직장인 80%다. (가) 응답자마다 무게를 달아 동네 평균을 추정하시오. 대학생과 직장인의 무게는 각각 얼마인가? (나) 같은 설문으로 이웃 동네(대학생 20%, 직장인 50%, 어르신 30%)의 평균도 추정하려 한다. 무게를 다시 달면 되는가?

선생님 (평상)
선생님
학교 앞에서 물었으니 대학생이 많이 잡혔죠. 그냥 평균을 내면 얼마가 나오고, 무엇을 고쳐야 할까요?
김민준 (평상)
김민준
그냥 평균은 (80 × 2 + 20 × 1)/100 = 1.8잔이요. 대학생이 동네 비율보다 네 배 많이 잡혔으니 높게 나왔겠네요.
이서연 (평상)
이서연
주사위랑 똑같이 하면 돼. 대학생 무게는 동네 비율 ÷ 표본 비율 = 0.2/0.8 = 0.25, 직장인은 0.8/0.2 = 4. 무게를 곱해 더하면 (80 × 0.25 × 2 + 20 × 4 × 1)/100 = 1.2잔.
선생님 (평상)
선생님
무게의 합으로 나눠도 같나요?
이서연 (평상)
이서연
무게의 합이 80 × 0.25 + 20 × 4 = 100이라 같아요. 동네 평균 0.2 × 2 + 0.8 × 1 = 1.2와도 맞고요.
선생님 (질문)
선생님
그럼 (나), 이웃 동네는요?
김민준 (평상)
김민준
어르신 무게를 크게 주면 되죠. 30%인데 표본에는…
김민준 (당황)
김민준
0명이네요. 0.3 ÷ 0이라 무게를 정할 수가 없어요.
이서연 (평상)
이서연
무게를 아무리 크게 줘도 곱할 사람이 없으면 0이야. 표본에 한 번도 안 나온 쪽은 어떤 무게로도 못 되살려.
김민준 (깨달음)
김민준
그럼 찌그러진 주사위에 6이 아예 없었다면, 3.5는 영영 못 맞히는 거네요.
선생님 (흐뭇함)
선생님
그래요. 가중치로 고칠 수 있는 것은 ‘드물게’ 뽑히는 쪽까지예요. ‘한 번도 뽑힐 수 없는’ 쪽은 못 고쳐요.

정리 (가) 대학생 무게 0.25, 직장인 무게 4. 추정 (80×0.25×2+20×4×1)/100=1.2(80 \times 0.25 \times 2 + 20 \times 4 \times 1)/100 = 1.2잔(그냥 평균은 1.8잔). (나) 안 된다. 표본에 0명인 어르신은 무게가 「30% ÷ 0」이라 정해지지 않고, 곱할 응답도 없다. 가중치 보정은 샘플을 뽑은 분포가 0이 아닌 곳에서만 된다 — πθ(y)>0\textcolor{#1565c0}{\pi_\theta}(\textcolor{#1c9c60}{y}) > 0 인 모든 y\textcolor{#1c9c60}{y} 에서 πold(y)>0\textcolor{#6f6f78}{\pi_\text{old}}(\textcolor{#1c9c60}{y}) > 0 이어야 한다.

문제 2 — 일곱 번째에 6이 나오면

본문의 찌그러진 주사위로 여섯 번 던져 [1,1,2,1,3,4][1, 1, 2, 1, 3, 4]가 나온 데 이어, 일곱 번째에 6이 나왔다. (가) 무게의 합으로 나누는 추정값을 다시 구하시오. (나) 이 한 번의 6이 추정에서 차지하는 몫(그 무게 ÷ 무게의 합)은 얼마인가? 일곱 번 가운데 한 번이니 1/7인가? (다) 일곱 번째가 1이었다면 추정값은 얼마인가? 한 번의 던지기가 추정을 얼마나 흔드는가?

위젯의 ‘문제 2’ 단추가 이 일곱 번을 넣는다. 풀이를 마친 뒤 수치판과 견주어 보라.

김민준 (자신만만)
김민준
무게 3.33을 더하면 돼요. 분자는 12.92 + 3.33 × 6 = 32.92, 분모는 4.86 + 3.33 = 8.19, 그러니까 4.02. 2.66보다 3.5에 가까워졌으니 6이 나와서 좋아졌네요.
선생님 (질문)
선생님
가까워졌나요? 4.02는 3.5의 어느 쪽이죠?
김민준 (평상)
김민준
위쪽이요. 0.84 모자라던 게 한 번 만에 0.52 넘치는 쪽으로 건너갔네요.
이서연 (평상)
이서연
그래도 무게의 합으로 나누니까 튀는 걸 눌러 주잖아. 횟수로 나누면 32.92/7 = 4.70이니까.
선생님 (평상)
선생님
눌러 주긴 했죠. 그럼 (나), 그 6 한 번의 몫은요?
이서연 (평상)
이서연
3.33 ÷ 8.19 ≈ 0.41.
이서연 (놀람)
이서연
일곱 번 중 한 번인데 추정의 41%를 혼자 정하네요. 공정한 주사위였다면 1/7, 14%였을 텐데.
선생님 (평상)
선생님
(다)도 해 볼까요?
김민준 (평상)
김민준
일곱 번째가 1이면 분자 12.92 + 0.42 = 13.33, 분모 4.86 + 0.42 = 5.28, 2.53이요. 한 번이 1이냐 6이냐로 2.53과 4.02, 1.5 가까이 갈리네요.
이서연 (평상)
이서연
공정한 주사위를 그냥 평균 내면 1이냐 6이냐로 (6 − 1)/7 ≈ 0.71 갈릴 뿐인데, 두 배 넘게 흔들리는 거야.
선생님 (흐뭇함)
선생님
그게 보정의 값이에요. 치우침을 없애는 대신 흔들림을 받아요.
김민준 (평상)
김민준
조별과제 평가에서 한 명의 점수만 네 배로 쳐 주면, 그 사람 컨디션 따라 조 점수가 널뛰는 거랑 같네요.

정리 (가) 32.92/8.19≈4.0232.92 / 8.19 \approx 4.02 (횟수로 나누면 4.70). 3.5를 넘어 반대편으로 건너갔다. (나) 3.33/8.19≈0.413.33 / 8.19 \approx 0.41 — 1/7 ≈ 0.14의 세 배다. (다) 일곱 번째가 1이면 2.53. 한 번의 던지기로 추정이 1.5 가까이 갈린다(공정한 주사위의 그냥 평균은 0.71). 드문 샘플의 큰 가중치가 분산을 키운다.