跳到正文
LLM 推理
S04模型本身·206

采样

采样器是引擎里最便宜的部分,却是用户感受最强烈的部分。temperature、top-k、top-p、min-p 和各种惩罚,全都是在对 logits 动手术。

  • temperature
  • top-k
  • top-p
  • min-p
  • 重复惩罚
  • 随机种子

为什么重要

问题

模型丢给你 13 万个数字,而你需要一个 token。取最大的那个,也就是贪心解码,对代码和算术站得住脚,拿来写文章则很糟糕:它会循环、会产出最平庸的措辞,还会让模型显得比实际更差。

从完整 softmax 里采样是另一种失败。13 万 token 分布的尾部包含数万个单个概率可忽略、但合起来占几个百分点的 token。采样次数足够多,你一定会撞上一个;而一个荒谬的 token 会毁掉它之后的一切,因为自回归会以自己的错误为条件继续往下走。

所以每个引擎都会带一条过滤流水线。这条流水线只花微秒级时间;而用户说一个模型「感觉」好或不好,说的几乎就是它。

核心思路

解法

在采样前把尾巴截掉。四种标准截断,按应用顺序排列:

  • Temperature —— 在 softmax 之前把 logits 除以 T。小于 1 变尖锐,大于 1 变平坦。T → 0 就是贪心。
  • top-k —— 保留概率最高的 k 个 token。简单,但固定的 k 两头都不对:模型有把握时太宽松,模型不确定时又太苛刻。
  • top-p(nucleus) —— 保留累积概率达到 p 的最小集合。它能自适应分布的形状,也因此成了默认选项。
  • min-p —— 保留概率不低于 m × p_max 的 token。更新一些,在高 temperature 下比 top-p 表现更好:阈值随模型自身的置信度伸缩,而不是盯着一个绝对的概率质量目标。
完整的采样器python
def sample(logits, temperature=1.0, top_k=0, top_p=1.0, min_p=0.0, rng=None):
    if temperature <= 0:
        return int(logits.argmax())        # greedy: every other knob is ignored

    probs = softmax(logits / temperature)
    probs = apply_top_k(probs, top_k)
    probs = apply_top_p(probs, top_p)
    probs = apply_min_p(probs, min_p)
    probs /= probs.sum()                   # renormalise over survivors

    return int(rng.choice(len(probs), p=probs))

顺序是规范的一部分

惩罚作用于 logits,截断作用于概率,归一化放在最后。把 top-p 放到 temperature 之前,你得到的是一个参数名相同、行为不同的采样器。同样的 temperature=0.7, top_p=0.9 在两个服务同一份权重的引擎上表现出明显差异,原因就在这里。

工作原理

工作原理

示意图采样流水线
对 LOGITS 动刀 —— 顺序很重要logits[vocab]惩罚项重复、出现、温度logits / Ttop-k保留最高的 k 个top-p保留累积概率 pmin-p保留 p ≥ m · p_max重新归一化对幸存者做 softmax抽样带种子的 RNG算例 · TOP-P = 0.9the42%a28%cat14%dog8%sat4%qux2%zzz2%累积 = 0.92 ≥ 0.9到此为止;其余全部置零第 4 个 token 会被保留,尽管它把累积推过了 0.9 ——这个阈值是下界,不是上限。为什么 MIN-P 比 TOP-P 表现更好尖锐分布:p_max=0.9 → 保留约 1 个平坦分布:p_max=0.05 → 保留约 40 个min-p 的阈值会随模型自身的置信度一起缩放。
每一级都是词表上的一个掩码。token 一旦被置零,后面的级别再也救不回来:过滤器是按交集组合的。

三种惩罚,它们并不是一回事

  • 重复惩罚(乘性)把已出现 token 的正 logits 除以系数,负 logits 乘以系数。作用于整个上下文时,它会连 the 一起压制,把语法搞坏。
  • 存在惩罚(加性、恒定)对任何出现过至少一次的 token 减去一个常数。
  • 频率惩罚(加性、按次数缩放)减去的量与该 token 出现的次数成正比。

加性的那两个表现更好,因为它们不与 logit 的符号纠缠。如果只实现一个,就实现有界窗口上的频率惩罚:窗口通常取最近几百个 token,而不是整个上下文。

批量采样才是有意思的地方

真实引擎会一次性为 batch 中每个请求采样,而每个请求的采样参数是不同的。这就把一个标量操作变成了带掩码的、按行处理的 GPU kernel。vLLM 用一个 SamplingMetadata 结构保存每请求的 temperature、惩罚和种子,并按「需要哪些过滤器」把请求分组,使无操作路径保持廉价。

