조회: 이미 있는 응답의 확률은 한 번에 읽는다
로그확률을 구하려면 토큰마다 확률표가 필요하다. 그런데 DPO는 이 값을 한 번 구하고 끝내지 않는다. 쌍 하나마다 학습 모델과 레퍼런스로, 선호 응답과 비선호 응답에 대해 네 번 구하고, 쌍이 만 개면 학습 한 바퀴에 4만 번이다. 그렇다면 100토큰짜리 응답의 로그확률을 구하는 일은 그 응답을 생성하는 일만큼 비쌀까? 생성과 학습에서 모델이 하는 일을 나란히 놓아 보자.
생성은 뽑고, 학습은 읽는다
생성할 때는 확률표에서 토큰 하나를 실제로 뽑는다. 뽑은 토큰을 입력 끝에 붙여 다시 모델을 돌리고, 새 확률표에서 또 뽑는다. 앞 토큰이 뽑혀야 다음 입력이 생기므로 이 고리는 토큰 수만큼 차례로 돈다.
DPO 학습에서는 응답이 이미 주어져 있다. 데이터셋에
| 흐름 | |
|---|---|
| 생성 | 확률표 → 뽑기 → 토큰 출력 → 다시 확률표 → … (토큰 수만큼 반복) |
| DPO 학습 | 확률표 → 정해진 토큰의 확률 조회 → 로그를 더해 손실 계산 |
조회에는 기다릴 것이 없다. 다음 자리에 올 토큰을 모델이 뽑을 필요 없이 데이터셋에서 가져오면 되기 때문이다. 그래서 프롬프트와 응답 전체를 한꺼번에 입력으로 넣으면, 트랜스포머는 모든 자리의 확률표를 동시에 내놓는다. 각 자리가 자기 앞의 토큰만 보도록 가리는 장치(causal mask) 덕분에, 한꺼번에 계산해도 자리마다의 표는 앞 토큰들만 보고 만든 조건부 확률 그대로다. 이렇게 학습 때 앞 토큰(과거)을 모델이 뽑은 것이 아니라 정답(데이터)에서 가져오는 방식을 티처 포싱(teacher forcing)이라 한다. 미리 모아 둔 데이터셋만 쓰는(오프라인) DPO가 싸게 도는 이유 중 하나다.
역사: 티처 포싱
정답을 입력으로 밀어 넣는 방식은 순환 신경망(자리를 하나씩 지나며 내부 상태를 이어 가는 신경망) 시절에 이름을 얻었다. 1989년 로널드 윌리엄스(Ronald J. Williams)와 데이비드 집서(David Zipser)는 순환망을 멈추지 않고 돌리면서 그때그때 학습시키는 알고리즘을 Neural Computation에 내놓으며, 망이 제 출력 대신 정답(선생님 신호)을 다음 계산에 쓰게 하는 변형에 티처 포싱이라는 이름을 붙였다. 그들은 이 방식이 조던(Michael I. Jordan, 1986)과 피네다(Fernando Pineda, 1988)의 연구에서 이미 자주 쓰이고 있었다고 적었다.
왜 이 변형이 필요했는지는 두 사람이 같은 해 Connection Science에 실은 실험 보고에 드러난다. 그들은 순환망에게 0과 1을 번갈아 내게 하거나, 0, 0, 1, 1을 되풀이하게 하려 했다. 파라미터를 작은 무작위 값으로 시작한 망은 한 값에 가라앉아 멈춰 있었고, 제 출력을 그대로 다시 받는 원래 방식은 이 상태를 벗어나지 못했다. 0, 0, 1, 1 과제는 사실상 끝내 배우지 못했다. 정답을 밀어 넣은 쪽은 단위 하나짜리 과제를 10번이 안 되는 되풀이로, 단위 둘짜리 과제를 약 100번 만에 배웠다. 두 사람은 정답이 매 걸음 망의 상태를 제자리로 되돌려 주기 때문이라고 보았다. 박자가 반 박 어긋난 정답을 보여 주는 실험에서도, 원래 방식은 파라미터를 0과 1 사이의 어중간한 값 0.5 쪽으로 끌고 갔지만 정답을 밀어 넣은 쪽은 한 걸음 만에 박자를 맞췄다. 그들은 대가도 적었다. 정답을 밀어 넣고 찾은 답은, 오차를 0까지 줄이지 못하면 망이 제 출력으로 달릴 때 가장 좋은 답이라는 보장이 없다.
순환망은 정답을 넣어도 자리를 하나씩 차례로 계산해야 했다. 앞 자리의 내부 상태가 있어야 다음 자리를 계산할 수 있기 때문이다. 2017년 트랜스포머 논문은 순환망의 이 차례 계산이 한 예제 안에서 병렬로 계산하는 것을 막는다고 짚었고, 가리개(causal mask)를 쓴 트랜스포머는 모든 자리를 한꺼번에 계산한다.
문제 4 — 연쇄 문제 채점하기
시험지에 10문항이 있고, 문항마다 앞 문항의 답을 넣어야 풀린다(2번은 1번의 답을, 3번은 2번의 답을 쓴다). 학생은 한 문항을 푸는 데 3분이 걸린다. 조교는 학생이 모든 답을 적어 낸 답안지를 받아, 한 문항을 채점하는 데 1분이 걸린다. (가) 학생이 10문항을 다 푸는 데 몇 분이 걸리는가? (나) 조교 10명이 한 장의 답안지를 문항 하나씩 나눠 채점하면 몇 분이 걸리는가? (다) 학생 10명이 한 장의 시험지를 문항 하나씩 나눠 풀면 몇 분이 걸리는가?





정리 (가) 30분. (나) 1분. 앞 문항의 답이 답안지에 이미 적혀 있으므로 문항마다 따로, 동시에 채점할 수 있다. (다) 30분. 앞 답이 나와야 다음을 풀 수 있으므로 사람을 늘려도 차례로 풀어야 한다.
문제 5 — 확률표는 몇 개, 그중 몇 개를 쓰나
프롬프트











정리 (가) 140개. (나) 40~139번째 자리의 100개. 1~39번째 자리의 표는 프롬프트 토큰을, 140번째 자리의 표는 응답 다음 토큰을 맞히는 표라 쓰지 않는다(응답 토큰의 확률은 한 칸 앞 자리의 표에서 읽는다). (다) 100번(앞 토큰의 계산 결과를 저장해 두는 KV 캐시가 있어도 순차).