跳到主要内容

数据截至 (上游 commit 3adf61e154c3)

04 · 采样生成:自回归循环、temperature 与 top-k

这一章讲什么: 训好的模型怎么变成文本。sample.py 处理「加载 + 编解码」的外围杂事,GPT.generatemodel.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):

  1. 裁窗口model.py:314):idx_cond = idx if idx.size(1) <= block_size else idx[:, -block_size:]——生成超过上下文长度后,悄悄丢弃最左端。
  2. 只算最后一个位置self(idx_cond)targets=None 分支,前向里 lm_head 只作用于 x[:, [-1], :]model.py:189-191),这是推理侧的省时优化(第 1 章 §8)。
  3. temperaturemodel.py:318):logits / temperature。小于 1 分布变尖(更保守),大于 1 变平(更放飞);sample.py 默认 0.8(sample.py:17)。
  4. top-k 截断model.py:320-322):torch.topk 取出第 k 大的值当门槛,低于门槛的 logits 置 -inf,softmax 后概率归零。默认 top_k=200sample.py:18)。
  5. 采样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顶层脚本
加载自建 checkpointsample.py:35-46init_from == 'resume' 分支
加载 GPT-2sample.py:47-49GPT.from_pretrained
编解码选择sample.py:57-74load_metaencode/decode
FILE: promptsample.py:77-79start.startswith('FILE:')
生成循环model.py:305-330GPT.generate
推理省时前向model.py:189-191GPT.forwardtargets is None 分支
top-k 截断model.py:320-322torch.topk + -inf 掩码