Back to blog

Better & Faster Large Language Models via Multi-token Prediction

通过多 token 预测训练提升 LLM 性能和推理速度

Better & Faster Large Language Models via Multi-token Prediction

一、论文概述

项目内容
标题Better & Faster Large Language Models via Multi-token Prediction
作者Fabian Gloeckle, Badr Youbi Idrissi, Baptiste Rozière, David Lopez-Paz, Gabriel Synnaeve
机构Meta AI (FAIR)
论文arXiv:2404.19737
发布2024年4月30日
主题cs.CL (计算与语言)

二、核心思想

问题定义

当前大语言模型(如 GPT、Llama)使用 next-token prediction 损失进行训练。然而,next-token prediction 是一种低效的语言、世界知识和推理能力获取方式。Teacher forcing 依赖局部模式,忽略了”困难”决策点,导致 SOTA 模型需要比人类儿童多数个数量级的数据才能达到相同的流利程度。

解决方案概述

本文提出 multi-token prediction:在训练语料的每个位置,让模型同时预测接下来的 nn 个 token,使用 nn 个独立的输出头(output heads)在共享的模型主干(shared trunk)上并行操作。

关键发现:

  • 作为辅助训练任务,multi-token prediction 在代码和自然语言模型上均提升了下游能力,且无训练时间开销
  • 方法在更大模型规模上效果更显著
  • 13B 参数模型在 HumanEval 上多解决 12% 的问题,在 MBPP 上多解决 17%
  • 4-token prediction 模型推理速度提升 3 倍

核心性能

指标数值
MBPP Pass@1 提升+3.8% (7B 模型, n=4)
HumanEval Pass@1 提升+1.2% (7B 模型, n=4)
13B 模型 MBPP 增益+17%
13B 模型 HumanEval 增益+12%
推理加速最高 3× (token 模型), 6.4× (byte 模型)
训练开销无额外开销

三、技术架构

方法概览

Multi-token Prediction 概览

Figure 1: Multi-token prediction 概览。(Top) 训练时,模型通过共享主干和 4 个专用输出头同时预测 4 个未来 token。推理时仅使用 next-token 输出头,其他三个头可用于加速推理(最高 3 倍)。(Bottom) Multi-token prediction 在 MBPP 代码任务上显著提升 pass@1,且随模型规模增大效果更显著。

核心公式

标准 Next-token Prediction 损失: L_1 = -\sum_t \log P_\theta(x_{t+1} \mid x_{t:1}) \tag{1}

Multi-token Prediction 损失: L_n = -\sum_t \log P_\theta(x_{t+n:t+1} \mid x_{t:1}) \tag{2}

分解形式(假设共享主干 + 独立头): Ln=−∑t∑i=1nlog⁡Pθ(xt+i∣zt:1)⋅Pθ(zt:1∣xt:1)L_n = -\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})

实际架构: Pθ(xt+i∣xt:1)=softmax(fu(fhi(fs(xt:1))))P_\theta(x_{t+i} \mid x_{t:1}) = \text{softmax}(f_u(f_{h_i}(f_s(x_{t:1}))))

其中:

  • fsf_s = 共享 transformer 主干(shared trunk)
  • fhif_{h_i} = 第 ii 个独立输出头(transformer 层)
  • fuf_u = 共享 unembedding 矩阵
  • zt:1z_{t:1} = 主干产生的隐藏表示

内存高效实现

前向/反向传播顺序

Figure 2: n-token prediction 模型 (n=2) 的前向/反向传播顺序。通过按顺序对各头执行前向/反向传播,避免同时物化所有 unembedding 层梯度,降低峰值 GPU 内存使用。

关键设计: 在共享主干前向传播后,顺序计算每个独立输出头的前向和反向传播,在主干处累积梯度。

  • 朴素实现: 峰值 GPU 内存 O(nV+d)O(nV + d)
  • 优化实现: 峰值 GPU 内存 O(V+d)O(V + d)
  • 其中 VV = 词汇表大小,dd = 隐藏维度

伪代码:

z = model.shared(x)
d = z.detach()
d.requires_grad = True
for i in range(n):
    p = model.heads[i](d)
    loss(p, y[i]).backward()
z.backward(gradient=d.grad)

