ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

Transformers 文本生成策略实战:贪心搜索、采样、束搜索与 custom_generate 自定义生成方法

Transformers 文本生成策略实战:贪心搜索、采样、束搜索与 custom_generate 自定义生成方法 Transformers 文本生成策略实战贪心搜索、采样、束搜索与 custom_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解码策略decoding strategy决定了模型如何选择下一个要生成的 token。Transformers 提供了从贪心搜索、多项式采样到束搜索等多种内置解码方法并在此基础上提供了custom_generate机制允许你把任意解码逻辑如“模型不确定时继续思考”、“生成卡住时回滚”、处理特殊 token 的自定义逻辑、使用专用 KV cache 等打包成一个 Hub 仓库或本地目录注入到任何模型中。读完本文你将掌握 Transformers 内置解码策略的选择与调用方式并能够创建、测试和分发自己的自定义生成方法。解码策略决定生成质量解码策略直接影响生成文本的质量对短输出、低创造性要求的任务贪心搜索通常足够对需要多样性的创作类任务采样类方法更合适对图像描述、语音识别等“以输入为基准”的任务束搜索能取得更好的整体概率。官方文档 generation_strategies.md 将策略分为两大类基础解码方法Basic decoding methods贪心搜索、采样Sampling、束搜索Beam search是所有文本生成任务的起点自定义生成方法Custom generation methods通过custom_generate机制扩展的专用行为。在源码层面这些方法统一由GenerationMixin.generate入口调度。src/transformers/generation/utils.py 中的GENERATION_MODES_MAPPING维护了生成模式到具体实现函数的映射GENERATION_MODES_MAPPING { GenerationMode.SAMPLE: _sample, GenerationMode.GREEDY_SEARCH: _sample, GenerationMode.BEAM_SEARCH: _beam_search, GenerationMode.BEAM_SAMPLE: _beam_search, GenerationMode.ASSISTED_GENERATION: _assisted_decoding, # Deprecated methods GenerationMode.DOLA_GENERATION: transformers-community/dola, GenerationMode.CONTRASTIVE_SEARCH: transformers-community/contrastive-search, GenerationMode.GROUP_BEAM_SEARCH: transformers-community/group-beam-search, GenerationMode.CONSTRAINED_BEAM_SEARCH: transformers-community/constrained-beam-search, }从中可以看到两点事实贪心搜索和采样共用_sample实现区别仅在do_sample参数而 DOLA、对比搜索contrastive search、组束搜索group beam search、约束束搜索constrained beam search等已不在核心实现中而是被迁移到了custom_generate仓库由 _get_deprecated_gen_repo 处理并提示用户显式传入custom_generate仓库名v4.62.0 后将移除兼容逻辑。这正是自定义生成机制的典型应用案例。基础解码方法贪心搜索Greedy search贪心搜索是默认解码策略每一步都选取概率最高的 token。除非在GenerationConfig中另行指定该策略默认最多生成 20 个新 tokenmax_new_tokens默认值为 20。贪心搜索适合输出相对较短、不追求创造性的任务但生成较长序列时容易开始重复自己。import torch from transformers import AutoModelForCausalLM, AutoTokenizer from accelerate import Accelerator device Accelerator().device tokenizer AutoTokenizer.from_pretrained(meta-llama/Llama-2-7b-hf) inputs tokenizer(Hugging Face is an open-source company, return_tensorspt).to(device) model AutoModelForCausalLM.from_pretrained(meta-llama/Llama-2-7b-hf, dtypetorch.float16).to(device) # explicitly set to default length because Llama2 generation length is 4096 outputs model.generate(**inputs, max_new_tokens20) tokenizer.batch_decode(outputs, skip_special_tokensTrue) Hugging Face is an open-source company that provides a suite of tools and services for building, deploying, and maintaining natural language processing采样Sampling采样又称多项式采样multinomial sampling不是选概率最高的 token而是按照整个词表上的概率分布随机抽取一个 token——只要某个 token 概率非零就有机会被选中。采样类策略能减少重复、产生更有创造力和多样性的输出。启用方式do_sampleTrue且num_beams1。import torch from transformers import AutoModelForCausalLM, AutoTokenizer from accelerate import Accelerator device Accelerator().device tokenizer AutoTokenizer.from_pretrained(meta-llama/Llama-2-7b-hf) inputs tokenizer(Hugging Face is an open-source company, return_tensorspt).to(device) model AutoModelForCausalLM.from_pretrained(meta-llama/Llama-2-7b-hf, dtypetorch.float16).to(device) # explicitly set to 100 because Llama2 generation length is 4096 outputs model.generate(**inputs, max_new_tokens50, do_sampleTrue, num_beams1) tokenizer.batch_decode(outputs, skip_special_tokensTrue) Hugging Face is an open-source company \nWe are open-source and believe that open-source is the best way to build technology. Our mission is to make AI accessible to everyone, and we believe that open-source is the best way to achieve that.束搜索Beam search束搜索在每个时间步同时维护多条生成序列beam在若干步之后选取整体概率最高的序列。与贪心搜索不同它具备“向前看”的能力即使某条序列开头的 token 概率较低只要整体序列概率更高也可能被选中。它最适合以输入为基准input-grounded的任务例如图像描述、语音识别。也可以配合do_sampleTrue使用束搜索每步内部进行采样但束搜索仍会在步骤之间贪心地剪掉低概率序列。启用方式设置num_beams参数必须大于 1否则等价于贪心搜索。import torch from transformers import AutoModelForCausalLM, AutoTokenizer from accelerate import Accelerator device Accelerator().device tokenizer AutoTokenizer.from_pretrained(meta-llama/Llama-2-7b-hf) inputs tokenizer(Hugging Face is an open-source company, return_tensorspt).to(device) model AutoModelForCausalLM.from_pretrained(meta-llama/Llama-2-7b-hf, dtypetorch.float16).to(device) # explicitly set to 100 because Llama2 generation length is 4096 outputs model.generate(**inputs, max_new_tokens50, num_beams2) tokenizer.batch_decode(outputs, skip_special_tokensTrue) [Hugging Face is an open-source company that develops and maintains the Hugging Face platform, which is a collection of tools and libraries for building and deploying natural language processing (NLP) models. Hugging Face was founded in 2018 by Thomas Wolf]三种基础策略的参数速查策略关键参数适用场景局限贪心搜索默认max_new_tokens默认 20短输出、确定性任务长序列易重复采样do_sampleTrue, num_beams1创意性、多样性输出结果不可复现受随机性影响束搜索num_beams1可选do_sampleTrue图像描述、ASR 等以输入为基准的任务计算量随 beam 数增加自定义生成方法custom_generate当内置方法无法满足需求时——例如希望模型不确定时继续“思考”、生成卡住时回滚、用自定义逻辑处理特殊 token、或使用专用 KV cache——可以用custom_generate机制扩展生成行为。这是对 自定义模型代码 能力的进一步延伸同样要求设置trust_remote_codeTrue。该机制有两种使用形态形态一加载自带自定义生成方法的模型仓库如果某个模型仓库内置了自定义生成方法仓库内含custom_generate/目录加载它时generate会被自动覆盖。从源码看from_pretrained加载流程中会尝试调用load_custom_generate成功则用functools.partial替换self.generate见 GenerationMixin.from_pretrained# 加载自定义生成方法如果 pretrained_model_name_or_path 定义了它并覆盖 generate if hasattr(self, load_custom_generate) and trust_remote_code: try: custom_generate self.load_custom_generate( pretrained_model_name_or_path, trust_remote_codetrust_remote_code, **repo_loading_kwargs ) self.generate functools.partial(custom_generate, modelself) except OSError: # 不存在自定义 generate 函数 pass示例transformers-community/custom_generate_example仓库是Qwen/Qwen2.5-0.5B-Instruct的一份副本但附带了自定义生成代码——直接调用generate就会使用它from transformers import AutoModelForCausalLM, AutoTokenizer # transformers-community/custom_generate_example 是 Qwen/Qwen2.5-0.5B-Instruct 的副本 # 但带有自定义生成代码 - 调用 generate 即使用自定义生成方法 tokenizer AutoTokenizer.from_pretrained(transformers-community/custom_generate_example) model AutoModelForCausalLM.from_pretrained( transformers-community/custom_generate_example, device_mapauto, trust_remote_codeTrue ) inputs tokenizer([The quick brown], return_tensorspt).to(model.device) # 自定义生成方法是一个最简贪心解码实现运行时还会打印一条自定义消息 gen_out model.generate(**inputs) # 此时应能看到它的自定义消息✨ using a custom generation method ✨ print(tokenizer.batch_decode(gen_out, skip_special_tokensTrue)) The quick brown fox jumps over a lazy dog, and the dog is a type of animal. Is形态二任意模型通过custom_generate参数注入方法自定义生成方法还有一个关键特性它可以从任何模型加载。~GenerationMixin.generate提供了custom_generate参数任何人都能创建并分享可作用于任意 Transformers 模型的自定义生成方法用户无需安装额外的 Python 包from transformers import AutoModelForCausalLM, AutoTokenizer tokenizer AutoTokenizer.from_pretrained(Qwen/Qwen2.5-0.5B-Instruct) model AutoModelForCausalLM.from_pretrained(Qwen/Qwen2.5-0.5B-Instruct, device_mapauto) inputs tokenizer([The quick brown], return_tensorspt).to(model.device) # custom_generate 用 transformers-community/custom_generate_example 中定义的 # 自定义生成方法替换原有的 generate gen_out model.generate(**inputs, custom_generatetransformers-community/custom_generate_example, trust_remote_codeTrue) print(tokenizer.batch_decode(gen_out, skip_special_tokensTrue)[0]) The quick brown fox jumps over a lazy dog, and the dog is a type of animal. Is从源码看generate在处理任何输入准备之前“0.a”步骤就会拦截字符串形式的custom_generate收集除self、kwargs、trust_remote_code、custom_generate之外的全部参数交给load_custom_generate加载出的函数并把modelself一并转发见 generate 的 custom_generate 分支if custom_generate is not None and isinstance(custom_generate, str): global_keys_to_exclude {self, kwargs, global_keys_to_exclude, trust_remote_code, custom_generate} generate_arguments {key: value for key, value in locals().items() if key not in global_keys_to_exclude} generate_arguments.update(kwargs) custom_generate_function self.load_custom_generate( custom_generate, trust_remote_codetrust_remote_code, **kwargs ) return custom_generate_function(modelself, **generate_arguments)也就是说你的自定义generate收到的参数与原生generate完全一致只是把self换成了model并且可以访问GenerationMixin定义的所有属性和方法。使用自定义方法前应阅读该仓库的README.md确认是否有新的输入参数或输出类型差异如果没有可认为其行为与基础generate一致。以transformers-community/custom_generate_example为例其 README 声明了一个额外参数left_padding在 prompt 前添加若干 pad tokengen_out model.generate( **inputs, custom_generatetransformers-community/custom_generate_example, trust_remote_codeTrue, left_padding5 ) print(tokenizer.batch_decode(gen_out)[0]) The quick brown fox jumps over the lazy dog.\n\nThe sentence The quick依赖检查requirements 缺失时的报错如果自定义方法固定了当前环境不满足的 Python 依赖load_custom_generate会先执行check_python_requirements校验custom_generate/requirements.txt见 load_custom_generate 实现不满足则抛出异常。例如transformers-community/custom_generate_bad_requirements仓库定义了不可能满足的依赖运行会得到类似报错ImportError: Missing requirements in your local environment for transformers-community/custom_generate_bad_requirements: foo (installed: None) bar0.0.0 (installed: None) torch99.0 (installed: 2.6.0)按提示更新环境依赖即可消除该错误。相关行为在测试中有覆盖见 tests/generation/test_utils.py 中的test_custom_generate_bad_requirements等用例。创建自定义生成方法创建一个新生成方法需要建立一个新的模型仓库并推送以下文件你设计生成方法所用的模型custom_generate/generate.py——自定义生成方法的全部逻辑custom_generate/requirements.txt——可选的额外 Python 依赖及版本锁定README.md——添加custom_generate标签并记录新方法的所有新参数与输出类型差异。仓库结构如下your_repo/ ├── README.md # include the custom_generate tag ├── config.json ├── ... └── custom_generate/ ├── generate.py └── requirements.txt添加基础模型起点就是一个普通的模型仓库。应放入你设计该方法时使用的模型它与生成方法构成一个可独立工作的“模型-生成”对。加载该仓库的模型时你的自定义方法会覆盖generate但方法本身仍可按上文方式加载到任意其他 Transformers 模型上。如果只是复制现有模型from transformers import AutoModelForCausalLM, AutoTokenizer tokenizer AutoTokenizer.from_pretrained(source/model_repo) model AutoModelForCausalLM.from_pretrained(source/model_repo) tokenizer.save_pretrained(your/generation_method, push_to_hubTrue) model.save_pretrained(your/generation_method, push_to_hubTrue)generate.py这是方法的核心。它必须包含一个名为generate的函数且该函数必须以model作为第一个参数。model就是模型实例因此你可以访问模型的全部属性和方法包括GenerationMixin中定义的如基础generate方法。注意generate.py必须放在名为custom_generate的目录内不能放在仓库根目录——这些文件路径在框架中是硬编码的对应get_cached_module_file(..., module_filecustom_generate/generate.py)的调用。底层流程是当基础generate收到custom_generate参数时先检查 Python 依赖如有再定位generate.py中的自定义generate最后调用它。除用于触发该机制的trust_remote_code和custom_generate两个参数外其余收到的参数与model会全部转发给你的函数。因此你的generate可以混用原有参数与自定义参数甚至返回不同输出类型import torch def generate(model, input_ids, generation_configNone, left_paddingNone, **kwargs): generation_config generation_config or model.generation_config # 回落到模型的生成配置 cur_length input_ids.shape[1] max_length generation_config.max_length or cur_length generation_config.max_new_tokens # 自定义参数示例在 prompt 前添加 left_padding整数个pad token if left_padding is not None: if not isinstance(left_padding, int) or left_padding 0: raise ValueError(fleft_padding must be an integer larger than 0, but is {left_padding}) pad_token kwargs.pop(pad_token, None) or generation_config.pad_token_id or model.config.pad_token_id if pad_token is None: raise ValueError(pad_token is not defined) batch_size input_ids.shape[0] pad_tensor torch.full(size(batch_size, left_padding), fill_valuepad_token).to(input_ids.device) input_ids torch.cat((pad_tensor, input_ids), dim1) cur_length input_ids.shape[1] # 最简贪心解码循环 while cur_length max_length: logits model(input_ids).logits next_token_logits logits[:, -1, :] next_tokens torch.argmax(next_token_logits, dim-1) input_ids torch.cat((input_ids, next_tokens[:, None]), dim-1) cur_length 1 return input_ids推荐实践可以放心复用原生generate中参数校验与输入准备的逻辑如果使用了model上的私有方法/属性应在 requirements 中锁定transformers版本建议加入模型/输入校验甚至单独的测试文件方便用户在自己环境中做健全性检查。本地开发与相对导入自定义generate可以相对导入custom_generate目录内的代码例如存在utils.py时from .utils import some_function只支持与custom_generate同层的相对导入父目录/兄弟目录导入无效。另外custom_generate参数同样支持本地目录——任何包含custom_generate结构的目录都可以直接传入这是开发自定义方法时推荐的工作流gen_out model.generate(**inputs, custom_generatepath/to/local/dir, trust_remote_codeTrue)警告加载本地目录同样会执行其中的custom_generate/generate.py因此与 Hub 仓库一样必须trust_remote_codeTrue。请只对你自己编写或审查过的代码开启该选项。这一点在 tests/generation/test_utils.py 的本地目录相关测试中得到验证。requirements.txt可在custom_generate目录内提供requirements.txt指定额外 Python 依赖。这些依赖在运行时被检查缺失时会抛出异常提示用户更新环境即前文的ImportError行为。README.md模型仓库根目录的README.md通常描述模型但既然该仓库的核心是自定义生成方法强烈建议把重心转向方法本身的说明并记录相对于原生generate的输入/输出差异——用户可以聚焦“新在哪里”通用实现细节仍依赖 Transformers 文档。为便于发现建议在 README 顶部添加custom_generate标签--- library_name: transformers tags: - custom_generate --- (your markdown content here)README 推荐实践记录相对于原生generate的输入/输出差异提供自包含示例方便快速实验说明软性约束例如该方法只在某类模型家族上效果良好。复用generate的输入准备传入可调用对象如果你想新增一个解码循环但想保留generate里已有的输入准备逻辑batch 扩展、attention mask、logits processors、stopping criteria 等可以给custom_generate传一个可调用对象Callablegenerate会执行完整的标准准备流程然后调用你提供的可调用对象来运行解码循环从而只覆盖解码循环本身。此时generate会先执行完整的输入准备再调用可调用对象并自动比对可调用对象的签名以提取新增参数见 _extract_generation_mode_kwargs。def custom_loop(model, input_ids, attention_mask, logits_processor, stopping_criteria, generation_config, **model_kwargs): next_tokens input_ids while input_ids.shape[1] stopping_criteria[0].max_length: logits model(next_tokens, attention_maskattention_mask, **model_kwargs).logits next_token_logits logits_processor(input_ids, logits[:, -1, :]) next_tokens torch.argmax(next_token_logits, dim-1)[:, None] input_ids torch.cat((input_ids, next_tokens), dim-1) attention_mask torch.cat((attention_mask, torch.ones_like(next_tokens)), dim-1) return input_ids output model.generate( **inputs, custom_generatecustom_loop, max_new_tokens10, )提示如果发布custom_generate仓库你的generate实现内部同样可以定义一个可调用对象并传给model.generate()这样既能自定义解码循环又能享受 Transformers 内建的输入准备逻辑。如何发现自定义生成方法在模型库中搜索custom_generate标签即可找到全部自定义生成方法。除标签外官方还维护了两个精选集合社区贡献的方法集合以及包含“此前属于 transformers 内置、现迁移为 custom_generate 仓库的参考实现”的教程集合。如前文GENERATION_MODES_MAPPING所示DOLA、contrastive search、group beam search、constrained beam search 均已以这种形式迁移。策略选择与验证小结默认/短输出、要确定性贪心搜索默认max_new_tokens20注意显式设置上限很多模型默认生成长度远大于 20如 Llama2 为 4096要多样性/创造性do_sampleTrue, num_beams1输入基准任务描述、ASRnum_beams1需要特殊解码逻辑优先评估是否可只覆盖解码循环传 Callable复用输入准备确需完整替换时用 Hub 仓库/本地目录 trust_remote_codeTrue发布前核对 README 是否说明了新参数、输出差异与适用模型范围依赖锁定是否写入custom_generate/requirements.txt。实现与测试证据集中在 src/transformers/generation/utils.pygenerate入口、load_custom_generate、模式映射与废弃策略迁移逻辑和 tests/generation/test_utils.pycustom_generate的参数注入、模型仓库覆盖、依赖检查、trust_remote_code强制要求、本地目录与 Callable 等用例。深入理解常见解码策略的数学细节可参考官方博客 “How to generate text: using different decoding methods for language generation with 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创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表