ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

Transformers 文本生成实战:GenerationConfig 与 GenerationMixin.generate 完全指南

Transformers 文本生成实战:GenerationConfig 与 GenerationMixin.generate 完全指南 Transformers 文本生成实战GenerationConfig 与 GenerationMixin.generate 完全指南【免费下载链接】transformers Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers本文以 Hugging Face Transformers 仓库中的文本生成主类文档为主线系统讲解GenerationConfig配置体系、GenerationMixin.generate()生成入口与compute_transition_scores()评分工具的完整用法与底层实现。你将掌握如何用一行代码完成贪心、采样、束搜索、辅助解码等多种生成策略如何检查、临时修改、自定义并持久化生成配置以及如何基于源码理解生成参数的真实作用与校验规则从而在文本、视觉、语音等多模态模型上稳定复现可控的生成结果。一、生成Generation机制概览在 Transformers 中“生成”指模型以自回归auto-regressive方式逐 token 产出序列的过程。仓库为不同框架分别实现了独立的生成混入类Mixin它们共享同一套配置与接口约定PyTorchGenerationMixin实现于 src/transformers/generation/utils.py核心方法为GenerationMixin.generateTensorFlowTFGenerationMixin.generate官方文档中声明实现于TFGenerationMixinFlax/JAXFlaxGenerationMixin.generate实现于FlaxGenerationMixin。需要特别说明的是从当前仓库的源码结构看src/transformers/generation/目录下只存在 PyTorch 的实现文件utils.py、configuration_utils.py、logits_process.py、stopping_criteria.py、streamers.py 等且TFGenerationMixin/FlaxGenerationMixin在src/transformers全目录中未检索到定义。可以推断在当前版本快照中生成能力的统一实现以 PyTorch 的GenerationMixin为准TF/Flax 的相关文档条目属于跨框架设计的历史约定。因此本文的源码级讲解将聚焦 PyTorch 实现其参数语义对所有框架保持一致。无论使用哪个框架生成行为都由同一个类控制——GenerationConfig见 src/transformers/generation/configuration_utils.py。它承载了控制生成行为所需的全部参数输出长度、解码策略、logits 加工、缓存策略、输出结构、特殊 token 等。generate()调用时若未显式传入配置会按既定优先级自动装配默认配置。二、GenerationConfig生成行为的“总开关”GenerationConfig是一个继承自PushToHubMixin的配置类负责在生成任务中参数化generate调用。它的文档字符串明确列出了generate支持的五类基础生成方法生成方法触发条件贪心解码greedy decodingnum_beams1且do_sampleFalse多项式采样multinomial samplingnum_beams1且do_sampleTrue束搜索解码beam-search decodingnum_beams1且do_sampleFalse束搜索多项式采样beam-search multinomial samplingnum_beams1且do_sampleTrue辅助解码assisted decoding向.generate()传入assistant_model或prompt_lookup_num_tokens上述模式在源码中由GenerationMode枚举统一描述configuration_utils.py共包含CONTRASTIVE_SEARCH、GREEDY_SEARCH、SAMPLE、ASSISTED_GENERATION、DOLA_GENERATION、BEAM_SEARCH、BEAM_SAMPLE、CONSTRAINED_BEAM_SEARCH、GROUP_BEAM_SEARCH九种模式。GenerationConfig.get_generation_mode()方法会根据当前配置自动判定实际生效的模式configuration_utils.py例如num_beams为空或 1、do_sample非True时进入贪心搜索num_beams1且num_beam_groups1时进入分组束搜索传入assistant_model、use_mtp或prompt_lookup_num_tokens时扩展为辅助生成设置dola_layers时扩展为 DoLa 生成。2.1 加载配置from_pretrainedGenerationConfig.from_pretrained用于从模型仓库或本地目录实例化生成配置configuration_utils.py from transformers import GenerationConfig # 从 Hugging Face 模型仓库下载并缓存配置 generation_config GenerationConfig.from_pretrained(openai-community/gpt2) # 从本地目录加载目录内需包含 generation_config.json generation_config.save_pretrained(./test/saved_model/) generation_config GenerationConfig.from_pretrained(./test/saved_model/) # 支持自定义配置文件名称 generation_config.save_pretrained(./test/saved_model/, config_file_namemy_configuration.json) generation_config GenerationConfig.from_pretrained(./test/saved_model/, my_configuration.json)from_pretrained的完整签名还支持cache_dir缓存目录、force_download强制重新下载、local_files_only仅使用本地文件、tokenHub 鉴权、revision分支/标签/commit等参数。其内部通过cached_file完成“本地目录 → 缓存 → Hub”的逐级查找然后以 JSON 解析配置字典并通过from_dict实例化。一个实用的进阶用法是配合return_unused_kwargsTrue在加载时临时修改个别参数同时收集未被识别的键值防止拼写错误被静默吞掉 generation_config, unused_kwargs GenerationConfig.from_pretrained( ... openai-community/gpt2, top_k1, fooFalse, do_sampleTrue, return_unused_kwargsTrue ... ) generation_config.top_k 1 unused_kwargs {foo: False}2.2 从模型配置转换from_model_configfrom_model_config用于从PreTrainedConfig或配置字典构造GenerationConfig主要服务于旧版模型的兼容迁移configuration_utils.py。它的实现要点包括移除模型配置中的None值让GenerationConfig的默认值生效对多模态/编解码模型依次探测decoder、generator、text_config子配置补全仍处于默认值的生成参数若任一output_attentions/output_hidden_states/output_scores/output_logits被置为True则自动将return_dict_in_generate置为True。这一方法解释了为何老模型没有独立的generation_config.json时依然可以调用generate——生成参数会从模型配置中继承。2.3 保存配置save_pretrainedsave_pretrained将配置序列化为generation_config.json写入目标目录方便复现与分发configuration_utils.py。其行为值得注意的细节默认文件名常量GENERATION_CONFIG_NAME generation_config.json定义在 src/transformers/utils/init.py保存前会强制执行validate(strictTrue)见 configuration_utils.py任何参数组合错误都会抛出异常并拒绝保存避免坏配置被固化复用序列化时默认使用 diff 模式use_diffTrue即只写出与默认配置不同的字段to_diff_dict见 configuration_utils.py配置文件因此最小化且易读支持push_to_hubTrue将配置连同repo_id一起推送。2.4 参数速查完整控制面GenerationConfig的构造逻辑__init__见 configuration_utils.py逐项弹出并校验所有已知参数。按功能分类主要参数如下输出长度控制max_length生成序列总长度上限官方推荐改用max_new_tokens它忽略 prompt 长度语义更清晰max_length仅为向后兼容保留max_new_tokens忽略 prompt 中已有 token 数最多新生成的 token 数min_length/min_new_tokens序列最小长度min_new_tokens设置时优先于min_lengthearly_stopping束搜索类方法的停止条件取值True有num_beams个完整候选即停、False启发式停止、never严格束搜索直到不可能出现更优候选才停max_time生成允许的最大运行秒数秒级超时后仍会完成当前一轮stop_strings一个字符串或字符串列表模型一旦输出这些字符串即终止生成。生成策略do_sample是否采样否则使用贪心解码num_beams束搜索的束数1 表示不做束搜索use_mtp模型支持时是否启用多 token 预测Multi-Token Prediction。缓存控制use_cache是否复用历史 key/value 注意力缓存以加速解码cache_implementation缓存实现名可选dynamicDynamicCache、staticStaticCache、offloaded、offloaded_static、quantized不指定时使用模型默认缓存通常为DynamicCachecache_config传给 KV 缓存类的参数字典max_cache_len仅对静态缓存生效用于预分配缓存长度避免多次generate()调用触发重新分配与torch.compile重编译。logits 加工temperature调节下一 token 概率分布的软度默认 1.0top_ktop-k 过滤保留的最高概率 token 数量默认 50top_p核采样nucleus sampling仅保留累计概率达到top_p的最小 token 集合默认 1.0min_p最小 token 概率按最可能 token 的概率缩放典型取值 0.01–0.2top_h熵预算缩放因子控制采样时保留分布熵的比例取值 0–1越小输出越聚焦典型 0.3–0.6typical_p局部典型性采样保留局部典型性累计概率达到typical_p的最小集合epsilon_cutoff仅采样条件概率大于该值的 token论文建议值 3e-4–9e-4eta_cutoffeta 采样结合局部典型采样与 epsilon 采样建议值 3e-4–2e-3repetition_penalty重复惩罚系数1.0 表示无惩罚encoder_repetition_penalty对不在原始输入中的序列施加的指数惩罚1.0 表示无惩罚length_penalty束搜索的长度指数惩罚作用于序列分数length_penalty 0.0鼓励长序列 0.0鼓励短序列no_repeat_ngram_size大于 0 时同尺寸 n-gram 最多出现一次bad_words_ids禁止生成的 token id 列表的列表renormalize_logits应用全部 logits 处理器后是否重新归一化 logits官方强烈建议设为Trueforced_bos_token_id/forced_eos_token_id强制作为第一个/最后一个生成 token 的 id如 mBART 多语言模型强制首 token 为目标语言 token后者支持列表remove_invalid_values移除模型输出的nan/inf防止生成崩溃注意会拖慢生成exponential_decay_length_penalty(start_index, decay_factor)元组在生成超过start_index个 token 后施加指数增长的长度惩罚suppress_tokens/begin_suppress_tokens生成期/生成初期被抑制logits 置为-inf的 token 列表sequence_bias将 token 序列映射到偏置值的字典正偏置提高选中概率负偏置反之token_healing修复 prompt 尾部 token提升因贪心分词偏差受损的补全质量guidance_scaleclassifier-free guidanceCFG缩放系数 1启用 CFGwatermarking_config水印配置支持WatermarkingConfig与SynthIDTextWatermarkingConfig传入dict时会自动转换为前者。输出变量num_return_sequences每个批次元素独立返回的序列数output_attentions/output_hidden_states/output_scores/output_logits是否返回注意力张量、隐藏状态、预测分数、未处理的 logitsreturn_dict_in_generate是否返回ModelOutput而非仅返回生成序列要拿到生成缓存或上述output_*输出必须置为True。特殊 tokenpad_token_id/bos_token_id/eos_token_id填充、序列起始、序列结束 token 的 ideos_token_id支持列表多 EOS。编解码模型专属encoder_no_repeat_ngram_sizeencoder_input_ids中出现过的 n-gram 禁止在decoder_input_ids中重现decoder_start_token_id解码起始 token id支持传入长度为batch_size的列表以实现同一批次多目标语言。辅助生成assisted/speculative decoding专属is_assistant模型是否为草稿draft模型num_assistant_tokens每轮迭代中草稿模型先生成的投机 token 数默认 20num_assistant_tokens_schedule调度策略heuristic全部投机 token 正确则 2否则 -1跨调用持久、heuristic_transient同前但每次调用后重置、constant保持不变默认assistant_confidence_threshold草稿模型置信度阈值低于阈值提前停止本轮投机默认 0.4跨调用持久prompt_lookup_num_tokens以 prompt 检索方式输出候选 token 的数量无需草稿模型max_matching_ngram_sizeprompt 匹配考虑的最大 n-gram 尺寸默认 2assistant_early_exit支持提前退出的模型可作草稿模型使用assistant_lookbehind/target_lookbehind不同 tokenizer 投机解码时的 token 对齐回溯长度默认 10assistant_ensemble_weight静态集成验证权重取值(0.0, 1.0)用w * p_target (1 - w) * q_draft混合接受概率None保持无损解码speculation_type请求的投机类型如dflash。性能与编译compile_config使用可编译缓存时控制generate如何编译前向传播CompileConfig封装fullgraph、dynamic、backend默认inductor、mode默认reduce-overhead、optionsdisable_compile关闭前向传播的自动编译。2.5 默认参数与校验机制当某个字段保持None时生成循环会用GenerationConfig._get_default_generation_params()的默认值兜底configuration_utils.py{ max_length: 20, min_length: 0, do_sample: False, use_cache: True, early_stopping: False, num_beams: 1, temperature: 1.0, top_k: 50, top_p: 1.0, typical_p: 1.0, repetition_penalty: 1.0, length_penalty: 1.0, no_repeat_ngram_size: 0, encoder_no_repeat_ngram_size: 0, num_return_sequences: 1, output_scores: False, return_dict_in_generate: False, remove_invalid_values: False, epsilon_cutoff: 0.0, eta_cutoff: 0.0, encoder_repetition_penalty: 1.0, num_assistant_tokens: 20, num_assistant_tokens_schedule: constant, assistant_confidence_threshold: 0.4, assistant_lookbehind: 10, target_lookbehind: 10, }validate()方法configuration_utils.py在构造与更新时自动执行负责两类检查硬性错误抛异常如early_stopping不是布尔或never、max_new_tokens 0、cache_implementation非法、num_return_sequences num_beams、同时强制与抑制同一 token 等软性警告仅告警如do_sampleFalse却设置了非默认的temperature/top_p/top_k/min_p/typical_p等采样参数num_beams1却设置了early_stopping/length_penalty或return_dict_in_generateFalse却开启output_*标志。软警告机制引入了user_set_attributes追踪只有用户显式设置的冲突参数才会告警而从模型generation_config.json继承的值不产生噪音。此外validate()还会拦截把logits_processor、stopping_criteria、assistant_model、streamer等本应传给generate()的参数误放进GenerationConfig的常见错误configuration_utils.py。三、GenerationMixin.generate自回归生成的统一入口GenerationMixin是所有具备生成能力模型如LlamaForCausalLM的混入基类src/transformers/generation/utils.py它让模型在初始化时自动装载GenerationConfig并暴露generate系列公共方法。仓库中还提供了custom_generate机制当模型仓库定义了custom_generate/generate.py且开启trust_remote_code时可用自定义生成逻辑完全替代标准流程。3.1 generate 的完整签名generate的核心签名utils.pydef generate( self, inputsNone, generation_configNone, # 未传时按优先级自动装载 logits_processorNone, # 自定义 LogitsProcessorList stopping_criteriaNone, # 自定义 StoppingCriteriaList prefix_allowed_tokens_fnNone, # 束搜索每步允许 token 约束函数 synced_gpusNone, # FSDP/ZeRO-3 多卡时避免死锁 assistant_modelNone, # 投机解码草稿模型 streamerNone, # 流式输出 token negative_prompt_idsNone, # CFG 所需负向 prompt negative_prompt_attention_maskNone, custom_generateNone, # 自定义生成Hub 仓库名 / 本地路径 / Callable **kwargs, # 临时覆盖 generation_config 参数 )关键约定配置装配优先级generation_config显式传入 模型的generation_config.json 模型配置转换所得未指明的参数继承GenerationConfig默认值临时覆盖generate(inputs, num_beams4, do_sampleTrue)这类写法等价于在调用时用 kwargs 覆盖generation_config的对应字段输入格式decoder-only 模型传input_idsencoder-decoder 模型可传input_ids、input_values、input_features或pixel_values覆盖文本、语音、视觉多模态场景为None时以bos_token_id和 batch size 1 初始化返回值return_dict_in_generateTrue时返回GenerateDecoderOnlyOutput/GenerateEncoderDecoderOutput束搜索对应GenerateBeamDecoderOnlyOutput/GenerateBeamEncoderDecoderOutput否则返回torch.LongTensor。3.2 底层调用链与扩展机制从源码结构可以梳理出generate的典型执行脉络先完成配置装配与输入预处理再依据GenerationMode分派到贪心搜索、采样、束搜索、对比搜索、辅助生成等具体解码循环循环中通过logits_process.py中的各类LogitsProcessor如NoBadWordsLogitsProcessor、SequenceBiasLogitsProcessor、SuppressTokens、水印处理器WatermarkLogitsProcessor/SynthIDTextWatermarkLogitsProcessor逐 token 加工 logits通过stopping_criteria.py中的StoppingCriteriaList判定是否终止通过streamers.py中的BaseStreamer实现 token 级流式吐出。仓库测试 tests/generation/test_utils.py 覆盖了compute_transition_scores等核心方法的正确性验证可作为学习各参数组合行为的参考用例。四、compute_transition_scores回溯每个 token 的生成分数compute_transition_scores用于根据生成过程中的scores以及束搜索时的beam_indices快速还原每个被选中 token 的转移分数utils.pydef compute_transition_scores( self, sequences, # 生成的序列形状 (batch_size*num_return_sequences, seq_len) scores, # 每步每个词表 token 的转移分数log 概率元组长度 生成的 token 数 beam_indicesNone, # 束搜索时的束索引num_beams1 时必须提供 normalize_logitsFalse, # 是否在词表维度做 log_softmax 归一化 )官方示例完整演示了贪心与束搜索两种场景的用法。贪心场景不传beam_indices时自动假定恒选第一个束 from transformers import GPT2Tokenizer, AutoModelForCausalLM import numpy as np tokenizer GPT2Tokenizer.from_pretrained(gpt2) model AutoModelForCausalLM.from_pretrained(openai-community/gpt2) tokenizer.pad_token_id tokenizer.eos_token_id inputs tokenizer([Today is], return_tensorspt) outputs model.generate(**inputs, max_new_tokens5, return_dict_in_generateTrue, output_scoresTrue) transition_scores model.compute_transition_scores( ... outputs.sequences, outputs.scores, normalize_logitsTrue ... ) # decoder-only 模型 input_length 为 prompt 长度encoder-decoder 模型为 1 input_length 1 if model.config.is_encoder_decoder else inputs.input_ids.shape[1] generated_tokens outputs.sequences[:, input_length:] for tok, score in zip(generated_tokens[0], transition_scores[0]): ... # | token | token 字符串 | log 概率 | 概率 ... print(f| {tok:5d} | {tokenizer.decode(tok):8s} | {score.numpy():.3f} | {np.exp(score.numpy()):.2%}) | 262 | the | -1.414 | 24.33% | 1110 | day | -2.609 | 7.36% | 618 | when | -2.010 | 13.40% | 356 | we | -1.856 | 15.58% | 460 | can | -2.508 | 8.14%束搜索场景通过beam_indices反查每个 token 实际来自哪个束从而重建整条束路径的分数 outputs model.generate( ... **inputs, ... max_new_tokens5, ... num_beams4, ... num_return_sequences4, ... return_dict_in_generateTrue, ... output_scoresTrue, ... ) transition_scores model.compute_transition_scores( ... outputs.sequences, outputs.scores, outputs.beam_indices, normalize_logitsFalse ... ) # 对生成 token 的分数求和并施加长度惩罚可重建序列分数 output_length np.sum(transition_scores.numpy() 0, axis1) length_penalty model.generation_config.length_penalty reconstructed_scores transition_scores.sum(axis1) / (output_length**length_penalty) print(np.allclose(outputs.sequences_scores, reconstructed_scores)) True实现上有三个要点utils.py一是beam_indices缺省时构造“恒选第一个束”的等价索引因此贪心搜索无需显式传入二是把 scores 重塑为[batch*beam, 生成步数]再按束索引取值三是normalize_logitsTrue时在词表维度执行log_softmax——注意文档提示要精确重建束搜索的sequences_scores应使用normalize_logitsFalse。这一工具特别适合做生成质量分析、困惑度评估与采样参数调优。五、配套指南与下一步学习本文对应的日文主类文档位于 docs/source/ja/main_classes/text_generation.md其完整英文对照版为 docs/source/en/main_classes/text_generation.md。原文档明确指向的生成策略深度指南为 docs/source/en/generation_strategies.md其中涵盖各解码策略的对比、generate的端到端代码示例以及 token 流式输出TextStreamer/TextIteratorStreamer实现在 src/transformers/generation/streamers.py等进阶主题。若需要进一步深入底层建议按以下路径阅读当前仓库源码配置层src/transformers/generation/configuration_utils.py ——GenerationConfig全量参数、默认值与validate校验规则执行层src/transformers/generation/utils.py ——generate入口、compute_transition_scores及各解码循环加工层src/transformers/generation/logits_process.py —— 各类 logits 处理器与水印实现src/transformers/generation/stopping_criteria.py —— 停止准则测试层tests/generation/test_utils.py —— 官方行为验证用例。六、常见问题速查generate返回普通张量而拿不到 scores请设置return_dict_in_generateTrue与output_scoresTrue这是compute_transition_scores的前置条件束搜索时num_return_sequences超过num_beamsvalidate会直接抛异常二者需满足num_return_sequences num_beams设了temperature却仍在贪心解码do_sample必须为True否则相关采样参数只触发软警告并被忽略配置保存失败save_pretrained保存前执行严格校验请先根据报错修正参数组合如移除同时强制与抑制的 token显式覆盖而非依赖默认源码明确提示仍为None的字段会在生成循环中被默认值覆盖想使用非默认值务必在GenerationConfig或generatekwargs 中显式设置。【免费下载链接】transformers Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进