返回关卡地图
预判 0 / 3 闯关题:0 / 3 XP 0 / 300

Level 21 | 零基础导学关卡

解码策略

Decoding Strategies

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

本课官方 Notebook ↗
21_Decoding_Strategies.ipynb 推理优化解码Sampling
Mission 1 概念练习

Temperature:先保护除数,再保持 logits 的形状与索引

🎯 先猜一猜

两个 logits 是 4.0 和 3.1,temperature=0.5。

缩放后的分差是多少?

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

先补的知识

  • logits 是 [batch, vocab_size] 的浮点张量,每一列仍代表固定的 token_id;Temperature 不排序也不删除列。
  • Softmax 前除以 T:T 小于 1 会放大差距,T 大于 1 会缩小差距,T=1 不改变 logits。
  • temperature 是 Python float,作为除数必须设置极小正数下限,避免 0 导致无定义结果。
  • 逐元素除法不改变 logits 的 shape、dtype 或 device。

图解原理

Temperature 像调节分数表的对比度:低温把分差拉大,高温把分差压小。它只缩放数值,不负责把分数变成概率;Softmax 在完整解码流水线最后统一执行。

T < 1

除以小数,分差放大,采样更集中。

T = 1

数值不变,保留原始相对差距。

T > 1

除以大数,分差缩小,采样更分散。

输入logits: [B,V] 浮点张量;temperature: Python float安全下限max(temperature, 1e-6)计算logits / temp,逐元素缩放输出[B,V],token 列顺序、dtype、device 不变可见断言T=0.5 时 index 5 与 6 的分差变为原来的 2 倍

语法热身:用考试分数练安全除法和差值变化

import torch

ratings = torch.tensor([[8.0, 5.0, 2.0]])
temperature = 0.5
safe_temperature = max(temperature, 1e-6)
adjusted = ratings / safe_temperature

print(adjusted.shape)             # [1,3]
print(adjusted[0,0] - adjusted[0,1])  # 6.0

从语法例子迁移到 TODO 1

  • ratings:对应 logits,每一列的位置含义不能改变
  • safe_temperature:对应 temp,使用 Python max 设置正数下限
  • adjusted:对应返回值 logits / temp,不在这里调用 Softmax
  • 差值 6.0:原差值 3.0 除以 0.5 后翻倍,对应可见测试的检查方式

巩固一下

TODO 1 为什么不直接 return F.softmax(logits / temp, dim=-1)?

学完这一段,试着做

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

Mission 2 概念练习

Top-K:用第 K 大值做门槛,过滤但不删除词表列

🎯 先猜一猜

一行 logits 是 [4.0,1.0,3.0,2.0],K=2。

第 K 大门槛和最终保留值分别是什么?

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

先补的知识

  • 函数顶部已处理 top_k<=0 或 top_k>=vocab_size 的边界;这些情况直接返回原 logits。
  • torch.topk(logits, top_k, dim=-1) 会为每个 batch 行分别返回 [B,K] 的最大值和索引。
  • 取 values 的最后一个元素可得到第 K 大门槛;使用 [..., -1:] 保留 [B,1],便于广播。
  • 被过滤位置设为 -inf,而不是删除列;后续 Softmax 会给它们概率 0,原 token_id 索引仍然有效。

图解原理

Top-K 是“每一行单独划线”:先找到该行前 K 名中最低的分数,再把低于这条线的位置盖成 -inf。张量宽度不变,所以后面采样到的列号仍能直接当 token_id。

topk每行找前 K 大值 cutoff取第 K 大,保留末维 comparelogits < cutoff where低分位置替换为 -inf
对象shape作用
logits[B,V]原词表分数与 token 位置
Top-K values[B,K]每行最大的 K 个值
kth_values[B,1]每行广播门槛
过滤结果[B,V]低于门槛处为 -inf
可见测试没有门槛并列,因此恰好保留 3 个位置。一般情况下若多个值等于第 K 大门槛,“小于门槛才过滤”可能保留多于 K 个并列项。

