采样器:top-k / top-p / 温度
采样器只有 101 行(s2_sampler.cpp),却是整个引擎中唯一引入随机性的地方。Slow-AR 的语义 token、Fast-AR 的码本 token 都通过同一个 sample_token() 产生。本章逐段说明其实现,以及几个容易写偏的边界细节。
1. 整体流程
logits (n 个)
│ 全部与 id 配对,按 logit 降序排序
▼
sorted items
│ 在"排序后分布"上先做一次 softmax 得到概率
▼
top-k 截断 + top-p 截断(累积概率)
▼
候选集(空则退化为取最高分)
│ temperature ≤ 0 → 直接 argmax
▼
对候选集按温度重新 softmax → 离散分布采样
│ 概率和异常 → 退化为 argmax
▼
采样到的 token id
2. 为什么先排序、先 softmax
:45-57 做了两件准备:把 logits 变成 (logit, id) 对并降序排序,然后调 softmax_from_sorted_logits() 在全量分布上算一次概率。
这一步是为 top-p 服务的:top-p(nucleus sampling)要求”从最高分开始累加,直到累积概率超过 p”,没有排序和概率就无从谈起。单独的 softmax 实现直接以排序后首元素(即全局最大值)为平移基准(:29-34),数值稳定。
3. top-k 与 top-p 的联合截断
核心循环在 :62-71:
for (i = 0; i < items.size(); ++i) {
cumsum += sorted_probs[i];
bool remove_for_top_k = (i >= k);
bool remove_for_top_p = (i > 0 && cumsum > top_p);
if (remove_for_top_k || remove_for_top_p) continue;
filtered.push_back(items[i]);
}
三个细节值得注意:
- top_k ≤ 0 表示不限制:
k = min(top_k, n)仅在 top_k>0 时生效(:55); - top_p 先钳到 [0,1](:56),避免外部传入非法值;
- top-p 判断带
i > 0且用严格>:最高分 token 永远保留,即使它一个人的概率就超过 p。这防止候选集被清空,也符合”nucleus 至少含一个 token”的标准定义。GenerateParams默认 top_p=0.8、top_k=30。
截断后若候选仍为空(极端参数组合),:73-75 兜底为最高分 token。
4. 温度:两种用法
温度在代码里出现两次,语义不同:
temperature ≤ 0走贪心解码:不再做任何随机采样,直接返回候选集第一个(即全局最高分,:77-79)。这给了用户一个确定性输出的开关;- 正常温度只在候选集上重算 softmax(:81-85):把候选 logits 经 apply_softmax() 按温度缩放。
为什么要”先全量 softmax 做截断、再在候选集上按温度 softmax”?因为截断依据的是未经温度扭曲的真实概率,保证 top-p 语义稳定;最终温度只影响候选内部的相对权重。温度越低分布越尖锐(更稳、更确定),越高越发散。默认 0.8 是稳定性与自然度的折中;RAS 触发时则临时用 1.0(见第 06 章)。
5. 随机数来源
thread_local static std::mt19937 gen(std::random_device{}());
std::discrete_distribution<int32_t> dist(probs.begin(), probs.end());
return filtered[dist(gen)].second;
见 :94-98。两个工程细节:
thread_local:每个线程各自持有一个mt19937,多线程并发调用采样器不会产生数据竞争,也不需要加锁;- discrete_distribution 直接吃概率权重:无需手工做轮盘赌;返回的是候选集内下标,再映射回原始 token id(
.second)。
采样前还有最后一道归一化保险(:87-92):若概率总和 ≤ 0(NaN/Inf 等异常),退化为 argmax 而不是让 distribution 产生未定义行为。
6. 两个调用点的差异
| 调用点 | 传入的 n | 参数来源 |
|---|---|---|
| Slow-AR 语义采样(s2_generate.cpp:70) | 155776(已加 sem_mask) | 用户温度/top_p/top_k,或 RAS 的 1.0/0.9 |
| Fast-AR 码本采样(s2_generate.cpp:141) | 4096(fast_logits.size) | 固定使用用户参数 |
注意 Fast-AR 的 n 直接取 logits 向量长度而非 hparams——即使码本大小变化也能自适应。
7. 小结
采样器实现了教科书式的 top-k/top-p/温度采样,但在工程上补了三层保护:最高分永不被 top-p 剔除、温度 0 退化为贪心、概率异常退化为 argmax。配合生成主循环里的 sem_mask 和 RAS,引擎在”稳定不口吃”和”自然不死板”之间取得平衡。第 08 章进入分量最重的模块——音频 codec:1447 行的卷积编解码器和向量量化是如何用 ggml 搭出来的。
关键文件
| 位置 | 职责 |
|---|---|
| s2_sampler.cpp:42 | sample_token 主流程 |
| s2_sampler.cpp:62 | top-k/top-p 截断 |
| s2_sampler.cpp:8 | 带温度的 softmax |
| s2_sampler.cpp:94 | thread_local 随机源 |
| include/s2_sampler.h | SamplerParams 定义 |