- LLM의 표준 목표인 다음 토큰 예측을 여러 미래 토큰 동시 예측으로 바꾸면, 같은 데이터와 계산 예산에서도 코드·자연어 생성 성능을 더 끌어낼 수 있음
- 구조는 공유 Transformer 본체 위에 여러 출력 헤드를 두는 방식이며, 기본 추론에서는 다음 토큰 헤드만 써 기존 자기회귀 생성처럼 동작함
- 코드 모델에서는 13B 파라미터 모델이 비교 가능한 다음 토큰 모델보다 HumanEval을 12%, MBPP를 17% 더 많이 풀었고, 이득은 큰 모델에서 더 뚜렷함
- 추가 헤드는 자기 추측 디코딩에 활용되어 4-token prediction 모델은 최대 3×, 8-byte prediction 모델은 6.4× 추론 속도 향상을 보임
- 합성 과제에서는 induction heads와 알고리듬 추론에 유리했으며, 학습 시 teacher forcing과 생성 시 자기회귀 분포 차이를 줄이는 효과가 있을 가능성이 있음
멀티 토큰 예측 방식
- 기존 언어 모델링은 각 위치에서 다음 토큰 하나의 교차 엔트로피 손실을 최소화함
- 멀티 토큰 예측은 각 위치에서 다음 n개 토큰을 한꺼번에 예측하도록 학습 목표를 확장함
- 모델 구조는 세 부분으로 나뉨
- 공유 Transformer 본체가 관측된 컨텍스트의 잠재 표현을 만듦
- n개의 독립 출력 헤드가 각 미래 토큰을 병렬로 예측함
- 공유 unembedding matrix가 최종 토큰 확률을 계산함
- 가장 단순한 추론 방식은 다음 토큰 예측 헤드만 사용하는 일반 자기회귀 예측이며, 나머지 헤드는 버릴 수 있음
- 추가 출력 헤드는 blockwise parallel decoding이나 Medusa-like tree attention 같은 자기 추측 디코딩(self-speculative decoding) 에 활용 가능함
메모리 효율 구현
- 단순 구현에서는 각 헤드의 logit과 gradient를 모두 메모리에 올려야 해 GPU 메모리 사용량이 커짐
- 현재 LLM에서는 vocabulary 크기 V가 잠재 표현 차원 d보다 훨씬 커서 logit vector가 GPU 메모리 병목이 됨
- 제안 구현은 공유 본체의 forward pass 뒤에 각 출력 헤드의 forward/backward를 순차 실행함
- 한 헤드의 logit과 gradient는 다음 헤드로 넘어가기 전에 해제됨
- 본체에는 누적 gradient만 유지됨
- 이 방식은 peak GPU 메모리 사용량을 O(nV + d) 에서 O(V + d) 로 줄이며, 런타임 비용은 늘리지 않음
코드 모델 실험 결과
- 실제 데이터 실험은 다음 토큰 예측 모델과 n-token prediction 모델을 같은 파라미터 수로 비교함
- 미래 예측 헤드에 n−1개 레이어를 추가하면 공유 본체에서 n−1개 레이어를 제거함
- 300M부터 13B까지 여섯 크기의 모델을 최소 91B code tokens로 처음부터 학습함
- MBPP와 HumanEval 평가에서는 작은 모델이 기준 모델보다 나쁠 수 있었지만, 규모가 커질수록 멀티 토큰 예측이 앞섬
- 13B 모델은 비교 가능한 다음 토큰 모델보다 더 많은 문제를 해결함
- HumanEval에서 12% 더 많은 문제를 해결함
- MBPP에서 17% 더 많은 문제를 해결함
- 7B 모델을 200B code tokens로 학습한 ablation에서는 n=1, 2, 4, 6, 8을 비교함
- n=4가 HumanEval과 MBPP의 pass@1, pass@10, pass@100에서 일관되게 가장 좋음
- APPS/Intro에서는 n=6이 앞섬
- 최적 window size는 입력 데이터 분포에 따라 달라질 수 있음
추론 속도와 byte-level 모델
- 7B 4-token prediction 모델에 greedy self-speculative decoding을 적용하고, 학습에 쓰지 않은 코드·자연어 테스트 프롬프트에서 디코딩 속도를 측정함
- 결과는 코드에서 3.0×, 텍스트에서 2.7× 속도 향상을 보임
- 코드에서는 3개 제안 중 평균 2.5개 토큰이 수락된 토큰이었음
- 8-byte prediction 모델은 추론 속도에서 6.4× 향상을 기록함
- byte-level tokenization 실험에서는 7B byte-level transformer를 314B bytes, 약 116B tokens에 해당하는 데이터로 학습함
- 8-byte prediction 모델은 next-byte prediction 대비 더 많은 문제를 해결함
- MBPP pass@1에서 67% 더 많은 문제를 해결함
- HumanEval pass@1에서 20% 더 많은 문제를 해결함
- multi-byte prediction은 byte-level 모델을 더 효율적으로 학습시키는 경로가 될 수 있음
여러 epoch, 미세조정, 자연어 결과
- 같은 데이터로 여러 epoch 학습해도 멀티 토큰 예측은 다음 토큰 예측보다 일부 우위를 유지함
- MBPP pass@1은 +2.4%
- HumanEval pass@100은 +3.2%
- 나머지 지표는 유사함
- CodeContests 미세조정에서는 4-token prediction으로 사전학습한 7B 모델이 다음 토큰 기준 모델보다 pass@k 전반에서 우수함
- 4-token prediction 모델을 그대로 n′=4 loss로 미세조정한 경우도 기준 모델보다 좋음
- 추가 헤드를 제거하고 next-token target으로 미세조정한 경우가 전체적으로 가장 좋았음
- 자연어에서는 7B 모델을 200B tokens로 학습해 6개 표준 NLP benchmark를 평가함
- 2-token prediction 모델은 다음 토큰 기준 모델과 비슷함
- 4-token prediction 모델은 성능이 다소 하락함
- 더 큰 모델 크기가 필요할 수 있음
- 생성형 자연어 평가는 요약과 수학 과제로 나눠 수행됨
- 8개 summarization benchmark에서 n=2와 n=4 모델은 200B·500B tokens 학습 모두에서 ROUGE-L F1 기준 다음 토큰 기준 모델보다 높음
- GSM8K 8-shot 평가에서는 200B tokens에서 n=2가 기준 모델을 앞섰지만, 500B tokens 이후에는 패턴이 뒤집혔고 n=4는 전반적으로 더 나쁨
합성 과제에서 본 induction과 알고리듬 추론
- induction은 문장에 “AB”가 나온 뒤 나중에 “A”가 다시 나오면 이어서 “B”를 예측하는 패턴임
- children stories 데이터셋으로 1M~1B nonembedding parameters 모델을 학습하고, 무작위 2-token 이름을 넣은 테스트셋으로 induction capability를 측정함
- 30M 이하 작은 모델에서는 2-token prediction loss가 induction capability 형성을 크게 개선함
- 100M 이상에서는 이 이점이 사라짐
- 다항식 산술 과제에서는 F7[X]/(X5)에서 unary negation, addition, multiplication, composition을 포함한 표현식을 학습·평가함
- 멀티 토큰 예측은 task difficulty 전반에서 정확도를 높였고, out-of-domain generalization도 낮은 절대값이지만 크게 개선함
- 30M에서 100M으로 모델을 키우는 것보다 next-token prediction을 멀티 토큰 예측으로 바꾸는 효과가 더 컸음
왜 작동할 수 있는가
- 멀티 토큰 예측은 teacher forcing 학습과 inference-time autoregressive generation 사이의 분포 불일치를 완화할 수 있음
- 다음 토큰 예측은 짧은 범위의 예측에 집중하면서 긴 범위 의존성을 무시할 수 있음
- 멀티 토큰 예측은 뒤따르는 토큰들과 강하게 관련된 토큰에 더 큰 암묵적 가중치를 부여함
- 이를 choice point 강화로 해석할 수 있음
- 유용한 텍스트 생성은 choice point에서 올바른 결정을 고르는 데 좌우된다고 봄
- 정보이론적 전개에서는 2-token prediction이 X와 Y 사이의 mutual information 항 중요도를 next-token prediction보다 더 키우는 형태로 나타남
한계와 비용
- 남은 과제는 멀티 토큰 예측에서 n을 자동으로 고르는 방법, loss scale과 loss balancing 활용, vocabulary size 조정, embedding space에서 동작하는 보조 prediction loss 개발임
- 모든 실험 모델 학습에는 총 약 500K GPU hours가 사용됨
- 하드웨어는 A100-80GB와 H100임
- 추정 총 배출량은 약 50 tCO2eq이며, Meta의 sustainability program으로 100% offset됨
- 목표는 언어 모델의 compute와 data efficiency를 높이는 것이지만, rebound effects를 주의해야 하며 LLM의 사회적 장점과 위험을 함께 고려해야 함