第 15 课:推理与采样 —— 让模型「创造性」地生成
第 15 课:推理与采样 —— 让模型"创造性"地生成
代码位置:src/sample.rs 模型接口:src/model.rs 演示入口:src/main.rs
1. 本课要搞懂的问题
- 模型输出的 logits 是分数,怎么变成真正的文本?
- 为什么"每次都选概率最大的 token"(argmax)效果很差?
- temperature、top-k、top-p 分别解决了什么问题?它们在代码里怎么实现?
generate生成一句文本时,每步发生了什么?KV cache 模式和全量模式有什么不同?
2. 生成的基本流程
训练完成后(第 13-14 课),模型已经学会了"给前缀,预测下一个 token"。生成文本就是把这件事反复做:
"Once upon a"
│ 前向:model.forward(ctx, 1, T, ...)
▼
logits(每个 token 一个分数,未归一化)
│ 采样:sample_token(logits, T=0.8, top_k=10, top_p=0.9, rng)
▼
下一个 token id
│ 拼回上下文
▼
"Once upon a time" → 再前向 → 再采样 → …… 直到 max_new 个 token
关键在于中间那一步:logits 是分数,不是文本。怎么从分数里选 token,决定了生成质量。
3. 为什么不能直接 argmax
最朴素的想法:把 logits 做 softmax 变成概率,然后永远选概率最大的(这就是贪心解码 / argmax)。它的毛病很典型:
| 现象 | 原因 |
|---|---|
| 重复、呆板 | 一旦某个 token 概率最高,之后每一步都会倾向选它,陷入 "the the the the..." 的死循环 |
| 毫无惊喜 | 概率第二、第三的候选被完全无视,模型"不敢"走任何低概率但合理的路 |
| 容易跑偏 | 一步选错(概率 0.4 的次优解可能是对的),后续全部跟着错,且无法回头 |
大模型的真实需求是"在合理的候选里随机一点":既不能完全确定(呆板),也不能完全乱来(胡言乱语)。于是有了下面三个旋钮。
4. temperature:调节分布的"锐度"
softmax 的公式,引入温度 T 后:
pᵢ = exp(zᵢ / T) / Σⱼ exp(zⱼ / T)
| T 的取值 | 效果 | 直觉 |
|---|---|---|
| T < 1(如 0.8) | 分数差距被放大,分布更"尖" | 更确定、更保守、更连贯 |
| T = 1 | 标准 softmax | 默认 |
| T > 1(如 1.5) | 分数差距被压缩,分布更"平" | 更随机、更有创造力、更易出错 |
代码(src/sample.rs sample_token 第 1 步):
let scaled: Vec<f32> = logits
.iter()
.map(|&l| l / temperature.max(1e-5)) // 防止 T=0 除零
.collect();
数值例子:logits = [1.0, 2.0, 0.5, 0.1](词表 4 个 token):
| 温度 | scaled | softmax 概率 | token 1 概率 |
|---|---|---|---|
| T = 0.5 | [2.0, 4.0, 1.0, 0.2] | [0.11, 0.83, 0.04, 0.02] | 0.83(更"确定") |
| T = 1.0 | [1.0, 2.0, 0.5, 0.1] | [0.21, 0.57, 0.13, 0.09] | 0.57 |
| T = 2.0 | [0.5, 1.0, 0.25, 0.05] | [0.25, 0.41, 0.19, 0.16] | 0.41(更"随机") |
5. top-k:只在前 k 个里选
思想:分数排最后的那些 token 本来就是"凑数"的,干脆把它们从候选里删掉,只在前 k 个里分配概率。
代码(sample_token 第 2-3 步:先按分数从高到低排序,再截断):
let mut items: Vec<(usize, f32)> = scaled.iter().enumerate().map(|(i, &v)| (i, v)).collect();
items.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
if top_k > 0 && items.len() > top_k {
items.truncate(top_k); // top_k = 0 表示不启用
}
6. top-p(nucleus):按累积概率截断
思想:不看"固定个数",而是从高到低累加概率,直到累积概率达到 p,把后面的全部丢掉。候选集合大小随分布自动变化——分布集中时集合小,分布分散时集合大。
代码(sample_token 第 5 步,注意它作用在 softmax 之后的概率上):
if top_p < 1.0 {
let mut cum = 0.0;
let mut keep = items.len();
for (i, p) in probs.iter().enumerate() {
cum += p;
if cum >= top_p { keep = i + 1; break; }
}
items.truncate(keep); // 丢掉尾部
probs.truncate(keep);
let s: f32 = probs.iter().sum();
for p in probs.iter_mut() { *p /= s; } // 重新归一化
}
沿用上面的例子(T=1.0,概率 [0.21, 0.57, 0.13, 0.09],已按分数排序):
top_p = 0.9:
0.57 → 累计 0.57(< 0.9 继续)
0.21 → 累计 0.78(< 0.9 继续)
0.13 → 累计 0.91(≥ 0.9 停)→ 保留前 3 个,丢弃第 4 个
相比 top-k 的"一刀切固定个数",top-p 更聪明:概率分布尖的时候只留 1-2 个候选,平的时候留一堆。实践中 top-k 和 top-p 经常一起开(本项目 main.rs 里就是
top_k=10, top_p=0.9同时启用)。
7. 汇总:sample_token 的完整六步
src/sample.rs 的 sample_token 把上面所有招数串起来:
| 步骤 | 做什么 | 对应代码 |
|---|---|---|
| 1 | 除以 temperature 缩放 | l / temperature.max(1e-5) |
| 2 | 按分数从高到低排序 | items.sort_by(...) |
| 3 | top-k 截断(k>0 时) | items.truncate(top_k) |
| 4 | softmax 转概率(减最大值防溢出) | (*v - max).exp() 再归一化 |
| 5 | top-p 按累积概率截断并重新归一化 | cum >= top_p 时 truncate |
| 6 | 按概率随机采样 | rng.next_f32() 做轮盘赌 |
第 6 步的"轮盘赌"采样:
let mut u = rng.next_f32(); // 均匀随机数 [0, 1)
for (i, p) in probs.iter().enumerate() {
if u < *p { return items[i].0; } // 落进第 i 段的概率 = pᵢ
u -= p; // 没落进就减去这段,继续往后看
}
items.last().map(|(i, _)| *i).unwrap_or(0)
数学上:选到 token i 的概率正好等于 pᵢ。这保证了——概率高的 token 被选中的机会大,但概率低的也有机会被选中,这就是"随机性"的来源。
8. generate:一步一步生成整句文本
全量模式(无 KV cache)的流程,src/sample.rs:
let block_size = model.cfg.block_size;
let mut ids = tokenizer.encode(prompt); // 起始文本 → id 序列
let mut cache = model.new_kv_cache();
for _ in 0..max_new {
if use_kv_cache && cache[0].seq_len() >= block_size { break; } // KV 模式的长度上限
let start = ids.len().saturating_sub(block_size); // 只保留最近 block_size 个
let ctx = &ids[start..];
let logits = if use_kv_cache {
if cache[0].seq_len() == 0 {
model.forward(ctx, 1, ctx.len(), Some(&mut cache)) // 首次:喂整个 prompt
} else {
model.forward(&ids[ids.len() - 1..], 1, 1, Some(&mut cache)) // 之后:只算最新 1 个
}
} else {
model.forward(ctx, 1, ctx.len(), None) // 全量:每次重算整个上下文
};
let v = model.cfg.vocab_size;
let n = logits.numel();
let last_row = &logits.data()[n - v..]; // 取最后一个位置的 logits
let next = sample_token(last_row, temperature, top_k, top_p, rng);
ids.push(next); // 拼回上下文,进入下一轮
}
tokenizer.decode(&ids) // id 序列 → 文本
逐步拆解:
- 编码:
tokenizer.encode(prompt)把起始文本变成 id 序列。 - 截断上下文:
ctx = ids[ids.len() - block_size ..],上下文超过block_size时只保留最近的一段——模型"记不住"更早的历史。 - 前向拿 logits:全量模式下每次都把整个
ctx重新算一遍(慢,但无需额外内存)。 - 取最后一个位置:前向输出形状是
[B*T, vocab_size],这里 B=1,最后一个位置即data[n-v..],它对应"基于当前全部上下文预测的下一个 token"。 - 采样并拼接:
sample_token选出一个 id,ids.push(next),进入下一轮。 - 解码:循环结束后
tokenizer.decode(&ids)把整条序列还原成文本(包含 prompt 和生成的部分)。
9. KV cache 模式的上下文限制:block_size 截断
KV cache(第 18 课详解)的思路:生成第 N 个 token 时,前 N-1 个 token 的 K/V 不需要重算,缓存起来每步只算最新的 1 个 token,大幅加速推理。但它带来一个限制,代码注释里写得很明白:
// KV cache 模式:上下文总长达到 block_size 就停(缓存无法像全量模式那样截断历史)
if use_kv_cache && cache[0].seq_len() >= block_size {
break;
}
两种模式对"超长上下文"的处理对比:
| 模式 | 上下文超过 block_size 时 | 代价 |
|---|---|---|
全量(use_kv_cache = false) |
每次把 ctx 截断到最近 block_size 个 token,可以无限生成 |
每步都要重算整个上下文,慢 |
KV cache(use_kv_cache = true) |
缓存里已经累积了全部历史 K/V,无法丢弃,只能停止生成 | 生成的 token 总数被限制在 block_size 以内 |
所以 KV cache 在
max_new还没用完时可能提前break——这是"加速"换来的"长度上限"。main.rs 里两个生成演示(false和true)生成内容高度一致,正是"cache 只改计算方式、不改生成分布"的验证(提示词短时不会触发截断)。
10. 参数怎么配:一个经验表
| 场景 | temperature | top-k | top-p |
|---|---|---|---|
| 要求准确、连贯(代码、摘要) | 0.2 ~ 0.7 | 小(5~20) | 0.8 ~ 0.9 |
| 通用对话 | 0.7 ~ 0.9 | 30~50 | 0.9 ~ 0.95 |
| 创意写作、头脑风暴 | 0.9 ~ 1.2 | 大或关闭 | 0.95 ~ 1.0 |
main.rs 里的演示配置:temperature=0.8, top_k=10, top_p=0.9,是一个偏保守、够通顺的组合。
11. 动手练习
- 把
generate的temperature分别改成 0.1 和 2.0 跑一遍,观察输出从"复读机"到"胡言乱语"的变化。 - 把
top_k改成 1(等价于只在概率最高 token 附近贪心)再跑,观察重复现象。 - 把
top_p改成 1.0(关闭)但保留top_k=10,对比输出差异。 - 在
sample_token的轮盘赌循环里打印每次的u和命中的 token,手动验证"命中概率 ≈ pᵢ"。 - 用足够长的 prompt(超过
block_size,如 40 个字符)分别跑use_kv_cache = false和true,观察 KV 模式提前停止生成的现象。
12. 本课总结
- 生成 = 前向拿 logits → 采样选 token → 拼回上下文 → 循环
- argmax 不可取:必然重复、呆板、无法回头
- 三个旋钮:temperature 调分布锐度、top-k 按个数截断、top-p 按累积概率截断,最后按概率随机采样
sample_token六步:缩放 → 排序 → top-k → softmax → top-p → 轮盘赌generate全量模式每次重算上下文;KV cache 模式加速但受block_size长度限制,超限即停- 下一步(第 16 课):把所有零件拼起来,训练一个小 GPT 并生成有意义的长文本