推理加速

Multi-token prediction 训练的额外输出头可用于自推测解码(self-speculative decoding),无需额外的 draft 模型:

  1. Blockwise parallel decoding (Stern et al., 2018)
  2. Medusa-like tree attention (Cai et al., 2024)
解码方式代码加速文本加速接受 token 数
4-head 推测解码3.0×2.7×2.5/3 (代码)
8-byte 推测解码6.4×--

训练细节

配置说明
模型规模300M - 13B 参数
训练数据91B - 1T tokens (代码/自然语言)
公平比较相同参数量:添加 n-1 层输出头时,从主干移除 n-1 层
优化器Adam
硬件A100-80GB, H100
总训练~500K GPU 小时, ~50 tCO2eq

四、核心创新

创新点说明理论/实验依据
无开销辅助损失Multi-token prediction 作为辅助训练任务,无训练时间或内存开销内存高效实现:顺序前向/反向传播
规模效应方法在更大模型上效果更显著13B 模型增益远大于 300M 模型
自推测解码训练的额外头直接用于推理加速4-token 模型推理加速 3×
Byte-level 训练多字节预测解锁高效 byte-level 模型训练8-byte prediction MBPP +67%
选择点加权隐式为重要 token 分配更高权重信息论分析:I(X;Y)I(X;Y) 权重翻倍
归纳能力促进小模型的归纳头形成30M 以下模型归纳能力大幅提升

五、实验结果

5.1 规模效应

规模效应

Figure 3: 不同模型规模下 n-token prediction 在 MBPP 上的结果。300M-13B 参数模型在代码上训练。小模型上 multi-token prediction 弱于基线,但在大规模上显著超越。

模型规模MBPP Pass@1 增益HumanEval Pass@1 增益
0.3B--
0.6B+1.8%-
1.3B+4.7%+6.8%
3B+11.1%-
6.7B+23.9%-
13B+26.0%+12%

5.2 最优 n 值搜索

7B 模型, 200B tokens 代码训练:

nMBPP @1MBPP @10MBPP @100HumanEval @1HumanEval @10HumanEval @100
1 (baseline)30.053.873.722.836.462.0
230.355.176.222.238.562.6
433.855.976.924.040.166.1
631.953.973.120.638.463.9
830.752.273.420.036.659.6

关键发现: n=4 在 HumanEval 和 MBPP 上一致性最优。最优窗口大小可能依赖于数据分布。

5.3 Byte-level 训练

7B 模型, 314B bytes (≈116B tokens):

nMBPP @1HumanEval @1APPS/Intro @1
1 (baseline)19.318.10.1
832.3 (+67%)21.8 (+20%)1.2
1628.620.41.0
3223.017.20.6

关键发现: Multi-byte prediction 是解锁高效 byte-level 模型训练的有前途途径。8-byte prediction 模型接近 token-based 模型性能,且训练数据少 1.7×。

5.4 多 epoch 训练

1T tokens (4 epochs):

nMBPP @1HumanEval @1HumanEval @100
140.731.783.0
443.1 (+2.4%)31.686.2 (+3.2%)

Multi-token prediction 在多 epoch 训练中仍保持优势,但增益有所减小。

5.5 Finetuning 效果

CodeContests Finetuning

Figure 4: CodeContests 上的 finetuning 性能比较。4-token prediction 预训练模型的两种 finetuning 方式均超越 next-token 基线。

关键发现:

  • 4-token prediction 预训练 + next-token prediction finetuning 效果最佳
  • 符合”辅助任务预训练 + 任务特定 finetuning”的经典范式
  • CodeContests 是最具挑战性的编码基准

5.6 自然语言评估

摘要任务

Figure 6: 摘要任务性能。7B 模型在 200B 和 500B tokens 自然语言上训练,finetuning 后在 8 个摘要基准上评估 ROUGE-L F1。

训练 tokensn=1n=2n=4
200B26.226.726.7
500B27.127.427.4

关键发现:

  • 标准 NLP 基准(选择题、负对数似然)上,n=4 有轻微退化
  • 生成式评估(摘要)上,n=2 和 n=4 均有提升
  • 选择题基准不适合有效区分生成能力

六、消融实验(合成数据)

6.1 归纳能力

