跳到主要内容

数据截至 (上游 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)

三个细节值得记住:

  1. 第 ③ 步的 .to(copy=True, dtype=torch.float32)(:2893)不是随手写的。 注释(:2891-2892)说明:prefill 那一步的 logits 张量可能巨大(全 prompt × 词表),不 clone 的话 outputs 的引用会让大 tensor 活到下一轮;clone 只保留最后一个位置的小向量,同时统一升 fp32 保证 processor 数值稳定。循环尾部还有配对的 del outputs(:2938)。
  2. 已完成的序列用 mask 锁成 pad。 next_tokens = next_tokens * unfinished_sequences + pad_token_id * (1 - unfinished_sequences)(:2925-2926)——batch 里先到 EOS 的序列继续跟着跑(保持 batch 形状),但产出的 token 被强制写成 pad,下游靠 attention mask 忽略。
  3. 流式输出在循环内同步推。 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):

参数处理器位置
temperatureTemperatureLogitsWarper:238
top_pTopPLogitsWarper:473
top_kTopKLogitsWarper:542
min_pMinPLogitsWarper:704
repetition_penaltyRepetitionPenaltyLogitsProcessor:306
no_repeat_ngram_sizeNoRepeatNGramLogitsProcessor: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-Kgeneration/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)实现的是「小模型起草、大模型验证」:

  1. CandidateGenerator 先产出一串候选 token——默认用 assistant_model(小模型)跑几步(AssistedCandidateGenerator,generation/candidate_generator.py:80);不给小模型时也能用 PromptLookupCandidateGenerator(:1018)直接从 prompt 里 n-gram 匹配找候选(适合代码补全这类高重复场景)。
  2. 主模型一次 forward 并行验证整串候选,从前往后收下所有「和主模型分布一致」的前缀,第一个分歧处以主模型为准。
  3. 收下的 token 一次性拼回,cache 里没被收下的部分回滚。

关键不变量:输出分布与主模型逐 token 生成严格一致(贪心下逐 token 相同),只是省了大模型的调用次数。约束写在函数开头:assisted 必须用 dynamic cache(:3618-3620 一带的检查),因为每轮收下的 token 数不固定。


7. 关键细节与坑

  • do_sample 默认不是 True。 GenerationConfig.__init__self.do_sample = kwargs.pop("do_sample", None)(generation/configuration_utils.py:401),模式判定用 self.do_sample is not True(:551)——不传就是贪心。想采样必须显式 do_sample=True,否则 temperature 等参数不会进处理器链。
  • batch 生成要左 padding。 decoder-only 模型若检测到右 padding,generate 会发警告并提示 padding_side='left'(:2530-2545 一带)——因为最后一个位置必须是有效 token,右 padding 会让 pad 混进上下文。
  • synced_gpus 是分布式防死锁开关。 FSDP/ZeRO-3 下各 rank 完成时间不同,synced_gpus=True 时完成的 rank 继续空转到 max_length(_sample docstring,:2809-2811),否则集合通信会挂。单卡用户无感。
  • return_dict_in_generate=True 的代价是内存。 scores/logits/attentions/hidden_states 每轮都追加进 tuple(:2900-2915),长生成 × 大词表下 scores 一项就很可观——默认关掉不是抠门。
  • deprecated 模式会联网。 走到 CONTRASTIVE_SEARCH 等模式时,custom_generate 会从 Hub 拉 transformers-community/* 仓库的代码执行——离线环境会直接失败,这不是 bug 是设计(:144-147)。

8. 代码地图

主题文件路径符号名
生成总控src/transformers/generation/utils.pyGenerationMixin.generate
模式表src/transformers/generation/utils.pyGENERATION_MODES_MAPPING
模式判定src/transformers/generation/configuration_utils.pyGenerationConfig.get_generation_mode
贪心/采样循环src/transformers/generation/utils.pyGenerationMixin._sample_prefill_update_model_kwargs_for_generation
beam searchsrc/transformers/generation/utils.py_beam_search_get_top_k_continuations_update_finished_beams_beam_search_has_unfinished_sequences
投机解码src/transformers/generation/utils.pycandidate_generator.py_assisted_decodingAssistedCandidateGeneratorPromptLookupCandidateGenerator
logits 处理器src/transformers/generation/logits_process.pyLogitsProcessorListTemperatureLogitsWarperTopPLogitsWarperTopKLogitsWarperMinPLogitsWarperRepetitionPenaltyLogitsProcessorNoRepeatNGramLogitsProcessor
停止条件src/transformers/generation/stopping_criteria.pyStoppingCriteriaListMaxLengthCriteriaMaxTimeCriteriaEosTokenCriteria
流式输出src/transformers/generation/streamers.pyBaseStreamerTextStreamer