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 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
而从类层次结构看,本仓库有一个值得注意的实现细节: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_beams | beam search 的束数,1表示不做束搜索;num_beam_groups与之配合可实现多样性组束搜索 |
| 策略选择 | use_mtp | 模型支持时是否启用 Multi-Token Prediction(多 token 预测) |
| 采样精细控制 | temperature | 调节下一个 token 概率分布的平滑程度,默认1.0 |
| 采样精细控制 | top_k/top_p | top-k 过滤保留概率最高的 k 个 token(默认 50);top-p 保留累计概率达到top_p的最小集合(默认 1.0) |
| 采样精细控制 | min_p/top_h/typical_p | 最小概率截断、熵预算缩放(保留重归一化熵不超过top_h×全分布熵的最短前缀)、局部典型性采样 |
| 采样精细控制 | epsilon_cutoff/eta_cutoff | 截断采样:epsilon 只保留条件概率大于阈值的 token;eta 采样是局部典型性与 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 | 缓存实现名称:"dynamic"(DynamicCache)、"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, # 已弃用参数(迁移到 Hub):num_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_threshold=0.4、num_assistant_tokens=20等字段共同决定 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_hub;to_diff_dict()/to_json_string(use_diff=True)等内部方法表明默认采用"与内置默认值之差"的精简 diff 形式落盘,使文件更小、更易读。update(**kwargs):就地更新任意字段,等价于config.max_new_tokens = 64这种属性赋值的批量形式。validate(strict=False, ...):从配置实例本身即可检测的错误参数化在此被拦截并抛异常;而像生成长度这类依赖模型与其他输入才能判断的问题,则推迟到generate运行时校验(见 configuration_utils.py#L647 起的注释)。get_generation_mode(assistant_model=None):综合各字段把当前参数映射为具体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_config=cfg, max_new_tokens=64) out_b = model.generate(input_ids, do_sample=True, temperature=0.8, top_p=0.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, inputs=None, # 输入张量;为 None 时可用 bos_token_id 初始化 generation_config=None, # GenerationConfig 实例,优先级最高 logits_processor=None, # 自定义 LogitsProcessorList stopping_criteria=None, # 自定义 StoppingCriteriaList prefix_allowed_tokens_fn=None, # 逐 token 约束函数 synced_gpus=None, # DeepSpeed/FSDP 多卡同步模式 assistant_model=None, # 草稿模型 → 触发 assisted decoding streamer=None, # BaseStreamer,配合 token 流式输出 negative_prompt_ids=None, # 无分类器引导(CFG)的负向 prompt negative_prompt_attention_mask=None, custom_generate=None, # 字符串或可调用对象:替换为自定义生成方法 **kwargs, # 其余参数一律并入/更新 GenerationConfig ) -> GenerateOutput | torch.LongTensor值得强调的设计是**kwargs与generation_config并存的参数解析机制:所有多余的命名参数(如max_new_tokens=64)会与传入的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_sampling; - prefill 与流水线优化:
_prefill、_split_model_outputs等。
以 beam search 为例,_check_early_stop_heuristic会把early_stopping的三种取值(True/False/"never")落实到"是否继续维护 running beam"的具体判断上,与上一节配置文档中的语义描述一一对应,这也解释了为什么early_stopping是"控制 beam 类方法停止条件"的核心开关。
3.4 输出结构与 compute_transition_scores
generate的返回类型被定义为联合类型:GenerateOutput = GenerateNonBeamOutput | GenerateBeamOutput,细分包括GenerateDecoderOnlyOutput/GenerateEncoderDecoderOutput/GenerateBeamDecoderOnlyOutput/GenerateBeamEncoderDecoderOutput(见 generation/utils.py)。所有输出 dataclass 的公共字段为sequences;当设置output_scores=True时还会携带sequences_scores、逐步scores、beam_indices;output_logits=True时携带未处理的逐 tokenlogits;加上注意力/隐状态类字段用于可视化与调试。
配套公开方法compute_transition_scores(sequences, scores, beam_indices=None, normalize_logits=False)的作用是:把generate(output_scores=True)返回的逐步得分还原为每个 token 的对数转移分数(logits 经 log-softmax),并给出与sequences对齐的逐位置分数,便于在 NLL 评估或置信度分析中直接使用。
3.5 token 流式输出(streaming)
generate的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_config=None, continuous_batching_config=None, workload_hints=None):真正的底层入口,负责初始化ContinuousBatchingManager并返回之;continuous_batching_context_manager(...):围绕 manager 的上下文管理器封装(支持block、timeout、persistent_manager、warmup等控制),负责完整的生命周期;generate_batch(inputs, generation_config=None, continuous_batching_config=None, ...):最高层的便捷函数,内部包裹上述上下文管理器,直接返回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_size | 256 | 每个 KV 缓存块容纳的 token 数 |
num_blocks | None | KV 缓存块总数;为None时依据 GPU 显存自动推断 |
max_batch_tokens | None | 单批最大 token 数,同样可自动推断 |
max_memory_percent | None | 用于 KV 缓存的空闲显存上限比例;自动解析为 0.9(无 logits 处理)或 0.8(有 logits 处理),为词表大小的临时张量留出余量 |
max_requests_per_batch | None | 单批最大请求数,无 workload hints 时回退到 1024 |
max_blocks_per_request | None | 用于flash_attn_with_kvcache快速 decode 路径的块表定维;设为 0 会禁用快速解码路径 |
allow_block_sharing | True | 是否允许块共享(前缀缓存前提);只能"允许"不能"强制",短 prompt 长生成的场景可考虑关闭 |
use_async_batching | None | 是否启用异步双缓冲,消除连续批循环的 CPU 开销(代价是显存翻倍),None时自动检测 |
use_cuda_graph | None | 是否启用 CUDA graphs,可为二元组(varlen 路径 / 快速解码路径),None自动推断 |
q_padding_interval_size/kv_padding_interval_size | 0 | CUDA graphs 的 query / KV 填充粒度(token 数),0 表示采用代码内预设 |
varlen_compile_config/decode_compile_config | None | 两条执行路径(varlen prefill / 静态 decode)各自的torch.compile配置 |
default_compile_level | 0 | 默认编译级别(0~3),越高性能越好但 warmup 越久 |
scheduler_type | "fifo" | 调度器类型,与下文的 Scheduler 家族对应 |
safety_margin | None | 调度安全边际,低于该空闲块比例即停止调度新的 prefill |
return_logprobs | False | 是否随生成结果返回 log 概率 |
seed | None | 采样种子,None时随机 |
cpu_offload_space | 0.0 | KV 缓存 CPU 交换空间(GiB),0 关闭 offload;超额时按cpu_offload_space_safety_threshold(默认 0.8×系统内存)钳制 |
max_queue_size | 0 | 服务场景下的请求队列上限,0 表示不限 |
per_request_processors/drop_unsupported_processors | False/True | 是否允许每请求独立 logits 处理器参数(如各自的 temperature);对不支持的处理器的处置策略 |
disable_nccl_graph_mixing | True | 关闭 NCCL 的图混合安全网(连续批场景不需要,能带来 TP 性能提升) |
cpu_group_timeout | 300.0 | CPU 通信超时(秒) |
4.3 实战:Manager 生命周期与 generate_batch
低层用法是手动管理 manager 生命周期,适合对请求有精细控制的服务场景:
# manager 由 mixin 提供,不要直接 new manager = model.init_continuous_batching( continuous_batching_config=ContinuousBatchingConfig(block_size=256, return_logprobs=True) ) manager.warmup() manager.start() rid = manager.add_request( input_ids=tokenized_prompt, request_id="req-001", max_new_tokens=64, streaming=True, # logit_processor_kwargs={"temperature": 0.8} # per_request_processors=True 时生效 ) for output in manager.request_id_iter(rid): ... # 逐条消费流式 GenerationOutput manager.stop(block=True) # 或 hard_stop=True 立即终止并 fail 所有待处理请求更高层的generate_batch则把启动、预热、停止全部折叠进一次调用,适合离线批量推理:
results = model.generate_batch( inputs=[tokenized_a, tokenized_b, tokenized_c], max_new_tokens=32, progress_bar=True, ) # -> 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 offload,distributed.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_type="fifo"),按照请求到达顺序处理,且解码请求优先于 prefill 请求——先到先服务保证公平性,decode 优先则保证已开始生成的低延迟不被新来的长 prompt 拖累。其默认安全边际为 0.15(15% 的空闲块,见 scheduler.py#L331)。PrefillFirstScheduler:与 FIFO 相反的策略,优先处理"被切分过的 prefill 请求"(即大 prompt 被分块前向的延续片段),确保这些半成品先被完成,再处理新的解码请求,从而避免大量请求卡在部分 prefill 状态(见 scheduler.py#L380)。
关于safety_margin,基类注释给出了非常精确的口语化定义:safety_margin=0.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.py:
GenerationConfig与ContinuousBatchingConfig的定义与校验实现; - src/transformers/generation/utils.py:
GenerationMixin、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),仅供参考