返回关卡地图
预判 0 / 2 闯关题:0 / 2 XP 0 / 200

Level 23 | 零基础导学关卡

投机解码

Speculative Decoding

先通过图解和小例子理解原理,再做预测与练习;不懂的地方可以反复尝试,最后把理解写进官方 Notebook。

本课官方 Notebook ↗
23_Speculative_Decoding.ipynb 推理优化解码Speculative Decoding
Mission 1 概念练习

先对齐当前位置:p 不小于 q 时,当前草稿必接受

🎯 先猜一猜

位置 0 的草稿 token 是 10,目标概率 p=0.8,草稿概率 q=0.5。

TODO 1 应执行什么?

先猜一个答案,再对照下面的讲解。可以随时修改选择,猜错会得到解释。

先补的知识

  • draft_probs 和 target_probs 都是 [K, vocab_size];第 i 行对应第 i 个草稿位置,列 token_id 对应该 token 的概率。
  • draft_tokens 在可见测试中是 Python 列表 [10,20,30,40];token_id 可直接作为张量列索引。
  • p 和 q 已通过 [i, token_id].item() 取成 Python float,后续比较不是张量广播。
  • 接受概率是 min(1,p/q)。当 p>=q 时结果为 1,因此不需要调用随机数。

图解原理

验证时必须同时对齐“位置 i”和“该位置草稿选中的 token_id”。我们只比较这个候选 token 在两个模型下的概率,不比较整行最大值。若目标模型给它的概率至少和草稿模型一样高,草稿没有过度推荐它,可以直接接受。

位置 i取 draft_tokens[i] 同一列读取 target_probs[i, token_id] 同一列读取 draft_probs[i, token_id] 比较p >= q 时直接 append
draft_probs[K,V] 浮点张量;草稿模型分布target_probs[K,V] 浮点张量;目标模型分布draft_tokens长度 K 的 token_id 序列;可见测试使用 Python listp/q当前 i、当前 token_id 对应的 Python floataccepted_tokens按原顺序追加的 Python list

p=0.8,q=0.5

p>=q,接受概率为 1,直接追加 token。

不要比较错误位置

target_probs[i].max() 回答的是“目标模型最喜欢谁”,不是“草稿 token 能否接受”。

语法热身:逐项核验候选商品的两方评分

import torch

reviewer_a = torch.tensor([[0.2, 0.8], [0.7, 0.3]])
reviewer_b = torch.tensor([[0.1, 0.9], [0.6, 0.4]])
chosen_items = [1, 0]
approved = []

for i, item_id in enumerate(chosen_items):
    p = reviewer_a[i, item_id].item()
    q = reviewer_b[i, item_id].item()
    if p >= q:
        approved.append(item_id)

从语法例子迁移到 TODO 1

  • i:对应草稿位置,选择概率矩阵的行
  • item_id:对应 token_id,选择该行中的词表列
  • reviewer_a:对应 target_probs,取目标概率 p
  • reviewer_b:对应 draft_probs,取草稿概率 q
  • approved.append:对应 accepted_tokens.append(token_id)

巩固一下

为什么读取概率要写 [i, token_id]?

学完这一段,试着做

用一个小动作确认自己理解了;最后再进入官方题目。

Mission 2 概念练习

p 小于 q 时掷一次硬币;一旦拒绝,后续草稿全部停止

🎯 先猜一猜

前两项已接受;第三项 p=0.1、q=0.9、r=0.9。

函数最终应做什么?

先猜一个答案,再对照下面的讲解。可以随时修改选择,猜错会得到解释。

先补的知识

  • 只有 p<q 才进入随机分支,此时接受概率 p/q 位于 0 到 1 之间。
  • torch.rand(1) 返回 shape [1] 的张量;.item() 把单元素张量转为 Python float,便于与 p/q 比较。
  • 接受时继续 for 循环验证下一个草稿 token;拒绝时必须 break,而不是 continue。
  • 后续草稿建立在前面 token 已接受的前缀上;当前 token 被拒绝后,后续条件上下文已经不成立。

图解原理

当小模型比大模型更自信时,需要按 p/q 打折:随机数落在接受区间内就保留,否则拒绝。拒绝不仅影响当前 token,还切断整条草稿链,因为后面的 token 是基于这个未被接受的前缀生成的。

第 1 项:p/q=0.8,r=0.5

0r=0.50.8

r 落在 [0,0.8) 接受区间内,追加 token 并继续。

第 2 项:p/q≈0.11,r=0.9

00.11r=0.9

r 超出接受区间,拒绝并 break;第 3 项不再验证。

1r = torch.rand(1).item()只在 p<q 的分支生成随机数
2if r < p / q成功则 append 当前 token
3else: break失败立即终止整个 for 循环
可见测试把 torch.rand 临时替换成固定返回 0.5、0.9 的函数。因此随机调用次数和顺序也是测试契约:p>=q 的第 0 项不能提前消耗随机数。

语法热身:按通过率逐关检查,失败立即停止

pass_rates = [0.8, 0.2, 0.9]
draws = [0.5, 0.7, 0.1]
passed = []

for level, (rate, draw) in enumerate(zip(pass_rates, draws)):
    if draw < rate:
        passed.append(level)
    else:
        break

print(passed)  # [0]

从语法例子迁移到 TODO 2

  • rate:对应 p/q,本轮草稿 token 的接受概率
  • draw:对应 r,由 torch.rand(1).item() 生成
  • passed.append:对应接受后 append token_id
  • break:对应拒绝当前 token 后停止验证后续草稿
  • passed:对应最终 accepted_tokens,保持前缀顺序

巩固一下

为什么拒绝分支是 break 而不是 continue?

学完这一段,试着做

用一个小动作确认自己理解了;最后再进入官方题目。

把理解变成自己的代码

准备好,去官方题目试一试

在本页用小例子建立直觉,再去官方 Notebook 完成实现。先运行你自己的测试,遇到困难时再查看官方提示与参考答案。

前往本课官方 Notebook ↗

这节课会遇到的代码对象

speculative_decode_step

先找题目里的输入、输出与 TODO,再把本课的手算过程对应进去;以官方题目中的函数说明和测试为准。

练习来源:Datawhale 官方仓库 · 518cc45。这里的入口直接打开官方版本,不读取或分享你的本地 Notebook。

完成 2 道闯关题后,本关即算完成;作业 checklist 用来辅助你回 notebook 练习。
进度只保存在当前浏览器 localStorage,分享 HTML 不会带走你的记录。
上一关:L22 下一关:L24