ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

Transformers 文本生成核心类全面解析:GenerationConfig、GenerationMixin 与 ContinuousMixin 实战指南

Transformers 文本生成核心类全面解析:GenerationConfig、GenerationMixin 与 ContinuousMixin 实战指南 Transformers 文本生成核心类全面解析GenerationConfig、GenerationMixin 与 ContinuousMixin 实战指南【免费下载链接】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自动回归文本生成是当前所有生成式大模型推理与训练的枢纽。在本仓库中这一能力被收敛为少数几个设计精良的主类main classes承载全部生成参数与校验逻辑的GenerationConfig、被注入到每个可生成模型中的GenerationMixin以及面向连续批处理continuous batching加速的ContinuousMixin、ContinuousBatchingManager与调度器家族。本文以 docs/source/en/main_classes/text_generation.md 为骨架结合本仓库中src/transformers/generation/目录的真实实现源码为你梳理这些类的职责边界、参数语义、方法调用链与实战用法。读完本文你将能熟练地参数化驱动一次高质量文本生成、创建与保存自定义生成配置、深度定制 beam search / assisted decoding 等策略并初步掌握如何用连续批处理把多路请求跑在同一张 GPU 上。一、文本生成 API 的顶层架构从 generate 到三个关键类在 Transformers 中文本生成并不是某个具体模型内置的黑盒而是通过配置对象 Mixin 注入的方式为所有支持生成的模型统一提供的。理解这种设计是掌握整套 API 的第一步。model.generate()是唯一的对外入口其行为完全由GenerationConfig实例参数化——包括解码策略贪婪、采样、beam search、组束搜索、assisted decoding 等、生成长度约束、KV 缓存策略以及各种 logits 后处理算子。从本仓库的GenerationConfig类文档看同一个generate调用在参数组合下可以对应多种生成模式greedy decodingnum_beams1且do_sampleFalsemultinomial sampling多项式采样num_beams1且do_sampleTruebeam-search decoding束搜索num_beams1且do_sampleFalsebeam-search multinomial samplingnum_beams1且do_sampleTrueassisted decoding辅助解码向.generate()传入assistant_model或prompt_lookup_num_tokens而从类层次结构看本仓库有一个值得注意的实现细节GenerationMixin直接继承自ContinuousMixin见 src/transformers/generation/utils.py#L359而ContinuousMixin定义在连续批处理实现文件 src/transformers/generation/continuous_batching/continuous_api.py#L1118 中。这意味着任何继承了GenerationMixin的模型同时自动获得了连续批处理能力init_continuous_batching/continuous_batching_context_manager/generate_batch三个入口无需额外混入其他类。此外还有一批服务于连续批处理流水线的配套类负责请求排队的Scheduler、FIFOScheduler、PrefillFirstScheduler以及作为用户界面层存在的ContinuousBatchingManager。二、GenerationConfig生成行为的总控制台GenerationConfig是一个持有某次生成任务全部配置的类。它与模型自身的PreTrainedConfig相互独立通常以 JSON 形式保存在模型仓库或本地目录中源码注释中提到的generation_config.json即指模型侧的标准文件名见 configuration_utils.py 相关说明。实践中一个generate调用的参数优先级可总结为显式传入的generation_config实例 调用时传入的 kwargs 模型自带的generation_config.json 类内置默认值。2.1 核心参数速查按用途分组从GenerationConfig的类文档src/transformers/generation/configuration_utils.py#L126 起与字段默认实现出发可以将浩繁的参数归纳为以下几组便于记忆与检索分组关键参数作用说明输出长度控制max_length/min_length控制包括 prompt 在内的总序列长度为向后兼容保留官方推荐使用下面的*_new_tokens变体输出长度控制max_new_tokens/min_new_tokens只统计新生成的 token 数量忽略 prompt 长度是推荐的长度控制方式输出长度控制early_stopping仅对 beam 类方法有效True表示凑齐num_beams个完整候选即停止False采用启发式提前停止never表示跑完经典 beam search 的完整流程输出长度控制max_time允许计算运行的最大秒数到点后仍会完成当前这一步输出长度控制stop_strings字符串或字符串列表一旦模型输出到其中任一字符串即终止生成策略选择do_sample是否使用采样否则走贪婪解码策略选择num_beamsbeam search 的束数1表示不做束搜索num_beam_groups与之配合可实现多样性组束搜索策略选择use_mtp模型支持时是否启用 Multi-Token Prediction多 token 预测采样精细控制temperature调节下一个 token 概率分布的平滑程度默认1.0采样精细控制top_k/top_ptop-k 过滤保留概率最高的 k 个 token默认 50top-p 保留累计概率达到top_p的最小集合默认 1.0采样精细控制min_p/top_h/typical_p最小概率截断、熵预算缩放保留重归一化熵不超过top_h×全分布熵的最短前缀、局部典型性采样采样精细控制epsilon_cutoff/eta_cutoff截断采样epsilon 只保留条件概率大于阈值的 tokeneta 采样是局部典型性与 epsilon 采样的混合惩罚与去重repetition_penalty/encoder_repetition_penalty重复惩罚系数1.0表示不惩罚encoder 变体惩罚未出现在原输入中的序列惩罚与去重length_penalty/exponential_decay_length_penalty长度惩罚对 beam 结果按长度加权指数衰减形式的长度惩罚惩罚与去重no_repeat_ngram_size/encoder_no_repeat_ngram_size禁止出现指定大小的 n-gram 重复encoder 变体针对输入侧KV 缓存use_cache是否复用过去 key/value 注意力以加速解码KV 缓存cache_implementation缓存实现名称dynamicDynamicCache、static、offloaded、offloaded_static、quantized等KV 缓存cache_config/max_cache_len缓存类的参数字典max_cache_len仅用于静态缓存预分配固定长度以避免反复重分配与torch.compile重编译2.2 内置默认值_get_default_generation_params()一个容易忽略但很重要的实现细节是值为None的字段不会参与后续逻辑。GenerationConfig类文档明确指出凡保持None的字段会在生成循环中被_get_default_generation_params()返回的默认值覆盖如果想要不同的取值必须在 config 中显式设置。该方法位于 src/transformers/generation/configuration_utils.py#L601其默认参数化如下{ 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, bad_words_ids: None, num_return_sequences: 1, output_scores: False, return_dict_in_generate: False, forced_bos_token_id: None, forced_eos_token_id: None, remove_invalid_values: False, exponential_decay_length_penalty: None, suppress_tokens: None, begin_suppress_tokens: None, 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, # 已弃用参数迁移到 Hubnum_beam_groups / diversity_penalty num_beam_groups: 1, diversity_penalty: 0.0, }注意其中do_sample、temperature、top_k、top_p等字段默认是有值的不是None因此不会被覆盖。同时源码也给出提示num_assistant_tokens_schedule的默认调度为constant配合assistant_confidence_threshold0.4、num_assistant_tokens20等字段共同决定 assisted decoding 中草稿模型assistant model每轮建议的 token 数量与接受/拒绝判定。2.3 五个关键方法docs/source/en/main_classes/text_generation.md为该类标注了五个 autodoc 入口本仓库中的实现如下from_pretrained(pretrained_model_name_or_path, ...)从模型仓库、本地目录或 JSON 文件加载GenerationConfig。支持cache_dir、force_download、local_files_only、token、revision、subfolder等仓库加载参数说明它的加载链路与模型/分词器走的是同一套 Hub 读取基础设施。from_model_config(model_config)当模型没有独立的generation_config.json时从PreTrainedConfig或其 dict派生一份生成配置保证零配置也能 generate。save_pretrained(save_directory, ...)把当前配置保存到目录序列化为 JSON支持push_to_hubto_diff_dict()/to_json_string(use_diffTrue)等内部方法表明默认采用与内置默认值之差的精简 diff 形式落盘使文件更小、更易读。update(**kwargs)就地更新任意字段等价于config.max_new_tokens 64这种属性赋值的批量形式。validate(strictFalse, ...)从配置实例本身即可检测的错误参数化在此被拦截并抛异常而像生成长度这类依赖模型与其他输入才能判断的问题则推迟到generate运行时校验见 configuration_utils.py#L647 起的注释。get_generation_mode(assistant_modelNone)综合各字段把当前参数映射为具体GenerationMode是后续generate路由到_sample、_beam_search还是_assisted_decoding分支的判决书。一个完整的查看默认值 → 微调 → 落盘 → 复用流程示例from transformers import AutoModelForCausalLM, GenerationConfig model AutoModelForCausalLM.from_pretrained(你的模型ID) # 1) 查看模型默认生成配置 cfg model.generation_config # 等价于 GenerationConfig.from_model_config(model.config) # 2) 参数化一次生成两种等价方式 out_a model.generate(input_ids, generation_configcfg, max_new_tokens64) out_b model.generate(input_ids, do_sampleTrue, temperature0.8, top_p0.9) # kwargs 覆盖 # 3) 创建一份自定义配置并保存到本地 custom GenerationConfig.from_model_config(model.config) custom.max_new_tokens 128 custom.repetition_penalty 1.1 custom.save_pretrained(./my_gen_config/) # 生成 generation_config.json三、GenerationMixin注入每个模型的生成引擎GenerationMixin的类文档src/transformers/generation/utils.py#L359将其定位为包含自回归文本生成全部函数的 mixin供模型类混入使用。继承它意味着模型在初始化时具备加载GenerationConfig等生成相关行为并能调用generate等公开方法。3.1 何时该继承、何时不该继承文档给出了极具参考价值的三类模型界定LlamaForCausalLM这类纯因果解码器模型应直接继承GenerationMixin以获得generate与相关公开方法BlipForQuestionAnswering这类拥有自定义generate且接口与GenerationMixin.generate大致一致多几个参数、输出结构相同、内部又间接调用GenerationMixin.generate的模型也应继承以便享受代码库中全套生成相关的自动化机制BarkModel虽然内部某个子模型调用了GenerationMixin.generate但其对外generate接口与GenerationMixin.generate并不一致因此不应继承以免破坏generate的接口约定。由此可提炼出仓库的规则是否继承GenerationMixin取决于模型的对外generate是否与标准接口保持近似共享同一套签名与输出。3.2 generate 的完整签名与能力边界本仓库中generate的签名见 generation/utils.py为def generate( self, inputsNone, # 输入张量为 None 时可用 bos_token_id 初始化 generation_configNone, # GenerationConfig 实例优先级最高 logits_processorNone, # 自定义 LogitsProcessorList stopping_criteriaNone, # 自定义 StoppingCriteriaList prefix_allowed_tokens_fnNone, # 逐 token 约束函数 synced_gpusNone, # DeepSpeed/FSDP 多卡同步模式 assistant_modelNone, # 草稿模型 → 触发 assisted decoding streamerNone, # BaseStreamer配合 token 流式输出 negative_prompt_idsNone, # 无分类器引导CFG的负向 prompt negative_prompt_attention_maskNone, custom_generateNone, # 字符串或可调用对象替换为自定义生成方法 **kwargs, # 其余参数一律并入/更新 GenerationConfig ) - GenerateOutput | torch.LongTensor值得强调的设计是**kwargs与generation_config并存的参数解析机制所有多余的命名参数如max_new_tokens64会与传入的GenerationConfig合并未显式传入时则由_prepare_generation_config依据加载链自动补齐。与此同时generate在方法体内还完成了大量隐含的准备工作输入准备_prepare_model_inputs、_maybe_initialize_input_ids_for_generation、_prepare_attention_mask_for_generation等无输入时以bos_token_id播种组件装配_get_logits_processor把 config 参数编译成一串 logits 处理器、_get_stopping_criteria装配停止条件、_merge_criteria_processor_list用户自定义列表与默认列表合并、_get_candidate_generator为 assisted decoding 构造候选生成器缓存准备_prepare_cache_for_generation依据cache_implementation构造DynamicCache/StaticCache等可配置static缓存并借助max_cache_len预分配以复用torch.compile图长度校验_prepare_generated_length/_validate_generated_length会特别区分模型自带的默认 max_length如 Llama2 默认 4096与用户显式设置避免意外截断特殊 token 处理与 token 自愈_prepare_special_tokens、heal_tokens切分 token 后修复不完整 token。3.3 生成模式的分发与内部循环get_generation_mode判定模式后generate会把控制权交给对应的内部方法。从本仓库源码结构可以确认以下分派目标贪婪/采样_sample内部依据do_sample选择 argmax 还是多项式采样束搜索_beam_search并配有一系列 beam 维护辅助函数_gather_beams、_check_early_stop_heuristic、_update_finished_beams、_get_top_k_continuations等assisted decoding / 投机采样_assisted_decoding与顶层辅助函数_speculative_samplingprefill 与流水线优化_prefill、_split_model_outputs等。以 beam search 为例_check_early_stop_heuristic会把early_stopping的三种取值True/False/never落实到是否继续维护 running beam的具体判断上与上一节配置文档中的语义描述一一对应这也解释了为什么early_stopping是控制 beam 类方法停止条件的核心开关。3.4 输出结构与 compute_transition_scoresgenerate的返回类型被定义为联合类型GenerateOutput GenerateNonBeamOutput | GenerateBeamOutput细分包括GenerateDecoderOnlyOutput/GenerateEncoderDecoderOutput/GenerateBeamDecoderOnlyOutput/GenerateBeamEncoderDecoderOutput见 generation/utils.py。所有输出 dataclass 的公共字段为sequences当设置output_scoresTrue时还会携带sequences_scores、逐步scores、beam_indicesoutput_logitsTrue时携带未处理的逐 tokenlogits加上注意力/隐状态类字段用于可视化与调试。配套公开方法compute_transition_scores(sequences, scores, beam_indicesNone, normalize_logitsFalse)的作用是把generate(output_scoresTrue)返回的逐步得分还原为每个 token 的对数转移分数logits 经 log-softmax并给出与sequences对齐的逐位置分数便于在 NLL 评估或置信度分析中直接使用。3.5 token 流式输出streaminggenerate的streamer参数与BaseStreamer抽象类配合实现边生成边吐出已解码文本。本仓库在 src/transformers/generation/streamers.py 中提供了TextStreamer终端直接打印、TextIteratorStreamer提供__iter__/__next__适合在另一个线程中消费、AsyncTextIteratorStreamer异步版支持__aiter__/__anext__、以及用于 assisted decoding 草稿输出的AccelerateTextStreamer等实现。典型使用模式是把TextIteratorStreamer传入generate(streamer...)由独立线程执行生成、主线程持续迭代输出 token从而实现类 Chat 的实时流式效果。四、连续批处理ContinuousMixin 与 ContinuousBatchingManager如果generate解决的是单路请求如何更聪明地解码那么连续批处理解决的就是多路并发请求如何共享同一块 GPU 显存与计算资源。传统静态批处理需要等待所有请求凑齐后同时推进、以最慢者为准连续批处理则让每个请求在自己的节奏上完成 prefill 与 decode一旦某请求生成完毕立即让出缓存块给新请求。4.1 三个嵌套的入口ContinuousMixin的类文档continuous_api.py#L1118明确指出它有三级嵌套入口修改任意一层都应同步其余两层init_continuous_batching(generation_configNone, continuous_batching_configNone, workload_hintsNone)真正的底层入口负责初始化ContinuousBatchingManager并返回之continuous_batching_context_manager(...)围绕 manager 的上下文管理器封装支持block、timeout、persistent_manager、warmup等控制负责完整的生命周期generate_batch(inputs, generation_configNone, continuous_batching_configNone, ...)最高层的便捷函数内部包裹上述上下文管理器直接返回dict[str, GenerationOutput]请求 ID 到生成结果的映射。同时需要注意ContinuousBatchingManager不应被直接构造——源码注释明确要求只能通过上述三个 mixin 方法创建continuous_api.py#L574-L581。Manager 内部管理一个后台生成线程、输入/取消队列、输出路由器与批处理器并提供add_request/add_requests/get_result/request_id_iter/cancel_request/register_result_handler等用户界面方法。初始化 Manager 时switch_to_cb_friendly_attn会把模型的注意力实现切换为带paged|前缀的分页版本如paged|eager、paged|sdpa若检测到模型支持 flash attention 则优先切换到flash_attention_2/3因为代码中的告警信息明确提示连续批处理在 flash attention 下效果要好得多。同时模型会进入eval()模式。4.2 ContinuousBatchingConfig连续批处理专用配置ContinuousBatchingConfig是一个独立的 dataclass见 configuration_utils.py#L1656只负责 KV 缓存与批处理机制层面的参数与GenerationConfig关注的解码策略正交。关键字段字段默认值含义block_size256每个 KV 缓存块容纳的 token 数num_blocksNoneKV 缓存块总数为None时依据 GPU 显存自动推断max_batch_tokensNone单批最大 token 数同样可自动推断max_memory_percentNone用于 KV 缓存的空闲显存上限比例自动解析为 0.9无 logits 处理或 0.8有 logits 处理为词表大小的临时张量留出余量max_requests_per_batchNone单批最大请求数无 workload hints 时回退到 1024max_blocks_per_requestNone用于flash_attn_with_kvcache快速 decode 路径的块表定维设为 0 会禁用快速解码路径allow_block_sharingTrue是否允许块共享前缀缓存前提只能允许不能强制短 prompt 长生成的场景可考虑关闭use_async_batchingNone是否启用异步双缓冲消除连续批循环的 CPU 开销代价是显存翻倍None时自动检测use_cuda_graphNone是否启用 CUDA graphs可为二元组varlen 路径 / 快速解码路径None自动推断q_padding_interval_size/kv_padding_interval_size0CUDA graphs 的 query / KV 填充粒度token 数0 表示采用代码内预设varlen_compile_config/decode_compile_configNone两条执行路径varlen prefill / 静态 decode各自的torch.compile配置default_compile_level0默认编译级别0~3越高性能越好但 warmup 越久scheduler_typefifo调度器类型与下文的 Scheduler 家族对应safety_marginNone调度安全边际低于该空闲块比例即停止调度新的 prefillreturn_logprobsFalse是否随生成结果返回 log 概率seedNone采样种子None时随机cpu_offload_space0.0KV 缓存 CPU 交换空间GiB0 关闭 offload超额时按cpu_offload_space_safety_threshold默认 0.8×系统内存钳制max_queue_size0服务场景下的请求队列上限0 表示不限per_request_processors/drop_unsupported_processorsFalse/True是否允许每请求独立 logits 处理器参数如各自的 temperature对不支持的处理器的处置策略disable_nccl_graph_mixingTrue关闭 NCCL 的图混合安全网连续批场景不需要能带来 TP 性能提升cpu_group_timeout300.0CPU 通信超时秒4.3 实战Manager 生命周期与 generate_batch低层用法是手动管理 manager 生命周期适合对请求有精细控制的服务场景# manager 由 mixin 提供不要直接 new manager model.init_continuous_batching( continuous_batching_configContinuousBatchingConfig(block_size256, return_logprobsTrue) ) manager.warmup() manager.start() rid manager.add_request( input_idstokenized_prompt, request_idreq-001, max_new_tokens64, streamingTrue, # logit_processor_kwargs{temperature: 0.8} # per_request_processorsTrue 时生效 ) for output in manager.request_id_iter(rid): ... # 逐条消费流式 GenerationOutput manager.stop(blockTrue) # 或 hard_stopTrue 立即终止并 fail 所有待处理请求更高层的generate_batch则把启动、预热、停止全部折叠进一次调用适合离线批量推理results model.generate_batch( inputs[tokenized_a, tokenized_b, tokenized_c], max_new_tokens32, progress_barTrue, ) # - dict[str, GenerationOutput]4.4 底层机制速览从 src/transformers/generation/continuous_batching/ 目录的结构可以推断整套流水线由若干专注的组件拼装而成cache.py负责分页 KV 缓存与显存求解infer_max_batch_tokens_and_num_blocks通过激活峰值求解二元内存分配cache_manager.py提供块分配器与多种块管理策略scheduler.py负责请求排队与逐批挑选input_outputs.py负责批张量的搬运与 CUDA graph 缓冲model_runner.py执行批量前向与采样并支持torch.compile/CUDA graph 捕获offloading_manager.py管理 CPU offloaddistributed.py提供张量并行TP下的通信同步。配置解析集中在initialization.py的resolve_continuous_batching_config中完成。感兴趣的读者可以在仓库中找到两条通往真实运行的路径一键脚本 examples/pytorch/continuous_batching_simple.py 与完整示例 examples/pytorch/continuous_batching.py以及自动化测试 tests/generation/test_continuous_batching.py。五、调度器家族Scheduler、FIFOScheduler 与 PrefillFirstScheduler当多个请求同时处于不同阶段有的还在 prefill、有的在 decode、有的排队等待谁来决定下一批处理谁答案就是调度器。文档中的三个类构成抽象基类 两种策略的清晰结构全部位于 src/transformers/generation/continuous_batching/scheduler.py。Scheduler抽象基类定义请求的生命周期管理——从加入 waiting 队列、被schedule_batch挑中进入 active、到finish_request时释放缓存块。每个批次的挑选受两个预算约束token_budget本批最多处理的 token 数与cache_budget本批最多读取的 KV 缓存条目数。核心方法schedule_batch返回被调度请求列表、是否能走 decode 快速路径、总 query token 数与最大 KV 读取长度。它还管理请求取消set_request_cancellation/clear_cancelled_requests。FIFOScheduler默认调度器对应ContinuousBatchingConfig.scheduler_typefifo按照请求到达顺序处理且解码请求优先于 prefill 请求——先到先服务保证公平性decode 优先则保证已开始生成的低延迟不被新来的长 prompt 拖累。其默认安全边际为 0.1515% 的空闲块见 scheduler.py#L331。PrefillFirstScheduler与 FIFO 相反的策略优先处理被切分过的 prefill 请求即大 prompt 被分块前向的延续片段确保这些半成品先被完成再处理新的解码请求从而避免大量请求卡在部分 prefill 状态见 scheduler.py#L380。关于safety_margin基类注释给出了非常精确的口语化定义safety_margin0.1意味着当空闲块不足 10%即已分配超过 90%时停止调度新的 prefill 请求设为0.0表示完全不设边际——这是平衡抢占显存的新请求与保护正在解码的存量请求的关键旋钮。六、延伸阅读路径docs/source/en/main_classes/text_generation.md本文对应的 API 索引页包含各主类的 autodoc 声明docs/source/en/generation_strategies.md文本生成策略指南讲解如何检查模型默认生成配置、如何临时修改参数、如何创建并保存自定义配置以及 token 流式输出等关联特性src/transformers/generation/configuration_utils.pyGenerationConfig与ContinuousBatchingConfig的定义与校验实现src/transformers/generation/utils.pyGenerationMixin、generate主循环与内部解码分支、输出 dataclass 定义src/transformers/generation/logits_process.py 与 src/transformers/generation/stopping_criteria.py_get_logits_processor/_get_stopping_criteria装配的处理器组件实现src/transformers/generation/continuous_batching/连续批处理流水线的完整源码examples/pytorch/continuous_batching_simple.py 与 tests/generation/test_continuous_batching.py可运行示例与测试用例。掌握本文介绍的这条主线——GenerationConfig定义做什么、GenerationMixin.generate决定怎么做、ContinuousMixin与调度器回答多路并发时如何高效地一起做——你就拥有了阅读任何模型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创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表