推理服务 · 解码

投机解码

小模型起草,大模型验证

目标
GPT-2 small 12 层
草稿
distilgpt2 6 层
每轮猜测数
4
这次运行
12 个词元 目标模型 5 趟
输出
与 GPT-2 完全相同

全部步骤

  1. 便宜地猜,一趟检查

    大模型的解码步受限于读取权重,所以一趟检查五个词元的开销和产生一个差不多。投机解码利用了这一点:一个小的草稿模型便宜地猜出接下来几个词元,大模型一次验证所有猜测。

    起草 k 个词元 · 1 趟验证

  2. 一轮

    真实的一轮:distilgpt2(6 层)一个接一个地猜出 4 个词元;GPT-2 small(12 层)对它们只跑一趟,并在每个位置写下自己的选择。猜测保留到第一个不一致之前,再加上 GPT-2 在那里的选择,所以每轮至少得到一个词元。

    保留一致的前缀 + 目标模型的 1 个词元

  3. 一次完整运行

    整次运行:用 5 趟而不是 12 趟 GPT-2 得到 12 个新词元,而且文本与 GPT-2 单独贪心生成的完全一样,因为每个词元都经过了 GPT-2 的检查。

    12 个词元 · 目标模型 5 趟 · 文本相同

  4. 采样时

    如果是采样而不是取最高的词元,草稿词元 x 以概率 min(1, p(x)/q(x)) 被接受,其中 p 是目标模型的概率,q 是草稿的概率。若被拒绝,就从 p 超出 q 的那部分中抽取替代词元。这样输出的词元恰好服从目标模型的分布。

    以概率 min(1, p(x) / q(x)) 接受 x

  5. 快了多少

    如果每个猜测以概率 α 被接受,一趟目标模型平均得到 (1 − α^(k+1)) / (1 − α) 个词元。草稿还必须便宜:distilgpt2 有 GPT-2 一半大,让这次运行更慢;真实系统用的草稿要小 10 到 100 倍。

    每趟词元数 = (1 − α^(k+1)) / (1 − α)

代码

# one round of greedy speculative decoding
draft = []
for _ in range(k): draft.append(small(ctx + draft).argmax())    # k cheap passes
picks = large(ctx + draft).argmax(-1)[-k-1:]                     # one pass checks all k
n = 0
while n < k and draft[n] == picks[n]: n += 1
ctx += draft[:n] + [picks[n]]                                    # at least one new token
# sampling: accept x with probability min(1, p(x) / q(x))
if random() < min(1, p[x] / q[x]): keep(x)
else: x = sample(normalize(max(0, p - q))); stop

延伸阅读