数据截至 (上游 commit b6c0bfe04c82)
03 · generate():解码总控与三种循环
这一章讲什么:
model.generate(...)从调用到逐 token 吐出的全过程。读完你会知道:temperature/top_p 这些参数到底在哪一步生效、beam search 在 v5 里怎么实现(提示:老版本里的BeamSearchScorer已经没了)、以及投机解码怎么塞同一个循环里。
1. 它要解决的小问题
自回归生成的骨架人人都知道:forward → 取最后一个位置的 logits → 选 token → 拼回去 → 再来。难的是同样的骨架要兼容:贪心/采样/beam search 三种选 token 方式、几十个采样调节参数、流式输出、提前停止、多卡同步、以及「用小模型起草、大模型验证」的投机解码。
generate() 的答案是:总控只做翻译和分派,算法各自是一个自包含循环。
2. 总控:generate() 的四步翻译
GenerationMixin.generate(generation/utils.py:2261)的核心动作(均在 generation/utils.py):
| 步 | 干什么 | 位置 |
|---|---|---|
| ① 定模式 | generation_config.get_generation_mode():看 num_beams/do_sample 决定走哪种解码 | :2474,实现在 generation/configuration_utils.py:534 |
| ② 选循环 | 查 GENERATION_MODES_MAPPING 表拿到本类上的解码方法 | :2474-2482 |
| ③ 组装处理器 | 每个采样参数 → 一个 LogitsProcessor,装进 LogitsProcessorList | _get_logits_processor,:1123 |
| ④ 组装停止条件 | 每个停止参数 → 一个 StoppingCriteria | _get_stopping_criteria,:1358 |
模式分派就是一张表(generation/utils.py:138-147):
| GenerationMode | 调用的方法 |
|---|---|
SAMPLE / GREEDY_SEARCH | _sample |
BEAM_SEARCH / BEAM_SAMPLE | _beam_search |
ASSISTED_GENERATION | _assisted_decoding |
DOLA / CONTRASTIVE_SEARCH / GROUP_BEAM_SEARCH / CONSTRAINED_BEAM_SEARCH | 不在库内,映射到 transformers-community/* Hub 仓库 |
最后一行是 v5 的大变化:四种小众解码策略被移出库,变成按 repo id 动态拉取的社区插件(:144-147);走到这些模式时 generate 会用 custom_generate 递归调用自己(:2489-2504)。核心库只保留三种循环,其它都外挂了。
custom_generate 参数本身也是公开扩展点(:2477-2478):传一个可调用对象或 Hub repo id,就完全替换解码方法。
3. 主循环 _sample:五步一个 token
_sample(generation/utils.py:2783)同时服务贪心与采样(差别只在选 token 那两行)。一轮的骨架:
prefill(整段 prompt 一次 forward,KV 进 cache) (:2868-2873)
│
▼ while 还有未完成序列
① 组输入 prepare_inputs_for_generation:只喂最新 1 个 token (:2877-2881)
② forward outputs = model_forward(**model_inputs) (:2881)
③ 取 logits outputs.logits[:, -1],clone 成 fp32 (:2893)
④ 加工 next_token_scores = logits_processor(input_ids, ·) (:2896)
⑤ 选 token 采样: multinomial(softmax) / 贪心: argmax (:2917-2922)
拼回 input_ids = cat([input_ids, next_tokens]) (:2929)
判停 stopping_criteria 逐条问,全停则退出 (:2933)
三个细节值得记住:
- 第 ③ 步的
.to(copy=True, dtype=torch.float32)(:2893)不是随手写的。 注释(:2891-2892)说明:prefill 那一步的 logits 张量可能巨大(全 prompt × 词表),不 clone 的话outputs的引用会让大 tensor 活到下一轮;clone 只保留最后一个位置的小向量,同时统一升 fp32 保证 processor 数值稳定。循环尾部还有配对的del outputs(:2938)。 - 已完成的序列 用 mask 锁成 pad。
next_tokens = next_tokens * unfinished_sequences + pad_token_id * (1 - unfinished_sequences)(:2925-2926)——batch 里先到 EOS 的序列继续跟着跑(保持 batch 形状),但产出的 token 被强制写成 pad,下游靠 attention mask 忽略。 - 流式输出在循环内同步推。
streamer.put(next_tokens.cpu())(:2930-2931)每轮把新 token 推给用户侧(比如打字机效果),streamer.end()在循环外(:2940-2941)。
模型侧每轮怎么滚动 KV,是 prepare_inputs_for_generation + _update_model_kwargs_for_generation(:940)的事,见第 4 章。
4. 处理器链:参数即对象
_get_logits_processor(:1123)读起来像一份「参数 → 处理器」对照表:每个非默认的 GenerationConfig 字段 append 一个处理器。常用的几个(均在 generation/logits_process.py):
| 参数 | 处理器 | 位置 |
|---|---|---|
temperature | TemperatureLogitsWarper | :238 |
top_p | TopPLogitsWarper | :473 |
top_k | TopKLogitsWarper | :542 |
min_p | MinPLogitsWarper | :704 |
repetition_penalty | RepetitionPenaltyLogitsProcessor | :306 |
no_repeat_ngram_size | NoRepeatNGramLogitsProcessor | :1073 |
每个处理器的接口就一句话:__call__(input_ids, scores) -> scores,返回改过的分数。用户想加约束(比如 JSON schema 约束、禁止某些词),就是自己写一个这种 callable 传给 generate(logits_processor=[...]),它会被拼到链尾。 链本身 LogitsProcessorList(:63)只是个会依次调用的 list。
停止条件同理(generation/stopping_criteria.py):MaxLengthCriteria(:62)、MaxTimeCriteria(:93)、EosTokenCriteria(:543),装进 StoppingCriteriaList(:618);每轮 stopping_criteria(input_ids, scores) 返回一个 bool 向量,unfinished_sequences &= ~... 更新完成状态(:2933)。
5. beam search:v5 的自包含实现
先破一个旧印象: 老版本里 beam search 的灵魂是独立的 BeamSearchScorer 类;在本 commit 里它已经不存在(全库 grep 不到),逻辑被收编进 _beam_search(:3208)及其三个静态助手。新的分工:
| 助手 | 干什么 | 位置 |
|---|---|---|
_get_top_k_continuations | 把分数 reshape 成 (batch, num_beams*vocab),跨所有 beam 全局取 top-K | generation/utils.py:3077 |
_update_finished_beams | 把撞到 EOS/停止条件的 beam 收进完成列表 | :3153 |
_beam_search_has_unfinished_sequences | 判停:活 beam 数够不够、分数还有没有被反超的可能 | :3055 |
每轮的关键一步是把 log_probs reshape 成 (batch_size, num_beams * vocab_size)(:3420)——beam 间的竞争因此变成一次扁平的 top-K,不用逐 beam 循环。选完 beam 后,cache 要按新 beam 归属重排:_gather_beams + cache.reorder_cache(beam_idx)(:3477-3480 一带,注意 RAG/RecurrentGemma 这类特殊模型有自己的 _reorder_cache)。
6. 投机解码:同一循环,候选先行
_assisted_decoding(:3562)实现的是「小模型起草、大模型验证」:
CandidateGenerator先产出一串候选 token——默认用assistant_model(小模型)跑几步(AssistedCandidateGenerator,generation/candidate_generator.py:80);不给小模型时也能用PromptLookupCandidateGenerator(:1018)直接从 prompt 里 n-gram 匹配找候选(适合代码补全这类高重复场景)。- 主模型一次 forward 并行验证整串候选,从前往后收下所有「和主模型分布一致」的前缀,第一个分歧处以主模型为准。
- 收下的 token 一次性拼回,cache 里没被收下的部分回滚。
关键不变量:输出分布与主模型逐 token 生成严格一致(贪心下逐 token 相同),只是省了大模型的调用次数。约束写在函数开头:assisted 必须用 dynamic cache(:3618-3620 一带的检查),因为每轮收下的 token 数不固定。