Hugging Face Transformers 文本生成策略完全指南:从贪心搜索到辅助解码 Hugging Face Transformers 文本生成策略完全指南从贪心搜索到辅助解码【免费下载链接】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本文基于 Transformers 仓库 generation_strategies.md 官方文档展开系统讲解generate()方法背后默认文本生成配置的运作机制、六种主流解码策略贪心搜索、多项式采样、波束搜索、波束多项式采样、辅助解码及其核心参数并演示如何将自定义生成配置随模型保存与共享、如何通过streamer实现逐词流式输出。读完本文你将能够根据任务场景摘要、翻译、开放式生成正确选择解码策略并调优参数同时理解这些配置在 generation/configuration_utils.py 与 generation/utils.py 中的底层实现原理。文本生成与解码策略概述文本生成是开放式文本生成、摘要、翻译等众多自然语言处理任务的核心环节同时也支撑着以文本为输出的多模态应用例如语音转文本、图像转文本。在 Transformers 中能够进行文本生成的模型包括 GPT-2、XLNet、OpenAI GPT、CTRL、Transformer-XL、XLM、BART、T5、GIT、Whisper 等。针对不同任务官方文档给出了使用 [~generation.GenerationMixin.generate] 方法生成文本输出的典型示例文本摘要在pipeline中调用generate完成摘要生成图像描述GIT 模型将图像特征与文本联合生成描述语音转录Whisper 模型将音频转为文本。generate()的输入取决于模型的模态modality通常由AutoTokenizer或AutoProcessor等模型对应的预处理器类返回。当模型的预处理器产生多种输入如图像像素值、input_ids、注意力掩码等时需要把所有输入一并传给generate()每个模型预处理器的详细说明见对应模型的文档。从概率分布中选择下一个 token 的过程被称为解码decodinggenerate()使用的解码策略可以被完全自定义。修改解码策略不会改变任何可训练参数的值但会对生成文本的质量产生显著影响——例如减少文本中的重复、让文本更连贯。本指南将依次覆盖三部分内容默认的文本生成配置、常见的解码策略及其关键参数、以及如何把自定义生成配置保存并共享到 Hub。默认文本生成配置模型的解码策略由它的**生成配置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时只会显示与默认生成配置不同的字段值未列出的字段均采用默认值。例如上面的输出仅显示了bos_token_id与eos_token_id两个字段其余全部是库内默认值。关于默认行为有两个重要的实现事实需要了解输出长度上限默认生成配置将输出大小限制为与输入提示组合后最多 20 个 token以避免触发资源限制。在 configuration_utils.py 中可以看到默认的max_length即为 20max_length: 20其含义是输入 输出的总长度若希望只限制新增 token 的数量应使用max_new_tokens。默认解码策略默认是贪心搜索greedy search即每一步都选择概率最高的 token 作为下一个 token这是最简单的解码方式。对于许多任务和小规模输出它都能工作得很好但当用于生成长输出时贪心搜索容易产生高度重复的结果。在 generation/utils.py 的GenerationMixin类文档字符串中明确给出了解码策略与参数组合的对应关系num_beams1且do_sampleFalse→贪心解码greedy decodingnum_beams1且do_sampleTrue→多项式采样multinomial samplingnum_beams1且do_sampleFalse→波束搜索解码beam-search decodingnum_beams1且do_sampleTrue→波束多项式采样beam-search multinomial sampling传入assistant_model或prompt_lookup_num_tokens→辅助解码assisted decoding。自定义文本生成直接传参覆盖通过向generate()直接传入参数名与取值可以覆盖generation_config中的对应设置 my_model.generate(**inputs, num_beams4, do_sampleTrue) # doctest: SKIP默认解码策略虽在多数任务中表现良好但以下几个参数是最常被微调的max_new_tokens要生成的新 token 的最大数量即输出序列的大小不包含提示中的 token。它与max_length的区别在于max_length统计输入与输出的总长度而max_new_tokens只统计新生成的 token 数更符合直觉建议优先使用。num_beams将num_beams设为大于 1 的值即可从贪心搜索切换到波束搜索。该策略在每个时间步评估多个假设hypotheses最终选择整个序列上概率最高的假设从而避免贪心搜索因起始 token 概率偏低而忽略高概率序列的问题。do_sample设为True后启用基于概率分布的采样类解码策略包括多项式采样、波束多项式采样、Top-K 采样、Top-p核采样等。这些策略在整词表概率分布上按各自规则挑选下一个 token含各自策略专属的调节项如top_k、top_p、temperature。num_return_sequences每个输入要返回的候选序列数量。它仅适用于支持多候选序列的解码策略如波束搜索、各类采样变体贪心搜索、对比搜索等只返回单条输出序列的策略不能使用该参数。从实现角度看generate()是 generation/utils.py 中GenerationMixin的核心入口约 L2388 起内部根据上述参数组合分派到_sample、_beam_search、_assisted_decoding等具体解码循环生成配置类GenerationConfig定义于 configuration_utils.pyL100 起并在__init__中逐一解析max_length、max_new_tokens、num_beams、do_sample、top_k、top_p、temperature等字段。保存自定义解码策略并随模型共享如果你在特定生成配置下调优了模型并希望与他人共享可以按以下步骤操作创建 [GenerationConfig] 类的一个实例指定解码策略参数调用 [GenerationConfig.save_pretrained] 保存生成配置注意保留config_file_name参数的默认值即generation_config.json将push_to_hub设为True把配置上传到模型的仓库。 from transformers import AutoModelForCausalLM, GenerationConfig model AutoModelForCausalLM.from_pretrained(my_account/my_model) # doctest: SKIP generation_config GenerationConfig( ... max_new_tokens50, do_sampleTrue, top_k50, eos_token_idmodel.config.eos_token_id ... ) generation_config.save_pretrained(my_account/my_model, push_to_hubTrue) # doctest: SKIP关于save_pretrained的实现细节从 configuration_utils.pyL890 起的源码可以看到config_file_name默认为generation_config.json常量GENERATION_CONFIG_NAME保存前会调用self.validate(strictTrue)对配置做严格校验若存在非法字段组合将直接抛出ValueError并拒绝保存——这是为了防止坏配置被保存后反复复用保存内容通过to_json_file(..., use_diffTrue)写入即采用 diff 形式只记录与默认值的差异因此加载后打印出来的也是差异字段当push_to_hubTrue时会以save_directory的目录名为默认repo_id创建仓库并上传文件。在同一目录保存多个生成配置也是支持的利用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_beams4, ... early_stoppingTrue, ... decoder_start_token_id0, ... eos_token_idmodel.config.eos_token_id, ... pad_tokenmodel.config.pad_token_id, ... ) # Tip: add push_to_hubTrue to push to the Hub translation_generation_config.save_pretrained(/tmp, translation_generation_config.json) # You could then use the named generation config file to parameterize generation generation_config GenerationConfig.from_pretrained(/tmp, translation_generation_config.json) inputs tokenizer(translate English to French: Configuration files are easy to use!, return_tensorspt) outputs model.generate(**inputs, generation_configgeneration_config) print(tokenizer.batch_decode(outputs, skip_special_tokensTrue)) [Les fichiers de configuration sont faciles à utiliser!]对应地from_pretrainedconfiguration_utils.py L950 起是一个类方法支持从模型仓库 id 或本地目录加载配置同时接受cache_dir、force_download、local_files_only、token、revision、subfolder等参数kwargs中与配置属性同名的键会覆盖已加载的值。流式输出Streaminggenerate()通过其streamer输入支持流式输出。streamer接受任何实现了put()与end()两个方法的类的实例内部用put()推送新 token用end()标记生成结束。注意流式类streamer的 API 仍在开发中未来可能发生变化。实践中你可以为各种目的编写自己的流式类仓库也提供了开箱即用的基础流式类。例如用 [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_tensorspt) streamer TextStreamer(tok) # Despite returning the usual output, the streamer will also print the generated text to stdout. _ model.generate(**inputs, streamerstreamer, max_new_tokens20) An increasing sequence: one, two, three, four, five, six, seven, eight, nine, ten, eleven,从源码看流式体系定义在 generation/streamers.pyBaseStreamerL28 起是基类声明了put(value)与end()两个接口直接继承它即可实现自定义流式类TextStreamerL42 起把 token 累积在token_cache中用启发式策略确定可打印文本遇到换行符立即冲刷缓存遇到 CJK 字符立即打印否则一直打印到最后一个空格为止避免打印不完整的单词end()负责冲刷剩余缓存并打印换行TextIteratorStreamer、AsyncTextIteratorStreamer等变体则在生成的同时把文本放入队列/异步迭代器便于与 Web 框架或异步应用集成。解码策略详解特定的generate()参数组合最终沉淀为generation_config用于启用特定的解码策略。下面逐一介绍控制解码策略的参数及其用法。贪心搜索Greedy Searchgenerate()默认使用贪心搜索解码因此无需传任何参数即可启用等价于num_beams1且do_sampleFalse。 from transformers import AutoModelForCausalLM, AutoTokenizer prompt I look forward to checkpoint distilbert/distilgpt2 tokenizer AutoTokenizer.from_pretrained(checkpoint) inputs tokenizer(prompt, return_tensorspt) model AutoModelForCausalLM.from_pretrained(checkpoint) outputs model.generate(**inputs) tokenizer.batch_decode(outputs, skip_special_tokensTrue) [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_sampleTrue且num_beams1。 from transformers import AutoTokenizer, AutoModelForCausalLM, set_seed set_seed(0) # For reproducibility 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_tensorspt) outputs model.generate(**inputs, do_sampleTrue, num_beams1, max_new_tokens100) tokenizer.batch_decode(outputs, skip_special_tokensTrue) [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.]由于引入了随机性官方示例统一使用set_seed(...)固定随机种子以保证结果可复现——在实际工程中若需要稳定的输出请同样设置随机种子。采样质量还可通过temperature温度调节分布陡峭程度、top_k只从概率最高的 K 个 token 中采样与top_p核采样只从累计概率达到 p 的最小 token 集合中采样进一步控制。波束搜索解码Beam-search Decoding与贪心搜索不同波束搜索解码在每个时间步保留多个假设最终选择在整个序列上概率最高的假设。它的优势在于能够发现那些起始 token 概率较低、因而被贪心搜索忽略的高概率序列通常配合early_stoppingTrue使用。启用方式将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_tensorspt) model AutoModelForCausalLM.from_pretrained(checkpoint) outputs model.generate(**inputs, num_beams5, max_new_tokens50) tokenizer.batch_decode(outputs, skip_special_tokensTrue) [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\nI have]注意波束搜索的计算成本与num_beams大致成正比——跟踪的假设越多每一步需要计算的前向次数越多。在 generation/utils.py 中波束搜索由_beam_search方法约 L3362 起实现它会维护 beam 数量个假设并按整体序列概率进行剪枝与排序。波束多项式采样Beam-search Multinomial Sampling顾名思义该策略结合了波束搜索与多项式采样。启用方式num_beams设为大于 1 的值且do_sampleTrue。 from transformers import AutoTokenizer, AutoModelForSeq2SeqLM, set_seed set_seed(0) # For reproducibility prompt translate English to German: The house is wonderful. checkpoint google-t5/t5-small tokenizer AutoTokenizer.from_pretrained(checkpoint) inputs tokenizer(prompt, return_tensorspt) model AutoModelForSeq2SeqLM.from_pretrained(checkpoint) outputs model.generate(**inputs, num_beams5, do_sampleTrue) tokenizer.decode(outputs[0], skip_special_tokensTrue) Das Haus ist wunderbar.示例使用了 T5AutoModelForSeq2SeqLM演示了在翻译任务上的应用在多个波束内分别做采样兼顾波束搜索的全局最优倾向与采样的多样性。辅助解码Assisted Decoding辅助解码是对上述解码策略的一种加速改造使用一个共享同一分词器理想情况下小得多的辅助模型assistant model先贪心地生成若干候选 token然后主模型通过**一次前向传播single forward pass**验证这些候选 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_tensorspt) model AutoModelForCausalLM.from_pretrained(checkpoint) assistant_model AutoModelForCausalLM.from_pretrained(assistant_checkpoint) outputs model.generate(**inputs, assistant_modelassistant_model) tokenizer.batch_decode(outputs, skip_special_tokensTrue) [Alice and Bob are sitting in a bar. Alice is drinking a beer and Bob is drinking a]从实现上看辅助解码对应 generation/utils.py 中的_assisted_decoding方法约 L3716 起主模型一次前向即可确认辅助模型提出的多个候选 token被接受的候选越多、加速收益越大仓库同时提供了prompt_lookup_num_tokens参数可以不依赖辅助模型、基于 n-gram 匹配来提出候选。使用采样模式时辅助解码同样支持temperature参数来控制随机性与多项式采样类似。不过在辅助解码中调低温度还有助于改善延迟——因为低温使分布更尖锐辅助模型提出的候选被主模型接受的概率更高从而减少拒绝重试的轮次。 from transformers import AutoModelForCausalLM, AutoTokenizer, set_seed set_seed(42) # For reproducibility 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_tensorspt) model AutoModelForCausalLM.from_pretrained(checkpoint) assistant_model AutoModelForCausalLM.from_pretrained(assistant_checkpoint) outputs model.generate(**inputs, assistant_modelassistant_model, do_sampleTrue, temperature0.5) tokenizer.batch_decode(outputs, skip_special_tokensTrue) [Alice and Bob are going to the same party. It is a small party, in a small]小结与进一步阅读本指南介绍了启用各种解码策略所需的核心参数贪心搜索默认无需传参、多项式采样do_sampleTrue, num_beams1、波束搜索num_beams1、波束多项式采样num_beams1, do_sampleTrue与辅助解码assistant_model...。generate()方法还提供了更多高级参数用于进一步控制生成行为完整参数列表参见 API 文档Text Generation。选型建议速查追求稳定与可复现、输出较短时用默认贪心搜索开放式创意生成用采样配合top_k/top_p/temperature摘要、翻译等追求全局最优的任务用波束搜索可加early_stoppingTrue追求低延迟推理且具备小模型可用时优先尝试辅助解码。【免费下载链接】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),仅供参考