归纳能力

Figure 7: n-token prediction 模型的归纳能力。显示已提及的两 token 名字的第二个 token 的准确率。

实验设置: 在儿童故事数据集上训练 1M-1B 参数模型,通过替换角色名为随机生成的两 token 名字测量归纳能力。

关键发现:

  • 30M 参数以下: 2-token prediction 模型的归纳能力远超 next-token 模型
  • 100M 参数以上: 优势消失
  • 解释: Multi-token prediction 促进跨序列位置的信息传递,有利于归纳头的形成。一旦归纳能力形成,next-token prediction 即可学习该任务。

6.2 算法推理

多项式算术

Figure 8: 多项式算术任务准确率。Multi-token prediction 提升了所有难度级别的准确率,尤其在分布外泛化上效果显著。

实验设置: F7[X]/(X5)\mathbb{F}_7[X]/(X^5) 上的多项式算术(取反、加法、乘法、组合),训练时操作数 1-5,测试时扩展到 1-10。

关键发现:

  • Multi-token prediction 在分布外泛化上带来显著提升
  • 将模型从 30M 增大到 100M 的效果不如将 next-token 替换为 multi-token prediction

七、为什么有效?理论分析

7.1 Lookahead 强化选择点

选择点加权

Figure 9: Multi-token prediction 损失为重要 token 分配更高隐式权重。困难转换 “5→A” 的后果同样难以预测,因此通过其关联获得更高权重。

核心洞察: 并非所有 token 决策都同等重要。选择点(choice points)与高层语义属性相关,决定文本是否有用。

隐式权重分配:

  • 选择点: 权重 n(n+1)2\frac{n(n+1)}{2}
  • 非关键点: 权重 nn

7.2 信息论分析

设 XX = 下一个 token,YY = 下下个 token,CC = 上下文。

Next-token prediction: H(X)=H(X∣Y)+I(X;Y)H(X) = H(X|Y) + I(X;Y)

2-token prediction: H(X)+H(Y)=H(X∣Y)+2I(X;Y)+H(Y∣X)H(X) + H(Y) = H(X|Y) + 2I(X;Y) + H(Y|X)

关键发现: 2-token prediction 将 I(X;Y)I(X;Y)(互信息)的权重提高了 2 倍。Multi-token predictor 在预测与后续文本相关的 token 时更准确。

八、相关工作

相关工作与本文关系
ProphetNet (Qi et al., 2020)最早研究 multi-token prediction,但复制残差流 n 倍
Blockwise parallel decoding (Stern et al., 2018)首次提出推测解码方案,本文用 transformer 层替换线性头
Medusa (Cai et al., 2024)使用 top-k 预测的自推测解码,可与本文模型结合
XLNet (Yang et al., 2019)排列语言建模,但仅预测 15% token
UL2 (Tay et al., 2022)去噪任务混合,但仅 15-25% 掩码
Rho-1 (Lin et al., 2024)选择性 token 训练,与选择点加权相关

九、总结

核心贡献

  1. 简单有效的训练改进: Multi-token prediction 作为辅助损失,无训练开销
  2. 规模效应: 在更大模型上效果更显著,13B 模型代码任务提升 12-17%
  3. 推理加速: 自推测解码实现 3× 加速(token)或 6.4× 加速(byte)
  4. Byte-level 训练: 解锁高效 byte-level 模型训练
  5. 理论分析: 信息论解释为什么 multi-token prediction 有效

技术影响

  • 训练范式: 为 next-token prediction 之外的辅助损失开辟了新方向
  • 推理效率: 自推测解码无需额外 draft 模型,降低部署成本
  • Byte-level 模型: 使 byte-level LLM 训练变得可行
  • 归纳能力: 促进小模型的 in-context learning 能力

局限性

  • 自然语言标准基准(选择题)上 n=4 有轻微退化
  • 最优 n 值依赖于数据分布,尚无自动选择方法
  • Byte-level 模型仍需更多数据才能完全匹配 token-based 模型

十、参考资源

  • 论文: arXiv:2404.19737
  • 引用: arXiv:2404.19737 [cs.CL]
  • 环境影响: ~500K GPU 小时 (A100/H100), ~50 tCO2eq (100% 碳抵消)