논문: Better & Faster Large Language Models via Multi-token Prediction
한 번에 여러 개의 미래 토큰을 예측하도록 학습시키는 MTP를 정리한 글입니다. 뒤에 나오는 코드는 구조를 이해하기 위한 의사 코드로, 그대로 실행되지는 않습니다.

Abstract

  • 언어모델이 한 번에 여러 개의 미래 토큰을 예측하도록 학습시키는 것이 더 높은 샘플 효율성을 가져옴
  • $n$개의 독립적인 출력 헤드를 사용하여 그 뒤에 이어지는 $n$개의 토큰을 예측하도록 함으로써, 학습 시간의 추가적인 부담 없이 다운스트림 작업 수행 능력이 향상됨을 확인
  • 모델 규모가 커질수록 유용하며, 4개의 토큰 예측 방식으로 학습된 모델은 대규모 배치 크기에서도 추론 속도가 최대 3배 더 빠름

mtp_overview

공유 트렁크(Shared) 하나 위에 4개의 헤드가 올라가 있고, 각 헤드가 $t+1$부터 $t+4$까지를 각각 맡는 구조다. 추론 시에는 헤드를 버리거나, 반대로 속도를 최대 3배까지 끌어올리는 데 쓴다.

Method

Next-token prediction

기존 방식은 이전 토큰 $x_t$가 주어졌을 때 다음 토큰 $x_{t+1}$을 예측할 확률을 최대화한다.

\[L_1 = -\sum_t \log P_\theta(x_{t+1} \mid x_{t:1})\]

Multi-token prediction

MTP는 이전 토큰 $x_t$가 주어졌을 때, 미래 $n$개 토큰이 전부 나올 결합확률을 최대화한다.

\[L_n = -\sum_t \log P_\theta(x_{t+n:t+1} \mid x_{t:1})\]

여기서 $x_{t+n:t+1}$은 미래 $n$개 토큰을 묶은 것이다.

\[\begin{aligned} x_{t+n:t+1} &= (x_{t+1},\ x_{t+2},\ \ldots,\ x_{t+n}) \\ P_\theta(x_{t+n:t+1} \mid x_{t:1}) &= P_\theta(x_{t+1},\ x_{t+2},\ \ldots,\ x_{t+n} \mid x_{t:1}) \end{aligned}\]

문제는 $n$개의 토큰이 동시에 특정 조합으로 나올 결합확률은 일반적으로 구하기 힘들다는 점이다. $n$개 토큰을 한 번에 예측하는 head를 만들려면 출력 공간이 $V^n$이 되어야 한다. 어휘 크기가 32,000이고 $n$이 4라면 출력 공간이 $32000^4$이 되므로 현실적으로 불가능하다.

공유 트렁크와 n개의 헤드

그래서 입력 $x$에 대해서 압축된 $z$를 만들고, $z$를 가지고 여러 토큰을 예측하도록 한다.

  • $z$는 transformer의 output hidden state다
  • $z$만 가지고 독립적인 $n$개의 헤드에 입력하여 토큰을 병렬로 예측한다

두 가지 가정

  1. 원래 문맥 $x$를 다시 보지 않고 $z$가 $x$를 충분히 압축할 수 있다고 가정한다. (Markov 가정. 바로 직전 상태 하나만 알면 충분하다고 단순화)
  2. $z$만 주어지면 $x_{t+1}$, $x_{t+2}$ …는 서로 독립적으로 동시에 예측 가능하다고 가정한다. $n$개의 토큰을 순차적으로 예측할 필요 없이, $z$만 있으면 $n$개의 헤드가 각각 독립적으로 병렬 예측할 수 있다

일반 autoregressive 방식은 매 스텝마다 $x_t$를 매번 처리해야 한다. 하지만 여기서는 $x_t$를 $z$를 구할 때 한 번만 계산하고, 예측할 때는 $z$만 사용한다.

결합확률 분해

$z$가 주어졌을 때 미래 $n$개 토큰이 정답 레이블로 나올 결합확률을 높이는 것이 목표다. 하지만 결합확률을 직접 모델링하기 어렵기 때문에, 미래 각 토큰들은 $z$만 주어지면 서로 독립적으로 예측 가능하다는 조건부 독립을 도입해서 결합확률을 각 토큰별 확률의 곱으로 쪼갠다.

먼저 조건부 확률의 연쇄 법칙을 적용한다.

