2. VPG — 미분할 수 없는 점수를 올리는 법

소프트맥스 정책의 그래디언트: 평균보다 나은 답이 오른다

식은 얻었다. 그런데 이 식은 실제로 어떤 답을 올리고 어떤 답을 내릴까? 점수가 가장 높은 답만 오를까, 점수가 양수인 답은 다 오를까? 답이 셋뿐인 문제에서 손으로 계산해 보자.

장난감 예제: 세 개의 답

"대한민국의 수도는?"에 대해 모델이 세 가지 답만 할 수 있다고 하자. 정책은 로짓 zkz_k의 소프트맥스 πk=ezk/∑jezj\pi_k = e^{z_k} / \sum_j e^{z_j} 이다. 소프트맥스의 로그확률 미분은 깔끔하다:

∂log⁡πa∂zk=1[k=a]−πk\frac{\partial \log \textcolor{#1565c0}{\pi_a}}{\partial \textcolor{#cc00ff}{z_k}} = \mathbb{1}[k = a] - \textcolor{#1565c0}{\pi_k}
πa, πk답 a, k를 고를 확률 (소프트맥스 출력)zk답 k의 로짓 — 이 예제에서 학습하는 파라미터1[k=a]k=a이면 1, 아니면 0 \small\begin{array}{ll} \textcolor{#1565c0}{\pi_a},\ \textcolor{#1565c0}{\pi_k} & \text{답 }a\text{, }k\text{를 고를 확률 (소프트맥스 출력)} \\ \textcolor{#cc00ff}{z_k} & \text{답 }k\text{의 로짓 — 이 예제에서 학습하는 파라미터} \\ \mathbb{1}[k = a] & \text{}k = a\text{이면 1, 아니면 0} \end{array}

“뽑힌 답의 로짓은 (1−πa)(1-\textcolor{#1565c0}{\pi_a})만큼 올리고, 나머지는 자기 확률만큼 내린다.” 이것을 기대값으로 모으면 로짓별 정확한 그래디언트가 나온다.

∂J∂zk=πk (rk−J)\frac{\partial \textcolor{#d62728}{J}}{\partial \textcolor{#cc00ff}{z_k}} = \textcolor{#1565c0}{\pi_k}\,\big(\textcolor{#d9670b}{r_k} - \textcolor{#d62728}{J}\big)
J기대 보상 ∑kπkrkrk답 k의 보상πk답 k를 고를 확률zk답 k의 로짓 \small\begin{array}{ll} \textcolor{#d62728}{J} & \text{기대 보상 }\sum_k \textcolor{#1565c0}{\pi_k} \textcolor{#d9670b}{r_k} \\ \textcolor{#d9670b}{r_k} & \text{답 }k\text{의 보상} \\ \textcolor{#1565c0}{\pi_k} & \text{답 }k\text{를 고를 확률} \\ \textcolor{#cc00ff}{z_k} & \text{답 }k\text{의 로짓} \end{array}

처음 로짓은 (z1,z2,z3)=(−0.6, 0.9, 0)(\textcolor{#cc00ff}{z_1}, \textcolor{#cc00ff}{z_2}, \textcolor{#cc00ff}{z_3}) = (-0.6,\ 0.9,\ 0) 으로 둔다. 틀린 답(“부산입니다”)을 가장 믿는 모델에서 출발해야, 그래디언트가 무엇을 올리고 무엇을 내리는지 한눈에 보이기 때문이다. 이 로짓을 소프트맥스에 넣은 처음 확률이 아래 표의 0.14, 0.61, 0.25다.

답 보상 rk\textcolor{#d9670b}{r_k} 처음 확률 πk\textcolor{#1565c0}{\pi_k} ∂J/∂zk\partial \textcolor{#d62728}{J} / \partial \textcolor{#cc00ff}{z_k} 방향
서울입니다 1.0 0.14 0.14×(1.0−0.29)=+0.100.14 \times (1.0 - 0.29) = +0.10 ↑
부산입니다 0.0 0.61 0.61×(0.0−0.29)=−0.180.61 \times (0.0 - 0.29) = -0.18 ↓
서울이요 0.6 0.25 0.25×(0.6−0.29)=+0.080.25 \times (0.6 - 0.29) = +0.08 ↑

(처음 기대 보상 J=0.14×1+0.61×0+0.25×0.6≈0.29\textcolor{#d62728}{J} = 0.14 \times 1 + 0.61 \times 0 + 0.25 \times 0.6 \approx 0.29)

눈여겨볼 것: “서울이요”(0.6점)도 올라간다. 절대 점수가 아니라 현재 평균 J\textcolor{#d62728}{J}보다 나은가가 방향을 정한다. 그리고 가장 좋은 답인 "서울입니다"보다 "부산입니다"를 내리는 힘이 더 크다 — 확률이 큰 답일수록 그래디언트에 곱해지는 πk\textcolor{#1565c0}{\pi_k}가 크기 때문이다.

위젯에서 해 볼 것 두 가지:

VPG가 실제로는 쓸 수 없는 이유

장난감 예제에서는 답이 셋뿐이라 ∑y\sum_{\textcolor{#1c9c60}{y}}를 직접 계산했다. 언어모델의 응답 공간은 어휘 크기를 문장 길이만큼 거듭 곱한 크기다. 메타가 2023년에 공개한 언어모델 라마 2(Llama 2)는 어휘가 약 32,000개인데, 이런 모델이 500토큰짜리 응답을 쓴다면 가짓수는 3200050032000^{500}이다. 기대값을 정확히 계산하는 것은 불가능하다.

그러니 뽑아서 평균낸다. 몇 개의 응답을 샘플링하고, 각자의 R⋅∇log⁡π\textcolor{#d9670b}{R} \cdot \nabla \log \textcolor{#1565c0}{\pi}를 평균내면 기대값의 추정치가 된다. 이것이 REINFORCE다. 그리고 위젯에서 이미 봤듯, 추정치는 흔들린다.

문제 4 — 조별과제 기여도 고쳐 적기

세 조원 A, B, C의 조별과제 기여도를 합이 100%가 되게 50%, 30%, 20%로 적었다. C가 자기 몫은 30%라고 해서 C만 30%로 고쳐 적었다. 합을 100%로 맞추려고, A와 B는 둘 사이의 원래 비율(50 : 30)을 지킨 채 줄였다. (가) A와 B는 각각 몇 %가 되는가? (나) A와 B 가운데 누가 더 많이 깎였는가? 깎인 사람이 뭔가 잘못한 것인가?

선생님 (평상)
선생님
민준 학생, A와 B는 몇 %가 됐나요?
김민준 (평상)
김민준
C가 10%p 늘었으니까 A, B가 5%p씩 내놓으면 되죠. 45%랑 25%요.
이서연 (평상)
이서연
그러면 45 : 25 = 9 : 5 야. 원래 50 : 30 = 5 : 3 이 안 지켜져.
김민준 (난처함)
김민준
아, 남은 70%를 5 : 3으로 나눠야 하는구나. 70 × 5/8 = 43.75%, 70 × 3/8 = 26.25%.
선생님 (질문)
선생님
그럼 누가 더 많이 깎였죠? A가 뭘 잘못해서 깎였나요?
이서연 (평상)
이서연
A가 6.25%p, B가 3.75%p요. 몫이 큰 사람이 더 많이 깎여요. 잘못한 건 없고, 합이 100%로 묶여 있어서 C가 는 만큼 나머지가 원래 몫에 비례해 줄었을 뿐이에요.

정리 (가) A 43.75%, B 26.25%. (나) A가 6.25%p, B가 3.75%p로 몫이 큰 A가 더 많이 깎였다. 잘못한 사람은 없다. 합이 100%로 묶여 있어서 한 사람의 몫이 늘면 나머지는 원래 몫에 비례해 준다.

문제 5 — 한 샘플의 그래디언트

세 답의 로짓이 모두 0이고, 답 A를 뽑아 보상 R=1\textcolor{#d9670b}{R} = 1을 받았다. 이 한 샘플의 그래디언트 R⋅∇zlog⁡πA\textcolor{#d9670b}{R} \cdot \nabla_z \log \textcolor{#1565c0}{\pi_A}를 구하시오. 그래디언트 방향으로 한 걸음 가면 B, C의 확률은 어떻게 되는가? (위젯의 「문제 5」 단추는 세 로짓을 0으로 둔다. 「샘플 1개로 한 걸음」을 눌러 수치판의 그래디언트와 막대를 자기 풀이와 견주어 보라. 첫 샘플은 A다.)

김민준 (자신만만)
김민준
이건 바로 돌려봤어요. 로짓 0,0,0이면 확률이 1/3씩이고, 소프트맥스 로그확률 미분이 원-핫 빼기 확률이니까 (1 − 1/3, −1/3, −1/3) = (0.67, −0.33, −0.33). R=1 곱해도 그대로예요.
선생님 (평상)
선생님
좋아요. 그럼 한 걸음 가면 B와 C는요?
김민준 (평상)
김민준
둘 다 로짓이 내려가니까… 둘 다 확률이 떨어지죠.
이서연 (의심)
이서연
잠깐. B랑 C는 한 번도 안 뽑혔잖아. 안 뽑힌 답이 벌을 받는 거야?
선생님 (질문)
선생님
서연 학생, 세 성분을 더하면 얼마죠?
이서연 (평상)
이서연
0.67 − 0.33 − 0.33 = 0. 아, 로짓을 전부 같은 만큼 올리면 소프트맥스는 안 변하니까, 합이 0인 방향으로만 움직이는 거구나.
이서연 (깨달음)
이서연
그러니까 B, C를 따로 벌주는 게 아니라, A를 올린 만큼 나머지 몫이 줄어드는 거네. 확률 합이 1이니까.
김민준 (평상)
김민준
아까 조별과제 기여도 문제랑 같네요. 한 명 몫을 올려 적으면 합이 100%라 나머지가 자동으로 깎이는.

정리 R⋅∇zlog⁡πA=(+23,−13,−13)\textcolor{#d9670b}{R} \cdot \nabla_z \log \textcolor{#1565c0}{\pi_A} = (+\tfrac23, -\tfrac13, -\tfrac13). B, C의 확률은 줄지만 "벌"이 아니라 확률 합이 1이라는 제약 때문이다. 소프트맥스 그래디언트의 성분 합은 항상 0이다.

문제 6 — 모두가 +5점이라면

모든 답의 보상이 똑같이 +5+5라면 정확한 정책 그래디언트 ∇θJ\nabla_\theta \textcolor{#d62728}{J}는 얼마인가? 샘플 하나로 추정한 그래디언트는 어떤가? (위젯의 「문제 6」 단추는 세 답의 보상을 0으로 두고 상수 c=5c = 5 를 더해, 모든 답이 5점인 상황을 만든다. 두 방식으로 걸어 보며 견주어 보라.)

선생님 (평상)
선생님
이번엔 모든 답이 5점이에요. 정확한 그래디언트는요?
김민준 (평상)
김민준
모든 답이 양수 보상이니까 전부 올라가요. 식이 R·∇log π 인데 R이 다 양수잖아요.
이서연 (평상)
이서연
전부 올라갈 수는 없어. 확률 합이 1이야.
김민준 (당황)
김민준
어… 그럼 뭐가 올라가는데?
이서연 (평상)
이서연
식으로 보자. 기대값이 Σ π(y)·5·∇log π(y) = 5·Σ ∇π(y) = 5·∇(Σ π(y)) = 5·∇1 = 0. 정확한 그래디언트는 0이야.
선생님 (미소)
선생님
서연 학생 증명이 정확해요. 그럼 민준 학생 직관은 어디서 틀렸을까요? 샘플 하나만 보면요?
김민준 (생각)
김민준
샘플 하나만 보면… 뽑힌 답은 5배 세게 올라가고 나머지는 내려가요. 그러니까 샘플 하나로는 0이 아니고, 뭐가 뽑혔느냐에 따라 매번 다른 방향으로 크게 튀네요.
선생님 (평상)
선생님
바로 그거예요. 평균은 0인데 한 번 한 번은 크게 흔들려요. 보상에 상수를 더하면 방향은 안 바뀌고 흔들림만 커지는 거죠.
이서연 (평상)
이서연
수업에서 본 거랑 같다. 확률변수에 상수를 곱하면 평균은 그대로 0이어도 분산은 제곱으로 커지는 거.

정리 E[∇log⁡π]=∇∑yπ(y)=0\mathbb{E}[\nabla \log \textcolor{#1565c0}{\pi}] = \nabla \sum_{\textcolor{#1c9c60}{y}} \textcolor{#1565c0}{\pi}(\textcolor{#1c9c60}{y}) = 0 이므로 정확한 그래디언트는 0. 샘플 하나의 추정치는 0이 아니며, 상수가 클수록 크게 흔들린다(분산). 위젯의 「문제 6」 단추로 두 방식을 걸어 보고, cc를 줄였다 키웠다 하며 확인해 보라.

문제 7 — (킬러) loss = -R * logprob(y)

민준이 정책 그래디언트를 이렇게 구현했다: 모델로 응답 y\textcolor{#1c9c60}{y}를 뽑고, loss = -R * logprob(y) 를 역전파한다. 학습 로그에서 loss 값이 꾸준히 내려가니 "정책이 좋아지고 있다"고 보고했다. (가) 이 loss의 그래디언트는 −∇J-\nabla \textcolor{#d62728}{J}의 올바른 추정인가? (나) loss 값의 감소를 성능 향상의 증거로 볼 수 있는가? (다) y\textcolor{#1c9c60}{y}를 GPT-4가 만든 응답으로 바꾸면 무엇이 깨지는가?

민준이 노트북 화면을 돌린다. 학습 로그의 loss 가 매끄럽게 내려가고 있다.
김민준 (자신만만)
김민준
보세요, loss가 계속 떨어져요. 정책이 좋아지고 있는 거죠.
선생님 (질문)
선생님
이 loss가 1000 스텝 뒤에 0이 됐어요. 무슨 일이 일어났을까요?
김민준 (평상)
김민준
보상을 다 받았다…?
선생님 (평상)
선생님
R이 0에서 1 사이 점수라면, loss가 0이 되는 방법이 하나 더 있지 않나요?
김민준 (평상)
김민준
음… logprob(y)가 0이면요. 확률이 1인 응답만 뽑으면… 아.
이서연 (평상)
이서연
보상이 0인 응답만 뽑아도 loss는 0이야. R=0 이면 곱이 0이니까.
김민준 (놀람)
김민준
둘 다 loss 0이네? 그럼 loss 값으로는 좋아졌는지 모르는 거네요.
선생님 (평상)
선생님
맞아요. 그럼 (가)로 돌아가죠. 이 loss의 그래디언트는 −∇J 의 올바른 추정인가요?
이서연 (평상)
이서연
그건 맞다고 생각해요. y가 π_θ에서 뽑혔으면 −R·∇log π(y)의 기대값이 −∇J니까요. 그러니까… loss 자체도 −J 의 추정 아닌가요? 그래디언트가 같으면 함수도 상수 차이밖에 안 나잖아요.
선생님 (질문)
선생님
서연 학생, "그래디언트가 같다"는 건 어느 변수에 대한 그래디언트죠? 그리고 y는 어디서 왔죠?
이서연 (생각)
이서연
θ에 대해서요. y는 π_θ에서 뽑았는데… 역전파할 때 y는 이미 뽑힌 숫자로 고정돼 있어요. θ가 바뀌면 원래는 어떤 y가 뽑히는지도 바뀌어야 하는데, 이 loss는 그걸 모르고요.
이서연 (깨달음)
이서연
아, 이 loss는 “지금 θ에서, 방금 뽑은 y를 고정했을 때” 만든 임시 함수라서, 그 점에서의 기울기만 −∇J 랑 맞는 거네요. 함수값은 J랑 상관이 없고요.
선생님 (평상)
선생님
정확해요. 이런 걸 대리 손실(surrogate loss)이라고 불러요. 기울기를 빌려 쓰기 위한 발판이지, 점수판이 아니에요. 그래서 실무에서는 loss 대신 평균 보상을 따로 기록해요.
이서연 (평상)
이서연
해석학에서 접선이랑 같네요. 한 점에서 기울기는 곡선이랑 같지만, 접선의 높이가 멀리서도 곡선의 높이를 알려주진 않잖아요.
김민준 (평상)
김민준
이건 조교님이 말한 "채점 기준표 점수랑 실제 실력은 다르다"랑 비슷하다. 기준표 통과했다고 잘 배운 건 아니라고.
선생님 (평상)
선생님
그럼 (다). y를 GPT-4가 만든 응답으로 바꾸면요?
김민준 (평상)
김민준
코드는 그대로 돌아가요. logprob(y)는 어떤 문장이든 계산되니까.
이서연 (평상)
이서연
돌아가긴 하는데, 기대값이 π_θ가 아니라 GPT-4의 분포 아래에서 취해져. 로그 미분 트릭은 "π_θ로 가중한 합"을 "π_θ에서 뽑은 평균"으로 바꾼 거라서, 뽑는 분포가 달라지면 그 등식이 깨져.
선생님 (평상)
선생님
그래요. 그 차이를 보정하려면 샘플마다 두 분포의 확률 비율을 곱해야 해요. 임포턴스 샘플링이라고 부르죠. 오늘 둘 다 한 번씩 틀렸는데, 틀린 이유가 달랐죠.

정리 (가) 그렇다 — 단, y∼πθ\textcolor{#1c9c60}{y} \sim \textcolor{#1565c0}{\pi_\theta}일 때만. (나) 아니다. loss는 한 점에서 기울기만 맞춘 대리 손실이며, 값 자체는 성능을 뜻하지 않는다. 평균 보상을 따로 기록해야 한다. (다) 기대값을 취하는 분포가 바뀌어 그래디언트가 치우친다. 보정하려면 샘플마다 두 분포의 확률 비율(임포턴스 가중치)을 곱해야 한다.