레퍼런스 모델 πref (학습을 시작할 때 복사해 얼려 둔 모델)는 DPO가 처음 들여온 것이 아니다. RLHF(사람의 선호로 보상 모델을 학습하고, 그 보상으로 정책 — 응답을 뽑는 모델 — 을 강화학습하는 방식)에서 PPO는 보상에서 토큰마다 βlogπrefπθ 를 뺐다 — 원래 모델에서 멀어질수록 벌점을 주는 KL 벌점이다. 그 목표를 풀어 DPO를 유도하면, 이 KL 항은 사라지지 않고 손실 속의 Δref=logπref(yw)−logπref(yl) — 레퍼런스가 선호 응답을 비선호 응답보다 얼마나 더 좋아하는지 — 로 옮겨 온다. 그래서 DPO는 "레퍼런스 대비 마진을 벌린다"는 모양이 되었다.
여기서 자연스러운 질문이 생긴다.
레퍼런스가 애초에 선호와 비선호를 구분하지 못하면?πref(서울)=πref(부산)=0.5 이면 Δref=0 이다. "0보다 더 벌린다"는 무슨 뜻인가?
더 극단적으로, 레퍼런스가 오답을 선호하면?Δref<0 이다. 거기서 "더 벌린다"는 것은 무엇을 기준으로 하는가?
이 질문의 답이 레퍼런스의 역할을 명확히 해준다.
레퍼런스의 자리: 방향은 데이터가, 크기는 레퍼런스가
선호 라벨은 사람이나 심사 모델이 두 응답을 읽고 매길 뿐, 레퍼런스가 어느 쪽에 확률을 더 주는지는 보지 않는다. 그러니 수만 쌍을 모은 데이터에는 레퍼런스가 비선호 응답을 더 좋아하는 쌍이 섞이기 마련이다. 레퍼런스가 오답을 더 좋아하면, DPO로 학습한 모델도 오답 쪽으로 끌려갈까? 레퍼런스는 손실 식의 어디에 들어가 무엇을 바꾸는가?
역사: 원래 모델에 목줄을 매다
학습 중인 모델을 원래 모델에 KL 벌점으로 묶어 두는 방법은 DPO보다 먼저, 실패를 겪으며 나왔다. 2017년 재키스(Natasha Jaques)와 동료들은 음악 선율을 만드는 순환 신경망을 다듬고 있었다. 악보 데이터로만 학습한 모델은 같은 음을 지나치게 되풀이해서, 만든 선율의 음 가운데 63%가 그런 되풀이 구간에 들어 있었다. 그래서 「같은 음을 너무 오래 끌지 말 것」 같은 음악 이론 규칙을 보상으로 걸고 강화학습으로 고쳤다. 그런데 규칙 보상만으로 학습한 모델은 규칙은 지켰지만 데이터에서 배운 것을 잃었다. 데이터로 학습한 원래 모델이 보기에 이 모델이 고르는 음의 평균 확률은 약 0.026 — 38가지 선택지에서 아무렇게나 고르는 것과 다름없었다. 분자 구조를 글자로 적은 문자열을 만드는 실험에서는 모델이 보상 함수의 빈틈을 찔러, 원자 하나짜리 분자 'N’을 되풀이하거나 탄소를 'CCCCC…'처럼 늘어놓았다. 그들의 해법이 미리 학습한 모델에서 너무 멀어지지 않게 거는 KL 벌점이었다(Sequence Tutor, ICML 2017). 2019년 OpenAI의 지글러(Daniel Ziegler)와 동료들은 「Fine-Tuning Language Models from Human Preferences」에서 같은 벌점을 7억 7,400만 파라미터 GPT-2에 걸어 사람의 선호로 언어모델을 학습시켰다. 벌점 없이 긍정적인 이어 쓰기만 보상했더니 감정 판별 모델의 점수로는 99.97% 긍정인데 문장은 알아볼 수 없는 말이 되었다고 부록에 적었다. 「보상 − KL 벌점」은 이후 RLHF가 이어 쓰는 틀이 되었다. 2023년 DPO는 이 벌점을 따로 계산하는 대신 손실 속 Δref 로 흡수했다.
그래디언트를 두 조각으로 읽기
DPO 손실은 −logσ(βz) 이고, z=Δθ−Δref 는 학습 모델의 마진 Δθ=logπθ(yw)−logπθ(yl) 을 레퍼런스의 마진과 비교한 값이다. 이 손실의 그래디언트는
방향을 정하는 것은 괄호 안이다. yw 의 로그확률을 올리고 yl 의 로그확률을 내린다. 이 항에 레퍼런스 모델은 등장하지 않는다. 어느 쪽이 yw 인지는 데이터셋 — 사람이나 심사 모델(응답을 채점하는 다른 LLM)이 매긴 라벨 — 이 정한다. 레퍼런스가 "부산"을 더 좋아하든 말든, 그래디언트는 언제나 “서울” 쪽으로 민다.
PPO의 KL 벌점은 이 점이 다르다. PPO는 토큰마다 보상에서 βlogπrefπθ 를 뺀다. 모델이 레퍼런스보다 확률을 더 준 토큰일수록 보상이 더 깎이고, 레퍼런스보다 덜 준 토큰은 오히려 보상이 얹힌다. 그래서 그 벌점은 토큰 하나하나의 확률을 레퍼런스 쪽으로 되돌리는 힘이다. DPO의 괄호 안에는 그런 항이 없다.
레퍼런스가 들어가는 곳은 앞의 크기σ(−βz) 뿐이다. 이미 레퍼런스보다 마진을 충분히 벌렸다면(z 가 크면) 크기가 0으로 줄어들어 "이제 그만"이 되고, 아직 못 벌렸다면(z 가 작거나 음수면) 크기가 커서 "더 밀어라"가 된다. 정책 그래디언트를 로그확률 그래디언트 앞의 빈칸 Φ 를 채우는 일로 보면, DPO에서 Φ 자리에 들어가는 것은 βσ(−βz) — 레퍼런스가 정하는 가중치(곱하는 수 — 신경망 파라미터가 아니다)다.
인간 라벨 = 목적지를 가리키는 화살표 (방향) 레퍼런스 = 기둥에 묶인 목줄 (얼마나 멀리 갈 수 있나) β = 목줄을 감아 짧게 만드는 손잡이 — 클수록 목줄이 짧다
Δref 가 0일 때, 음수일 때
레퍼런스가 크기만 바꾼다면, 레퍼런스 마진 Δref 가 0이거나 음수일 때 달라지는 것은 무엇일까? 아직 두 응답을 전혀 구분하지 못하는 모델(Δθ=0)을 세 레퍼런스 아래에 놓고 손실을 재 보자. 레퍼런스가 바꾸는 것은 "Δθ 가 얼마가 되어야 충분한가"라는 기준이다.
레퍼런스
Δref
Δθ=0 인 모델의 손실 (β=0.1)
뜻
둘을 동등하게 봄
0
−logσ(0)=0.693
“레퍼런스도 모르고 나도 모른다. 벌려야 한다.”
오답을 선호
−1
−logσ(0.1)=0.644
동등하게 보는 것만으로 이미 레퍼런스보다 +1 앞서 있다
정답을 강하게 선호
+4
−logσ(−0.4)=0.913
레퍼런스보다 4만큼 뒤처져 있다
어느 경우든 그래디언트의 방향은 같다 — yw 쪽. 달라지는 것은 넘어야 할 기준뿐이다. 0에서 출발하면 조금만 벌려도 진전이고, 이미 +4에서 출발했다면 +4를 넘어야 진전이다.
아래 DPO 손실 위젯에서 Δref 슬라이더를 움직여 보라. 곡선 모양은 그대로 두고 자리만 옆으로 미끄러진다. 위젯의 단추는 아래 문제 1·2의 값을 불러온다.
문제 1 — 레퍼런스가 이미 정답을 알 때
레퍼런스가 정답을 강하게 선호하는 쌍(Δref=+4, β=0.1)에서 학습을 시작했다. 첫 스텝의 손실은 얼마인가? 레퍼런스가 이미 정답을 e4≈55 배 더 잘 만드니, 그래디언트는 0에 가까운가?
김민준
위 표의 셋째 줄이네요. Δref=+4 면 손실 0.913. 레퍼런스가 정답을 55배나 더 좋아하니 할 일이 거의 없고, 그래디언트도 0에 가깝겠죠.
선생님
민준 학생, 그 표의 모델은 Δθ 가 얼마였죠? 학습을 막 시작한 모델은 누구의 복사본이고요?
김민준
표는 Δθ=0 인 모델이었어요. 시작할 때 모델은 레퍼런스 복사본이니까 Δθ=+4, z=0. 손실은 −logσ(0)=0.693 이에요. 레퍼런스가 무엇을 알든 출발점 손실은 같구나.
선생님
그래요. 그럼 그래디언트는요?
김민준
크기가 βσ(0)=0.5β=0.05 라 살아 있어요. 방향은 데이터가 정하니까 yw 쪽이고요.
이서연
손실이 묻는 건 "레퍼런스보다 더 벌렸나"뿐이니까, 레퍼런스가 이미 아는 만큼은 출발선이지 도착점이 아닌 거네.
김민준
과제 채점에서 "기준 점수 대비 향상도 0"이랑 "0점"을 헷갈린 거네요. 작년에 95점 받은 학생도 올해 95점이면 향상도는 0이고요.
정리 출발점에서는 πθ=πref 라 z=0 이고, 손실은 레퍼런스와 상관없이 −logσ(0)≈0.693 이다. 그래디언트 크기 βσ(0)=β/2=0.05 로 살아 있다. 표의 0.913 은 Δθ=0 인 모델의 값이지 출발점의 값이 아니다. "레퍼런스 대비 차이가 0"과 "손실 0"은 다르다.
문제 2 — 레퍼런스가 틀린 쌍은 더 세게 밀리나
학습 도중 모델이 두 쌍에서 똑같이 Δθ=1 이다(β=0.5). 쌍 A는 레퍼런스가 정답을 선호하고(Δref=+3), 쌍 B는 오답을 선호한다(Δref=−3). (가) 두 쌍의 손실과 그래디언트 가중치 σ(−βz) 를 구하시오. 어느 쌍이 몇 배 세게 밀리는가? (나) 두 쌍에서 그래디언트는 각각 yw 와 yl 중 어느 쪽을 올리는가?
김민준
B가 세게 밀리겠죠. 레퍼런스가 틀렸으니 바로잡을 게 많잖아요.
선생님
민준 학생, 가중치 안의 z 는 두 쌍에서 각각 얼마죠?
김민준
z=Δθ−Δref 니까 A는 1−3=−2, B는 1−(−3)=4. 가중치는 A가 σ(1)=0.731, B가 σ(−2)=0.119…
김민준
반대네요. A가 약 6.1배 세게 밀려요. 손실도 A는 1.313, B는 0.127이고요.
이서연
레퍼런스가 오답을 좋아하면 기준이 낮아져. B는 Δθ=1 만으로 레퍼런스보다 4나 앞섰으니까 "이제 그만"에 가까운 거지.
선생님
그래요. (나)는요?
이서연
괄호 안은 ∇logπθ(yw)−∇logπθ(yl) 뿐이고, 레퍼런스는 앞의 σ 안에만 있어요. σ 는 항상 양수라 부호를 못 바꾸니까 두 쌍 다 yw 를 올리고 yl 을 내려요.
정리 (가) A: z=−2, 손실 1.313, 가중치 0.731. B: z=4, 손실 0.127, 가중치 0.119. A가 약 6.1배 세게 밀린다. 레퍼런스가 오답을 좋아할수록 기준이 낮아져 그 쌍은 덜 밀린다. (나) 두 쌍 모두 yw 를 올리고 yl 을 내린다. 레퍼런스는 양수 가중치 σ(−βz) 안에만 있어서 방향을 뒤집을 수 없다.