\[\begin{aligned} P(A,\ B \mid C) &= P(A \mid B,\ C) \cdot P(B \mid C) \\ P_\theta(x_{t+n:t+1},\ z_{t:1} \mid x_{t:1}) &= P_\theta(x_{t+n:t+1} \mid z_{t:1},\ x_{t:1}) \cdot P_\theta(z_{t:1} \mid x_{t:1}) \\ P_\theta(x_{t+n:t+1} \mid x_{t:1}) &\approx P_\theta(x_{t+n:t+1} \mid z_{t:1}) \cdot P_\theta(z_{t:1} \mid x_{t:1}) \end{aligned}\]

마지막 줄에서 $x_{t:1}$이 조건에서 빠지는 것이 1번 가정($z$가 $x$를 충분히 압축)이다. 이제 2번 가정(조건부 독립)을 적용하면 앞쪽 항이 곱으로 쪼개진다.

\[\begin{aligned} P_\theta(x_{t+n:t+1} \mid x_{t:1}) &\approx P_\theta(x_{t+n:t+1} \mid z_{t:1}) \cdot P_\theta(z_{t:1} \mid x_{t:1}) \\ P_\theta(x_{t+n:t+1} \mid z_{t:1}) &= \prod_{i=1}^{n} P_\theta(x_{t+i} \mid z_{t:1}) \end{aligned}\]

이것을 loss에 대입하면 곱이 로그 안에서 합으로 바뀐다.

\[\begin{aligned} L_n &= -\sum_t \log P_\theta(x_{t+n:t+1} \mid z_{t:1}) \cdot P_\theta(z_{t:1} \mid x_{t:1}) \\ &= -\sum_t \sum_{i=1}^{n} \log P_\theta(x_{t+i} \mid z_{t:1}) \cdot P_\theta(z_{t:1} \mid x_{t:1}) \end{aligned}\]

결국 $V^n$짜리 거대한 출력층 하나가 아니라, $V$짜리 헤드 $n$개의 loss 합으로 학습할 수 있게 된다.

헤드 구성

각 헤드의 예측은 다음과 같이 세 단계를 거친다.

\[P_\theta(x_{t+i} \mid x_{t:1}) = \mathrm{softmax}\big(f_u(f_{h_i}(f_s(x_{t:1})))\big)\]
  • $f_s$ = 트랜스포머 모델 (공유 트렁크)
  • $f_{h_i}$ = 트랜스포머 레이어 구조의 헤드
  • $f_u$ = hidden_dim에서 vocab_size로 선형 변환하는 lm_head

Contribution

논문에서는 학습을 통해 $z_t$가 $x_t$, $x_{t+1}$ …에 대한 정보까지 미리 예측해서 저장하도록 학습 압력을 줄 수 있다고 주장한다. 즉 “앞으로 이런 토큰들이 나올 것 같다”는 정보를 $z$에 녹여 넣게 학습이 되고, 그게 더 나은 표현 학습으로 이어진다는 것이다.

여러 미래 토큰을 동시에 맞추도록 훈련시키면 모델이 더 멀리 내다보는 표현을 학습하게 되어, 1개의 토큰만 생성할 때의 성능도 좋아진다는 주장이다.

Limitation

조건부 독립 가정

미래 토큰들이 서로 조건부 독립이라는 가정 자체가 한계다. 각 헤드는 이전 토큰이 무엇을 뽑았는지 전혀 알지 못한다.

예를 들어 원본 문장이 “오늘 날씨는 맑음 입니다.”일 때 다음과 같은 상황이 생길 수 있다.

입력:  [오늘]
예측:  [날씨는], [하루의], [날씨는], [맑음], [입니다]

“오늘”이라는 입력이 주어지면 그 뒤에 토큰을 $n$개 뽑아야 하는데, 각 헤드들은 $z$(“오늘”)만 가지고 예측하게 된다. 헤드끼리 서로 뭘 뽑았는지 알 수 없다.

학습 시와 추론 시

학습 시에는 각 헤드가 자기 위치의 정답과 비교해서 loss를 계산한다. 정답 시퀀스가 “오늘 날씨는 맑음 입니다”이고 [오늘]이 입력으로 들어오면 다음과 같다.

헤드1 -> 날씨는  -> loss 계산
헤드2 -> 맑음    -> loss 계산
헤드3 -> 입니다  -> loss 계산

