Transformers 文本生成解码策略完全指南:从默认配置到 Speculative Decoding 的工程实践
【免费下载链接】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
文本生成(Text Generation)是开放域写作、摘要、翻译等自然语言处理任务的核心能力,也是语音转文本、视觉转文本等多模态应用中输出侧的关键环节。在 🤗 Transformers 中,模型选择下一个输出 token 的过程被称为解码(Decoding),而generate()方法采用的解码策略直接决定生成文本的质量——是否重复、是否连贯、是否多样。本文以 docs/source/ko/generation_strategies.md 为主线,系统讲解默认生成配置、常见解码策略与核心参数、如何将自定义解码策略随模型保存与共享,以及流式输出与推测解码的工程用法,并结合 src/transformers/generation/ 下的源码剖析底层实现。读完本文,你将能够针对不同任务选择合适的解码策略,并掌握用GenerationConfig复现与共享生成效果的方法。
文本生成与generate()方法概述
在 Transformers 中,支持文本生成的模型(如 GPT2、XLNet、OpenAI GPT、CTRL、TransformerXL、XLM、Bart、T5、GIT、Whisper 等)都通过 [~generation.GenerationMixin.generate] 方法完成推理,该方法的实现位于 src/transformers/generation/utils.py 的GenerationMixin类中。generate()可以驱动多种任务:
- 文本摘要
- 图像描述
- 音频转录
generate()的输入由模型的前处理类(如AutoTokenizer或AutoProcessor)返回,输入形态取决于模型的数据类型。当模型的前处理组件产生多个输入类型时,需要将全部输入传递给generate(),具体细节可查阅各模型文档。
选择输出 token 的过程称为解码,generate()允许用户自定义解码策略。修改解码策略不会改变模型的可训练参数,却能显著影响输出质量——例如减少文本重复、让文本更连贯。这正是本文要深入探讨的主题。
默认文本生成配置
模型的解码策略定义在**生成配置(generation config)**中。当预训练模型在 [pipeline] 中用于推理时,模型内部会调用PreTrainedModel.generate()并应用默认生成配置;即使用户没有为模型保存自定义配置,默认配置也会生效。
显式加载模型时,可以通过model.generation_config检查模型的生成配置:
>>> from transformers import AutoModelForCausalLM >>> model = AutoModelForCausalLM.from_pretrained("distilbert/distilgpt2") >>> model.generation_config GenerationConfig { "bos_token_id": 50256, "eos_token_id": 50256, }打印model.generation_config时只会显示与默认值不同的字段,其余字段沿用库内默认值,因此输出通常很简短。GenerationConfig类定义在 src/transformers/generation/configuration_utils.py,它继承了PushToHubMixin,支持直接推送配置到 Hub。
关于默认行为有两点关键事实:
- 默认最大生成长度:默认生成配置将「提示词 + 输出」的总长度限制为20 个 token,以避免资源耗尽(在源码
_get_default_generation_params()中体现为max_length=20)。 - 默认解码策略:默认为贪心搜索(Greedy Search),即每一步选择概率最高的 token,是最简单的解码方式。
贪心搜索对输出较短、不追求创造性的任务表现良好;但在生成较长输出时会快速陷入重复。此外,若模型本身在config.json中声明了更长的默认生成长度(例如 Llama 2 系列为 4096),实际默认值以模型配置为准,这一点在 docs/source/en/generation_strategies.md 的示例中也有体现。
自定义文本生成:核心参数
用户可以直接向 [generate] 方法传入参数和值来覆盖generation_config:
>>> my_model.generate(**inputs, num_beams=4, do_sample=True) # doctest: +SKIP即使默认解码策略对多数任务够用,仍有几个高频调优参数值得掌握:
max_new_tokens:新生成 token 的最大数量,即不包含提示词的输出序列长度。除按长度停止外,也可以选择按总耗时停止生成(例如设置max_time),此时会借助 [StoppingCriteria](见 src/transformers/generation/stopping_criteria.py)在生成超过指定时间后中止。num_beams:设为大于 1 时,从贪心搜索切换为束搜索(Beam Search)。该策略在每个时间步同时评估多个假设(束),最终选出整体概率最高的完整序列。其优势在于能够发现那些因初始 token 概率较低而被贪心搜索忽略的高分序列。do_sample:设为True时启用基于概率分布的采样类策略,包括多项式采样(Multinomial Sampling)、束搜索多项式采样(Beam-Search Multinomial Sampling)、Top-K 采样和 Top-p 采样等。这些策略从整个词表上的概率分布中抽取下一个 token,并施加各自的调整。num_return_sequences:每个输入返回的候选序列数量。仅适用于支持多序列输出的解码策略(如束搜索的变体和采样);贪心搜索这类策略只返回单条序列。
源码佐证:generate()的完整签名(含generation_config、logits_processor、stopping_criteria、assistant_model、streamer等参数)定义在 src/transformers/generation/utils.py 的GenerationMixin.generate中。其中_get_logits_processor()会依据GenerationConfig组装出对应的 logits 处理器链——temperature对应TemperatureLogitsWarper、top_k对应TopKLogitsWarper、top_p对应TopPLogitsWarper、repetition_penalty对应RepetitionPenaltyLogitsProcessor等,这些实现集中在 src/transformers/generation/logits_process.py。_get_stopping_criteria()则负责组装MaxLengthCriteria、MaxTimeCriteria、StopStringCriteria、EosTokenCriteria等停止条件。
将自定义解码策略与模型一起保存
当希望共享带有特定生成配置的微调模型时,可按以下步骤操作:
- 创建 [
GenerationConfig] 类实例; - 设置解码策略参数;
- 使用 [
GenerationConfig.save_pretrained] 保存生成配置(config_file_name参数留空); - 设置
push_to_hub=True将配置上传到模型仓库。
>>> from transformers import AutoModelForCausalLM, GenerationConfig >>> model = AutoModelForCausalLM.from_pretrained("my_account/my_model") # doctest: +SKIP >>> generation_config = GenerationConfig( ... max_new_tokens=50, do_sample=True, top_k=50, eos_token_id=model.config.eos_token_id ... ) >>> generation_config.save_pretrained("my_account/my_model", push_to_hub=True) # doctest: +SKIP单目录多配置:config_file_name的用法
同一目录中可以保存多个生成配置,只需利用save_pretrained的config_file_name参数区分文件名,之后再用 [GenerationConfig.from_pretrained] 加载。这对「一个模型、多种用途」的场景非常实用——例如一份配置用采样实现创意文本生成,另一份用束搜索实现摘要。注意:向模型仓库追加配置文件需要相应的 Hub 权限。
>>> from transformers import AutoModelForSeq2SeqLM, AutoTokenizer, GenerationConfig >>> tokenizer = AutoTokenizer.from_pretrained("google-t5/t5-small") >>> model = AutoModelForSeq2SeqLM.from_pretrained("google-t5/t5-small") >>> translation_generation_config = GenerationConfig( ... num_beams=4, ... early_stopping=True, ... decoder_start_token_id=0, ... eos_token_id=model.config.eos_token_id, ... pad_token=model.config.pad_token_id, ... ) >>> # 提示:如需推送到 Hub,请加上 push_to_hub=True >>> translation_generation_config.save_pretrained("/tmp", "translation_generation_config.json") >>> # 使用指定的生成配置文件来参数化生成 >>> generation_config = GenerationConfig.from_pretrained("/tmp", "translation_generation_config.json") >>> inputs = tokenizer("translate English to French: Configuration files are easy to use!", return_tensors="pt") >>> outputs = model.generate(**inputs, generation_config=generation_config) >>> print(tokenizer.batch_decode(outputs, skip_special_tokens=True)) ['Les fichiers de configuration sont faciles à utiliser!']源码佐证:save_pretrained与from_pretrained的完整实现位于 src/transformers/generation/configuration_utils.py,前者支持通过config_file_name指定任意文件名(默认是generation_config.json),后者会按文件名加载并可通过from_model_config从模型配置自动构建。此外,GenerationConfig.validate()会在参数组合非法时(例如束搜索相关参数与采样参数冲突)抛出异常,保证配置在运行前就被校验。
流式输出(Streaming)
generate()通过streamer参数支持流式输出。streamer需要是一个实现了put()与end()方法的类的实例:put()用于追加新 token,end()用于标记文本生成结束。
注意:streamer 类的 API 仍在开发中,未来可能发生变化。
你可以为实现特定目的而编写自定义 streaming 类,也可以直接使用库内置的基础类。例如使用 [TextStreamer] 将generate()的输出逐词打印到屏幕:
>>> from transformers import AutoModelForCausalLM, AutoTokenizer, TextStreamer >>> tok = AutoTokenizer.from_pretrained("openai-community/gpt2") >>> model = AutoModelForCausalLM.from_pretrained("openai-community/gpt2") >>> inputs = tok(["An increasing sequence: one,"], return_tensors="pt") >>> streamer = TextStreamer(tok) >>> # streamer 除了返回常规输出外,还会把生成文本打印到标准输出(stdout) >>> _ = model.generate(**inputs, streamer=streamer, max_new_tokens=20) An increasing sequence: one, two, three, four, five, six, seven, eight, nine, ten, eleven,源码佐证:TextStreamer定义在 src/transformers/generation/streamers.py 第 42 行。它的put()方法维护一个token_cache列表,每当解码出完整单词(以空格为界的启发式规则,CJK 字符则即时输出)便调用on_finalized_text()打印;end()负责冲刷剩余缓存并输出换行。其子类TextIteratorStreamer(同文件第 157 行)把文本放入队列,供下游应用以迭代器方式非阻塞读取,非常适合构建 Gradio 等交互式演示。generate()内部的 streamer 驱动逻辑则由 src/transformers/generation/utils.py 中的GenerationStreamer与_sample/_beam_search等循环配合完成,GenerationStreamer会在每个 step 后检查是否需要向 streamer 投递 token,并统计num_tokens用于控制迭代。
解码策略详解
generate()参数(最终写入generation_config)的特定组合会激活特定的解码策略。下面逐一介绍常用策略的开启方式、适用场景与代码示例。若想系统理解各策略的数学原理,可参考经典博客How to generate text: using different decoding methods for language generation with Transformers。
贪心搜索(Greedy Search)
贪心搜索是generate()的默认策略,无需任何额外参数——等价于num_beams=1且do_sample=False,每步选取概率最高的 token。
>>> from transformers import AutoModelForCausalLM, AutoTokenizer >>> prompt = "I look forward to" >>> checkpoint = "distilbert/distilgpt2" >>> tokenizer = AutoTokenizer.from_pretrained(checkpoint) >>> inputs = tokenizer(prompt, return_tensors="pt") >>> model = AutoModelForCausalLM.from_pretrained(checkpoint) >>> outputs = model.generate(**inputs) >>> tokenizer.batch_decode(outputs, skip_special_tokens=True) ['I look forward to seeing you all again!\n\n\n\n\n\n\n\n\n\n\n']从输出可见,贪心搜索在生成长文本时很快陷入重复(连续的换行符),这正是它的典型短板。
多项式采样(Multinomial Sampling)
与贪心搜索总是选择最高概率 token 不同,多项式采样(又称祖先采样,Ancestral Sampling)根据模型给出的整个词表概率分布随机选取下一个 token。任何非零概率的 token 都有机会被选中,从而降低重复风险,产出更多样、更具创造性的文本。
启用方式:do_sample=True且num_beams=1。
>>> from transformers import AutoTokenizer, AutoModelForCausalLM, set_seed >>> set_seed(0) # 为了可复现性 >>> checkpoint = "openai-community/gpt2-large" >>> tokenizer = AutoTokenizer.from_pretrained(checkpoint) >>> model = AutoModelForCausalLM.from_pretrained(checkpoint) >>> prompt = "Today was an amazing day because" >>> inputs = tokenizer(prompt, return_tensors="pt") >>> outputs = model.generate(**inputs, do_sample=True, num_beams=1, max_new_tokens=100) >>> tokenizer.batch_decode(outputs, skip_special_tokens=True) ['Today was an amazing day because when you go to the World Cup and you don\'t, or when you don\'t get invited, that\'s a terrible feeling."']实践中常与temperature、top_k、top_p搭配使用:temperature越低分布越尖锐(趋近贪心),top_k只保留概率最高的 k 个 token,top_p(核采样)只保留累积概率达到 p 的最小 token 集合。这三者的实现分别对应 src/transformers/generation/logits_process.py 中的TemperatureLogitsWarper、TopKLogitsWarper与TopPLogitsWarper。
束搜索解码(Beam Search Decoding)
与贪心搜索不同,束搜索在每个时间步同时维护多个假设(束),最终选择整条序列上概率最高的假设。它可以「向前看」,识别出那些以低概率 token 开头、却会被贪心搜索忽略的高分序列,因此特别适合以输入为锚定的任务,如图像描述、语音识别、翻译等。
启用方式:num_beams设为大于 1(即追踪的假设数量)。
>>> from transformers import AutoModelForCausalLM, AutoTokenizer >>> prompt = "It is astonishing how one can" >>> checkpoint = "openai-community/gpt2-medium" >>> tokenizer = AutoTokenizer.from_pretrained(checkpoint) >>> inputs = tokenizer(prompt, return_tensors="pt") >>> model = AutoModelForCausalLM.from_pretrained(checkpoint) >>> outputs = model.generate(**inputs, num_beams=5, max_new_tokens=50) >>> tokenizer.batch_decode(outputs, skip_special_tokens=True) ['It is astonishing how one can have such a profound impact on the lives of so many people in such a short period of time."\n\nHe added: "I am very proud of the work I have been able to do in the last few years.\n\n"I have']束搜索的核心循环_beam_search()位于 src/transformers/generation/utils.py,它通过beam_scores累积各束的对数概率,并配合length_penalty(长度惩罚)与early_stopping(是否提前停止)决定最终取舍——early_stopping=True会在所有束都达到 EOS 时立即停止,early_stopping="never"则会等所有束都生成完毕。
束搜索多项式采样(Beam-Search Multinomial Sampling)
顾名思义,该策略结合了束搜索与多项式采样:num_beams大于 1 且do_sample=True。与纯束搜索相比,它允许在每步进行随机采样,但仍会在束之间剪掉低概率序列。
>>> from transformers import AutoTokenizer, AutoModelForSeq2SeqLM, set_seed >>> set_seed(0) # 为了可复现性 >>> prompt = "translate English to German: The house is wonderful." >>> checkpoint = "google-t5/t5-small" >>> tokenizer = AutoTokenizer.from_pretrained(checkpoint) >>> inputs = tokenizer(prompt, return_tensors="pt") >>> model = AutoModelForSeq2SeqLM.from_pretrained(checkpoint) >>> outputs = model.generate(**inputs, num_beams=5, do_sample=True) >>> tokenizer.decode(outputs[0], skip_special_tokens=True) 'Das Haus ist wunderbar.'推测解码(Speculative Decoding / Assisted Decoding)
推测解码(又称辅助解码)利用一个共享同一 tokenizer 的、小得多的辅助模型(assistant model)先生成若干候选 token,主模型通过单次前向传播同时验证这些候选,从而加速解码过程。当do_sample=True时,采用推测解码论文(论文编号 2211.17192)中提出的 token 验证与重采样方法。
当前辅助解码仅支持贪心搜索与采样两种模式,且不支持批处理输入。
启用方式:在generate()中传入assistant_model参数。
>>> from transformers import AutoModelForCausalLM, AutoTokenizer >>> prompt = "Alice and Bob" >>> checkpoint = "EleutherAI/pythia-1.4b-deduped" >>> assistant_checkpoint = "EleutherAI/pythia-160m-deduped" >>> tokenizer = AutoTokenizer.from_pretrained(checkpoint) >>> inputs = tokenizer(prompt, return_tensors="pt") >>> model = AutoModelForCausalLM.from_pretrained(checkpoint) >>> assistant_model = AutoModelForCausalLM.from_pretrained(assistant_checkpoint) >>> outputs = model.generate(**inputs, assistant_model=assistant_model) >>> tokenizer.batch_decode(outputs, skip_special_tokens=True) ['Alice and Bob are sitting in a bar. Alice is drinking a beer and Bob is drinking a']与采样配合使用时,同样可以通过temperature控制随机性;不过在推测解码中,降低temperature还有助于改善延迟(更低的温度使辅助模型的候选更容易被主模型接受,减少无效的候选生成)。
>>> from transformers import AutoModelForCausalLM, AutoTokenizer, set_seed >>> set_seed(42) # 为了可复现性 >>> prompt = "Alice and Bob" >>> checkpoint = "EleutherAI/pythia-1.4b-deduped" >>> assistant_checkpoint = "EleutherAI/pythia-160m-deduped" >>> tokenizer = AutoTokenizer.from_pretrained(checkpoint) >>> inputs = tokenizer(prompt, return_tensors="pt") >>> model = AutoModelForCausalLM.from_pretrained(checkpoint) >>> assistant_model = AutoModelForCausalLM.from_pretrained(assistant_checkpoint) >>> outputs = model.generate(**inputs, assistant_model=assistant_model, do_sample=True, temperature=0.5) >>> tokenizer.batch_decode(outputs, skip_special_tokens=True) ['Alice and Bob, who were both in their early twenties, were both in the process of']源码佐证:推测解码的候选生成逻辑位于 src/transformers/generation/candidate_generator.py 的AssistedCandidateGenerator(以及面向不同 tokenizer 的PromptLookupCandidateGenerator、AsyncAssistedCandidateGenerator等),generate()内部通过_get_candidate_generator()按配置选择候选生成器,再由_assisted_decoding()循环执行「辅助模型生成候选 → 主模型并行验证 → 接受匹配 token」的流程。主模型验证候选时逐 token 比对分布,一旦出现不一致即从该位置重新接管生成,从而在保证输出分布与原始解码一致的前提下获得加速。
总结与选型建议
- 短输出、确定性优先:默认贪心搜索即可;需要更高整体分数时切换束搜索(
num_beams>1)。 - 长文本、创意写作:使用采样(
do_sample=True),并配合temperature、top_k、top_p控制多样性与稳定性。 - 翻译 / 摘要 / 图像描述等以输入为锚定的任务:优先考虑束搜索及其
length_penalty、early_stopping调优。 - 对延迟敏感的服务场景:尝试推测解码(
assistant_model),并适度降低temperature。 - 配置共享与复现:用
GenerationConfig固化参数组合,通过save_pretrained(..., push_to_hub=True)随模型共享;config_file_name支持单模型多配置。 - 交互式体验:用
TextStreamer/TextIteratorStreamer实现逐词流式输出。
更完整的参数说明可查阅 docs/source/en/generation_strategies.md(英文版),所有解码策略的最终执行逻辑均可回溯到 src/transformers/generation/utils.py 的generate及其内部循环,以及 src/transformers/generation/logits_process.py、src/transformers/generation/stopping_criteria.py、src/transformers/generation/streamers.py 等模块,读者可按需深入研读。
【免费下载链接】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),仅供参考