语法热身:每位评委只保留得分最高的两个方案

import torch

ratings = torch.tensor([[2.0, 7.0, 5.0, 1.0],
                        [8.0, 3.0, 6.0, 4.0]])
best_values, _ = torch.topk(ratings, 2, dim=-1)
cutoff = best_values[..., -1:]
blocked = torch.tensor(float('-inf'), device=ratings.device)
screened = torch.where(ratings < cutoff, blocked, ratings)

print(cutoff.shape)   # [2,1]
print(screened.shape) # [2,4]

从语法例子迁移到 TODO 2

  • ratings:对应 logits,每个 batch 行独立筛选
  • best_values:对应 torch.topk 返回的前 K 大值
  • cutoff:对应 kth_values,末尾切片保留长度 1 的维度
  • blocked:对应 filter_value=-inf,并放在 logits.device
  • screened:对应过滤后的 logits,shape 与原索引不变

巩固一下

为什么 kth_values 使用 [..., -1:] 而不是 [..., -1]?

学完这一段,试着做

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

Mission 3 概念练习

Top-p:在排名坐标累计概率,再 scatter 回 token 坐标

🎯 先猜一猜

排序后的累计概率是 [0.54,0.76,0.86,0.94],top_p=0.8。

右移“累计概率大于阈值”的掩码后,应保留几个候选?

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

先补的知识

  • sorted_logits 与 sorted_indices 已由 Notebook 按最后一维降序得到;二者 shape 都是 [B,V]。
  • cumulative_probs 已经对排序后的 logits 做 Softmax 再 cumsum,累计发生在概率而不是原始分数上。
  • 第一个使累计概率超过 top_p 的 token 仍需保留,因此 remove 掩码必须向右平移一格,并把第一项设为 False。
  • 排序位置不是 token_id;过滤后必须根据 sorted_indices 沿 dim=-1 散射回原词表顺序。

图解原理

Top-p 同时使用两个坐标系:先在“从高概率到低概率”的排名坐标里决定候选集合,再回到“原始 token_id”坐标里采样。右移掩码是为了把首次跨过阈值的那一项也收入最小候选集合。

sort分数降序并记录原索引 softmax+cumsum得到累计概率 shift mask保留首次越界项 scatter恢复原 token_id 顺序

阈值 0.8 的排名坐标

0.540.760.86 首次越界0.94 后续

右移前第三项会被标记;右移后第三项保留,从第四项开始删除。

恢复 token 坐标

排序位置 0可能来自原 token 5排序位置 1可能来自原 token 6排序位置 2可能来自原 token 1scatter 后保留值回到列 5、6、1
右移时要读取 clone():左右切片有重叠,直接原地从同一张量复制可能在写入时污染后面还要读取的旧掩码。

语法热身:按累计权重筛商品,再恢复商品原编号

import torch

scores = torch.tensor([[0.2, 2.0, 1.0, 3.0]])
ranked, original_pos = torch.sort(scores, dim=-1, descending=True)
running = torch.cumsum(torch.softmax(ranked, dim=-1), dim=-1)
drop = running > 0.75
drop[..., 1:] = drop[..., :-1].clone()
drop[..., 0] = False
ranked[drop] = float('-inf')
restored = torch.zeros_like(scores).scatter_(-1, original_pos, ranked)

从语法例子迁移到 TODO 3

  • ranked:对应 sorted_logits,已经按分数降序
  • original_pos:对应 sorted_indices,记录每个排序值来自哪个 token 列
  • running:对应 cumulative_probs,shape 仍为 [B,V]
  • drop:对应 sorted_indices_to_remove,需要右移并保护首项
  • restored:对应 restored_logits,沿 dim=-1 恢复原词表顺序

巩固一下

最后 scatter_ 的直接目的是什么?

学完这一段,试着做

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

把理解变成自己的代码

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

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

前往本课官方 Notebook ↗

这节课会遇到的代码对象

apply_temperature · apply_top_k · apply_top_p · decode_next_token · autoregressive_decode · fake_logits

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

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

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