这也意味着,朴素实现——给每一行各排一次词表——每个请求每个 token 都要付出一次 13 万元素排序的代价。生产 kernel 用 radix-select 或对概率值做迭代二分来找阈值,从而避开完整排序。

动手观察

动手试试

模拟器采样流水线
存活 token 数
20 / 20
熵(bit)
2.85
采样结果
dog
p(采样结果)
9.9%
分布 —— 灰条 = 过滤前,彩条 = 过滤后
  • the*36.2%
  • a19.9%
  • cat*12.1%
  • dog9.9%
  • sat*6.0%
  • ran4.4%
  • jumped3.0%
  • quietly2.0%
  • onto1.6%
  • under1.3%
  • mat1.0%
  • log0.8%
  • sofa0.6%
  • roof0.4%
  • .0.3%
  • !0.2%
  • epistemic0.1%
  • borogove0.1%
  • qux0.0%
  • zzz0.0%

带 * 的 token 已在输出中出现过,因此重复惩罚会把它们的正 logits 除以惩罚系数。被划掉的行已被过滤器置零,无论换什么随机种子都不可能被采样到。

在当前设置下采样 400 次

这个经验直方图才是用户真正体验到的分布。过滤器不只是让输出「更好」,它让词表中整片区域变得不可达;top-k 设得过紧,模型就会开始循环。

把 temperature 设成 1.6,再对比 top-p 0.9 与 min-p 0.05。高 temperature 下分布被抹平,top-p 的累积目标会把几十个垃圾 token 一并卷进来,min-p 的相对阈值则守住了防线。一张截图就能看完 min-p 的全部论证。

亲手实现

实现

code/s04_sampling.py(节选)python
def apply_top_p(probs, top_p):
    if top_p >= 1.0:
        return probs
    order = np.argsort(-probs)              # descending
    cumulative = np.cumsum(probs[order])
    # keep everything up to AND INCLUDING the token that crosses top_p
    cutoff = int(np.searchsorted(cumulative, top_p)) + 1
    keep = order[:cutoff]
    out = np.zeros_like(probs)
    out[keep] = probs[keep]
    return out


def apply_min_p(probs, min_p):
    if min_p <= 0.0:
        return probs
    threshold = min_p * probs.max()         # relative to the top token
    return np.where(probs >= threshold, probs, 0.0)

那个会悄悄搞坏 top-p 的差一错误

跨过阈值的那个 token 必须保留,而不是丢弃。丢掉它,当最高 token 概率为 0.95 时,top_p=0.9 会保留零个 token,你的归一化随即除以零。每个实现里都有这个 + 1,现在你知道它是干什么的了。
在本地运行
实现完整流水线,然后在十几组参数设置下各采样 5 万次,打印经验熵、可达 token 数,以及一项卡方检验:采样器确实服从它声称实现的那个过滤后分布。
$ python code/s04_sampling.py
预期输出: 一张参数扫描表、带种子采样可复现的断言,以及一份演示:随着 temperature 上升,top-p 与 min-p 急剧分化。

只需要 NumPy — 查看环境准备.

生产实践

在生产环境中

  • vLLM —— 支持逐请求参数的批量 GPU 采样器,另有一条「快路径」:当 batch 内所有请求都是贪心时,完全跳过排序。
  • llama.cpp / picoLM —— 一串可组合的 sampler 结构,由用户按顺序配置,把「顺序」这个问题从隐式变成了显式。
  • 随机种子 —— 可复现性要求每请求独立的 RNG 流,而不是全局种子。用全局种子,你拿到的 token 取决于当时 batch 里恰好还有哪些别的请求,这会让 bug 报告完全无法复现。

练习

  1. 1
    把三种惩罚都实现出来,再构造一个 prompt:重复惩罚设为 1.3 时会因为压制虚词而可测量地破坏语法,这就是要给它加窗口的理由。
  2. 2
    把 top-p 里完整的 argsort 换成对概率阈值的二分搜索。验证输出完全一致,并在 12.8 万元素的向量上测加速比。
  3. 3
    实现每请求独立的带种子 RNG 流,并证明一个请求的输出与 batch 中还有什么无关。

继续学习

接下来

到这里,一个正确但很慢的引擎就完整了:分词、前向、采样、循环。余下每一章都是让它更快,而第一章就是可获得的最大单项收益。S05 引入 KV 缓存,把一个 O(N³) 的引擎变成 O(N²) 的。

自测

习题

先作答,再看解析。答错比答对更有价值,因为解析会指出你该回头重读哪一部分。

自测 1 题 / 共 6

为什么 temperature 升高时,min-p 比 top-p 退化得更优雅?

得分 0/6