数据截至 (上游 commit 3adf61e154c3)
04 · 采样生成:自回归循环、temperature 与 top-k
这一章讲什么: 训好的模型怎么变成文本。
sample.py处理「加载 + 编解码」的外围杂事,GPT.generate(model.py:305-330)是 25 行的教科书式自回归采样循环。读完你会知道 LLM 推理的最小闭环长什么样,以及它离一个真推理引擎差在哪。
1. 它要解决的小问题
训练产出的是一个「给一段上文、输出下一个 token 的概率分布」的函数。生成文本要把它变成循环:预测 → 采样一个 → 接回上文 → 再预测。这个循环里真正要回答的问题是:
- 概率分布怎么变成「一个具体的 token」?(采样策略:temperature、top-k)
- 上下文超过模型窗口怎么办?(裁左端)
- 权重和 tokenizer 怎么对上?(meta.pkl 或 tiktoken)
2. sample.py 的外围流程
sample.py 做三件事:
① 加载模型(sample.py:35-49)。 两条路:init_from='resume' 从 out_dir/ckpt.pt 读自建模型(先拿 checkpoint['model_args'] 重建结构再灌权重,sample.py:37-46);或以 gpt2* 开头直接 GPT.from_pretrained 拉 OpenAI 权重(sample.py:47-49)。读自建 checkpoint 时同样要剥 _orig_mod. 前缀(sample.py:42-45,与训练侧同款坑)。
② 选编解码器(sample.py:57-74)。 逻辑是「有 meta.pkl 用 meta,否则默认 GPT-2 BPE」:
- 字符级模型:checkpoint 的
config['dataset']指到含 meta.pkl 的数据目录,拿出stoi/itos手工映射(sample.py:61-68)。 - 其他:
tiktoken.get_encoding("gpt2"),编码时允许<|endoftext|>特殊 token(sample.py:71-74)。
③ 准备 prompt 并循环采样(sample.py:77-89)。 --start 直接是文本,或以 FILE: 前缀从文件读(sample.py:77-79);编码成一个 batch 为 1 的张量,然后 num_samples 次调 model.generate、解码、打印。
3. 核心机制:generate 的自回归循环
直觉: 模型一次前向只告诉你「下一个 token 的分布」。要生成 N 个 token,就把「前向 → 取样 → 拼接」重复 N 次。
原理演示(与源码同构的骨架):
# 示意,非源码
def generate(model, idx, max_new_tokens, temperature, top_k):
for _ in range(max_new_tokens):
idx_cond = crop_to_window(idx) # 超窗就裁左端
logits = forward_last_position(model, idx_cond) # 只要最后一个位置的 logits
logits = logits / temperature # 温度调锐度
logits = top_k_filter(logits, top_k) # 只留 top-k,其余置 -inf
probs = softmax(logits)
next_id = multinomial_sample(probs) # 按概率抽一个
idx = cat(idx, next_id)
return idx
真实实现逐段对照(model.py:305-330):
- 裁窗口(
model.py:314):idx_cond = idx if idx.size(1) <= block_size else idx[:, -block_size:]——生成超过上下文长度后,悄悄丢弃最左端。 - 只算最后一个位置:
self(idx_cond)走targets=None分支,前向里 lm_head 只作用于x[:, [-1], :](model.py:189-191),这是推理侧的省时优化(第 1 章 §8)。 - temperature(
model.py:318):logits / temperature。小于 1 分布变尖(更保守),大于 1 变平(更放飞);sample.py默认 0.8(sample.py:17)。 - top-k 截断(
model.py:320-322):torch.topk取出第 k 大的值当门槛,低于门槛的 logits 置-inf,softmax 后概率归零。默认top_k=200(sample.py:18)。 - 采样(
model.py:324-328):softmax 归一后torch.multinomial(probs, 1)按分布抽一个,拼到序列末尾。
整个函数包在 @torch.no_grad() 里(model.py:305),不建梯度图。
4. 坑与边界
- 没有 KV cache,生成是 O(T²) 的。 每产出一个 token 都把整段 上下文完整重算一遍(
model.py:312-316的循环)。真推理引擎会把每层的 K/V 缓存起来只算增量——这里不做,因为教学清晰度优先。长文本生成会明显变慢,这是本库「不是推理框架」的最直接体现。 - 没有 greedy/beam 选项。 唯一的解码策略是 multinomial 采样;想要近似贪心只能把 temperature 调得很小。
- 裁窗口是静默的。 prompt 超长或生成超窗时左端被丢,不告警(
model.py:314)。 - 采样结果不可复现,除非种子一致。
sample.py:19固定seed=1337,但换设备/版本不保证 bitwise 一致 (inferred:代码只设了种子,无任何确定性算子声明)。 - 编解码对错全看 meta 匹不匹配。 用 tiktoken 训的 checkpoint 若误配了 meta.pkl(或反之),生成出来就是乱码,且代码不会检查这种错配 (inferred:两路编解码只靠
load_meta布尔分流,sample.py:57-74)。
5. 本章代码地图
| 主题 | 文件 | 符号 |
|---|---|---|
| 采样入口 | sample.py | 顶层脚本 |
| 加载自建 checkpoint | sample.py:35-46 | init_from == 'resume' 分支 |
| 加载 GPT-2 | sample.py:47-49 | GPT.from_pretrained |
| 编解码选择 | sample.py:57-74 | load_meta、encode/decode |
| FILE: prompt | sample.py:77-79 | start.startswith('FILE:') |
| 生成循环 | model.py:305-330 | GPT.generate |
| 推理省时前向 | model.py:189-191 | GPT.forward 的 targets is None 分支 |
| top-k 截断 | model.py:320-322 | torch.topk + -inf 掩码 |