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),它们共享同一套配置与接口约定:
- PyTorch:
GenerationMixin,实现于 src/transformers/generation/utils.py,核心方法为GenerationMixin.generate; - TensorFlow:
TFGenerationMixin.generate,官方文档中声明实现于TFGenerationMixin; - Flax/JAX:
FlaxGenerationMixin.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 decoding) | num_beams=1且do_sample=False |
| 多项式采样(multinomial sampling) | num_beams=1且do_sample=True |
| 束搜索解码(beam-search decoding) | num_beams>1且do_sample=False |
| 束搜索多项式采样(beam-search multinomial sampling) | num_beams>1且do_sample=True |
| 辅助解码(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_beams>1且num_beam_groups>1时进入分组束搜索;传入assistant_model、use_mtp或prompt_lookup_num_tokens时扩展为辅助生成;设置dola_layers时扩展为 DoLa 生成。
2.1 加载配置:from_pretrained
GenerationConfig.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_name="my_configuration.json") >>> generation_config = GenerationConfig.from_pretrained("./test/saved_model/", "my_configuration.json")from_pretrained的完整签名还支持cache_dir(缓存目录)、force_download(强制重新下载)、local_files_only(仅使用本地文件)、token(Hub 鉴权)、revision(分支/标签/commit)等参数。其内部通过cached_file完成“本地目录 → 缓存 → Hub”的逐级查找,然后以 JSON 解析配置字典并通过from_dict实例化。
一个实用的进阶用法是配合return_unused_kwargs=True,在加载时临时修改个别参数,同时收集未被识别的键值,防止拼写错误被静默吞掉:
>>> generation_config, unused_kwargs = GenerationConfig.from_pretrained( ... "openai-community/gpt2", top_k=1, foo=False, do_sample=True, return_unused_kwargs=True ... ) >>> generation_config.top_k 1 >>> unused_kwargs {'foo': False}2.2 从模型配置转换:from_model_config
from_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_pretrained
save_pretrained将配置序列化为generation_config.json写入目标目录,方便复现与分发(configuration_utils.py)。其行为值得注意的细节:
- 默认文件名常量
GENERATION_CONFIG_NAME = "generation_config.json"定义在 src/transformers/utils/init.py; - 保存前会强制执行
validate(strict=True)(见 configuration_utils.py),任何参数组合错误都会抛出异常并拒绝保存,避免坏配置被固化复用; - 序列化时默认使用 diff 模式(
use_diff=True),即只写出与默认配置不同的字段(to_diff_dict,见 configuration_utils.py),配置文件因此最小化且易读; - 支持
push_to_hub=True将配置连同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_length;early_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:缓存实现名,可选"dynamic"(DynamicCache)、"static"(StaticCache)、"offloaded"、"offloaded_static"、"quantized",不指定时使用模型默认缓存(通常为DynamicCache);cache_config:传给 KV 缓存类的参数字典;max_cache_len:仅对静态缓存生效,用于预分配缓存长度,避免多次generate()调用触发重新分配与torch.compile重编译。
logits 加工
temperature:调节下一 token 概率分布的软度,默认 1.0;top_k:top-k 过滤保留的最高概率 token 数量,默认 50;top_p:核采样(nucleus sampling),仅保留累计概率达到top_p的最小 token 集合,默认 1.0;min_p:最小 token 概率,按最可能 token 的概率缩放,典型取值 0.01–0.2;top_h:熵预算缩放因子,控制采样时保留分布熵的比例,取值 0–1,越小输出越聚焦(典型 0.3–0.6);typical_p:局部典型性采样,保留局部典型性累计概率达到typical_p的最小集合;epsilon_cutoff:仅采样条件概率大于该值的 token,论文建议值 3e-4–9e-4;eta_cutoff:eta 采样,结合局部典型采样与 epsilon 采样,建议值 3e-4–2e-3;repetition_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,官方强烈建议设为True;forced_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_scale:classifier-free guidance(CFG)缩放系数,> 1启用 CFG;watermarking_config:水印配置,支持WatermarkingConfig与SynthIDTextWatermarkingConfig,传入dict时会自动转换为前者。
输出变量
num_return_sequences:每个批次元素独立返回的序列数;output_attentions/output_hidden_states/output_scores/output_logits:是否返回注意力张量、隐藏状态、预测分数、未处理的 logits;return_dict_in_generate:是否返回ModelOutput而非仅返回生成序列;要拿到生成缓存或上述output_*输出必须置为True。
特殊 token
pad_token_id/bos_token_id/eos_token_id:填充、序列起始、序列结束 token 的 id,eos_token_id支持列表(多 EOS)。
编解码模型专属
encoder_no_repeat_ngram_size:encoder_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 数,默认 20;num_assistant_tokens_schedule:调度策略,"heuristic"(全部投机 token 正确则 +2,否则 -1,跨调用持久)、"heuristic_transient"(同前但每次调用后重置)、"constant"(保持不变,默认);assistant_confidence_threshold:草稿模型置信度阈值,低于阈值提前停止本轮投机,默认 0.4,跨调用持久;prompt_lookup_num_tokens:以 prompt 检索方式输出候选 token 的数量(无需草稿模型);max_matching_ngram_size:prompt 匹配考虑的最大 n-gram 尺寸,默认 2;assistant_early_exit:支持提前退出的模型可作草稿模型使用;assistant_lookbehind/target_lookbehind:不同 tokenizer 投机解码时的 token 对齐回溯长度,默认 10;assistant_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")、options);disable_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_sample=False却设置了非默认的temperature/top_p/top_k/min_p/typical_p等采样参数,num_beams=1却设置了early_stopping/length_penalty,或return_dict_in_generate=False却开启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.py):
def generate( self, inputs=None, generation_config=None, # 未传时按优先级自动装载 logits_processor=None, # 自定义 LogitsProcessorList stopping_criteria=None, # 自定义 StoppingCriteriaList prefix_allowed_tokens_fn=None, # 束搜索每步允许 token 约束函数 synced_gpus=None, # FSDP/ZeRO-3 多卡时避免死锁 assistant_model=None, # 投机解码草稿模型 streamer=None, # 流式输出 token negative_prompt_ids=None, # CFG 所需负向 prompt negative_prompt_attention_mask=None, custom_generate=None, # 自定义生成(Hub 仓库名 / 本地路径 / Callable) **kwargs, # 临时覆盖 generation_config 参数 )关键约定:
- 配置装配优先级:
generation_config显式传入 > 模型的generation_config.json> 模型配置转换所得;未指明的参数继承GenerationConfig默认值; - 临时覆盖:
generate(inputs, num_beams=4, do_sample=True)这类写法,等价于在调用时用 kwargs 覆盖generation_config的对应字段; - 输入格式:decoder-only 模型传
input_ids;encoder-decoder 模型可传input_ids、input_values、input_features或pixel_values,覆盖文本、语音、视觉多模态场景;为None时以bos_token_id和 batch size 1 初始化; - 返回值:
return_dict_in_generate=True时返回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.py):
def compute_transition_scores( self, sequences, # 生成的序列,形状 (batch_size*num_return_sequences, seq_len) scores, # 每步每个词表 token 的转移分数(log 概率),元组长度 = 生成的 token 数 beam_indices=None, # 束搜索时的束索引,num_beams>1 时必须提供 normalize_logits=False, # 是否在词表维度做 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_tensors="pt") >>> outputs = model.generate(**inputs, max_new_tokens=5, return_dict_in_generate=True, output_scores=True) >>> transition_scores = model.compute_transition_scores( ... outputs.sequences, outputs.scores, normalize_logits=True ... ) >>> # 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_tokens=5, ... num_beams=4, ... num_return_sequences=4, ... return_dict_in_generate=True, ... output_scores=True, ... ) >>> transition_scores = model.compute_transition_scores( ... outputs.sequences, outputs.scores, outputs.beam_indices, normalize_logits=False ... ) >>> # 对生成 token 的分数求和并施加长度惩罚,可重建序列分数 >>> output_length = np.sum(transition_scores.numpy() < 0, axis=1) >>> length_penalty = model.generation_config.length_penalty >>> reconstructed_scores = transition_scores.sum(axis=1) / (output_length**length_penalty) >>> print(np.allclose(outputs.sequences_scores, reconstructed_scores)) True实现上有三个要点(utils.py):一是beam_indices缺省时构造“恒选第一个束”的等价索引,因此贪心搜索无需显式传入;二是把 scores 重塑为[batch*beam, 生成步数]再按束索引取值;三是normalize_logits=True时在词表维度执行log_softmax——注意文档提示,要精确重建束搜索的sequences_scores,应使用normalize_logits=False。这一工具特别适合做生成质量分析、困惑度评估与采样参数调优。
五、配套指南与下一步学习
本文对应的日文主类文档位于 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_generate=True与output_scores=True,这是compute_transition_scores的前置条件;- 束搜索时
num_return_sequences超过num_beams:validate会直接抛异常,二者需满足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),仅供参考