第 15 课:推理与采样 —— 让模型「创造性」地生成

2026-09-16 干徒
RustLLM推理

第 15 课:推理与采样 —— 让模型"创造性"地生成

代码位置:src/sample.rs 模型接口:src/model.rs 演示入口:src/main.rs

1. 本课要搞懂的问题

  1. 模型输出的 logits 是分数,怎么变成真正的文本?
  2. 为什么"每次都选概率最大的 token"(argmax)效果很差?
  3. temperature、top-k、top-p 分别解决了什么问题?它们在代码里怎么实现?
  4. 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.rssample_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_ptruncate
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 序列 → 文本

逐步拆解:

  1. 编码tokenizer.encode(prompt) 把起始文本变成 id 序列。
  2. 截断上下文ctx = ids[ids.len() - block_size ..],上下文超过 block_size 时只保留最近的一段——模型"记不住"更早的历史。
  3. 前向拿 logits:全量模式下每次都把整个 ctx 重新算一遍(慢,但无需额外内存)。
  4. 取最后一个位置:前向输出形状是 [B*T, vocab_size],这里 B=1,最后一个位置即 data[n-v..],它对应"基于当前全部上下文预测的下一个 token"。
  5. 采样并拼接sample_token 选出一个 id,ids.push(next),进入下一轮。
  6. 解码:循环结束后 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 里两个生成演示(falsetrue)生成内容高度一致,正是"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. 动手练习

  1. generatetemperature 分别改成 0.1 和 2.0 跑一遍,观察输出从"复读机"到"胡言乱语"的变化。
  2. top_k 改成 1(等价于只在概率最高 token 附近贪心)再跑,观察重复现象。
  3. top_p 改成 1.0(关闭)但保留 top_k=10,对比输出差异。
  4. sample_token 的轮盘赌循环里打印每次的 u 和命中的 token,手动验证"命中概率 ≈ pᵢ"。
  5. 用足够长的 prompt(超过 block_size,如 40 个字符)分别跑 use_kv_cache = falsetrue,观察 KV 模式提前停止生成的现象。

12. 本课总结

  • 生成 = 前向拿 logits → 采样选 token → 拼回上下文 → 循环
  • argmax 不可取:必然重复、呆板、无法回头
  • 三个旋钮:temperature 调分布锐度、top-k 按个数截断、top-p 按累积概率截断,最后按概率随机采样
  • sample_token 六步:缩放 → 排序 → top-k → softmax → top-p → 轮盘赌
  • generate 全量模式每次重算上下文;KV cache 模式加速但受 block_size 长度限制,超限即停
  • 下一步(第 16 课):把所有零件拼起来,训练一个小 GPT 并生成有意义的长文本
Rust 大语言模型 学习指南共 22 章
1从零用 Rust 实现大语言模型 —— 学习计划2第 1 课:张量 Tensor —— 一切的基础3第 2 课:自动微分 Autograd —— 让模型学会「自我修正」4第 3 课:张量运算扩展 —— 广播、归约、softmax、批量矩阵乘法5第 4 课:模块化重构 —— 项目结构分层与 Rc<RefCell> 架构6第 5 课:线性层与激活函数 —— 神经网络的「积木」7第 6 课:损失函数与优化器 —— 让模型知道「错在哪、怎么改」8第 7 课:第一个 MLP —— 教会神经网络算 XOR9第 8 课:BPE 分词器 —— 让模型「读懂」文字10第 9 课:注意力机制 —— 让 token 互相「看」11第 10 课:多头注意力 —— 让模型「多角度」看世界12第 11 课:位置编码与归一化 —— 让序列带上「位置感」13第 12 课:完整 GPT 模型 —— 把积木拼成能预测下一个词的模型14第 13 课:训练循环 —— 让模型真正开始学习15第 14 课:数据加载 —— 文本如何变成训练样本16第 15 课:推理与采样 —— 让模型「创造性」地生成本篇17第 16 课:训练小 GPT —— 看 loss 从 1.46 一路降到 0.1618第 17 课:AdamW 优化器 —— 给梯度下降装上「惯性」和「自适应步长」19第 18 课:KV Cache —— 让逐 token 生成不再重复计算20第 19 课:RoPE 旋转位置编码 —— 把「相对位置」揉进注意力21第 20 课:学习率调度与收尾 —— warmup、cosine decay 与全项目总结22第 21 课:GPU 加速训练与推理