추론 시에는 헤드들이 낸 예측이 서로 안 맞을 수도 있다. 따라서 후보들을 하나씩 검증하면서 맞는 데까지만 사용하고, 틀리기 시작하는 지점부터는 다시 예측한다.

디코딩 방법

1. Autoregressive

한 개의 헤드만 남기고 나머지 헤드는 버리고, Autoregressive와 같은 방식으로 디코딩한다.

2. Speculative decoding

  1. 현재 context를 $f_s$에 통과시켜 $z_{t:1}$을 얻는다. 이때 $z_t$는 마지막 토큰의 $z$다
  2. 이 $z$를 $n$개의 헤드에 각각 통과시키고 $f_u$로 logit화한 뒤 argmax를 취해 $n$개의 토큰을 생성한다
  3. 원래 context와 생성된 토큰을 하나로 concat해서 $f_s$에 통과시켜 $z_t$를 얻는다. 이때의 $z_t$는 모든 토큰 위치의 $z$다
  4. next-token head($f_{h_1}$)에 통과시키고 $f_u$로 logit화한 뒤, 각 위치마다의 argmax 토큰을 추출해서 앞서 생성한 $n$개의 토큰과 비교한다

코드 구현

모델 구조

공유 트렁크는 딱 한 번만 통과하고, 그 출력 $z$를 $n$개의 헤드가 각자 독립적으로 사용한다.

import torch
import torch.nn as nn

class MultiTokenPredictionModel(nn.Module):
    def __init__(self, hidden_dim, vocab_size, n_future_tokens=4, n_shared_layers=30):
        super().__init__()
        
        # 1. 공유 트렁크 (Shared Trunk) - 기존 트랜스포머와 동일
        #    딱 1개만 존재하며, 모든 헤드가 이 출력을 공유해서 사용
        self.shared_trunk = TransformerModel
        self.lm_head = nn.Linear(hidden_dim, vocab_size=32000, bias=False)
        
        # 2. n개의 독립적인 출력 헤드
        #    각 헤드는 (작은 트랜스포머 블록) + (선형층) 으로 구성
        #    헤드끼리 파라미터를 공유하지 않음 (독립적)
        self.n_future_tokens = n_future_tokens
        self.output_heads = nn.ModuleList([
            nn.ModuleDict({
                "transformer_block": TransformerBlock(hidden_dim),  # 헤드마다 별도 블록
                "unembed": self.lm_head                             # 헤드마다 별도 출력층(가중치는 공유함)
            })
            for _ in range(n_future_tokens)
        ])

    def forward(self, input_ids):
        # --- 무거운 연산: 트렁크는 딱 한 번만 통과 ---
        x = input_ids
        z = self.shared_trunk(x)  # 공유 은닉 표현 (shared hidden representation)
        
        # --- 가벼운 연산: 헤드 n개가 z를 각자 독립적으로 사용 ---
        predictions = []
        for i, head in enumerate(self.output_heads):
            h = head["transformer_block"](z)      # 헤드 i 전용 처리
            logits = head["unembed"](h)            # t+i+1 번째 토큰 예측
            predictions.append(logits)
        
        # predictions[0] -> t+1 예측
        # predictions[1] -> t+2 예측
        # ...
        # predictions[n-1] -> t+n 예측
        return predictions

Loss 계산

헤드 $i$는 $t+i+1$ 위치의 정답을 맞혀야 하므로, 정답 시퀀스를 그만큼 밀어서 정렬한 뒤 cross entropy를 계산한다. 마지막에 $n$개 헤드의 loss를 평균낸다.

def compute_loss(predictions, target_ids, n_future_tokens):
    total_loss = 0
    for i in range(n_future_tokens):
        # 헤드 i는 (t+i+1) 위치의 정답 토큰을 맞혀야 함
        shifted_targets = target_ids[:, i+1:]  # 정답을 i+1칸씩 밀어서 정렬
        logits = predictions[i][:, :shifted_targets.size(1)]
        
        loss_i = nn.functional.cross_entropy(
            logits.reshape(-1, logits.size(-1)),
            shifted_targets.reshape(-1)
        )
        total_loss += loss_i
    
    return total_loss / n_future_tokens  # n개 헤드 loss 평균

compute_loss가 앞에서 유도한 $L_n = -\sum_t \sum_{i=1}^{n} \log P_\theta(x_{t+i} \mid z_{t:1})$에 대응한다. 각 헤드의 cross entropy가 안쪽 합의 한 항씩을 담당한다.

카테고리:

업데이트:

댓글남기기