Transformers 生成实用工具全解析:generate() 输出、Logits 处理、停止条件与流式输出实战指南
本篇技术指南以 docs/source/ja/internal/generation_utils.md 为核心,系统讲解 transformers 项目中 GenerationMixin.generate() 所依赖的全部内部实用工具:生成输出类型(ModelOutput 子类)、LogitsProcessor 系列(logits 处理与采样 warper)、StoppingCriteria 系列(停止条件) 与 Streamers(流式输出)。读完本文,你将能够精确解读 generate() 的返回对象、按需组合 logits 处理器控制解码行为、自定义生成终止策略,并实现逐 token 流式输出,这些能力可直接落地到文本生成、对话系统与推理服务等实际场景。
一、生成实用工具概览:generate() 背后的四大组件
在 transformers 中,所有生成式模型(GPT、Llama、T5、Whisper 等)的自回归解码统一由 [~generation.GenerationMixin.generate] 驱动,其完整实现位于 src/transformers/generation/utils.py。围绕它,仓库在 src/transformers/generation/ 目录下按职责拆分出四个子模块:
| 模块文件 | 职责 | 核心类 |
|---|---|---|
| utils.py | generate() 主循环与输出数据结构 |
GenerateDecoderOnlyOutput、GenerateEncoderDecoderOutput、GenerateBeamDecoderOnlyOutput、GenerateBeamEncoderDecoderOutput |
| logits_process.py | 修改/重塑语言模型头的预测分数 | LogitsProcessor、LogitsProcessorList 及 28+ 个具体实现 |
| stopping_criteria.py | 决定何时终止生成(除 EOS 之外) | StoppingCriteria、StoppingCriteriaList、MaxLengthCriteria、MaxTimeCriteria |
| streamers.py | 逐 token 流式产出文本 | TextStreamer、TextIteratorStreamer、AsyncTextIteratorStreamer |
这四类组件恰好构成了解码流水线的四个可插拔环节:采词前(logits 处理)→ 采词(输出收集)→ 停不停(停止条件)→ 怎么吐(流式)。下文依次深入每一环,并给出源码级证据。
二、理解 generate() 的输出:四种 ModelOutput 子类
2.1 输出对象与示例
[~generation.GenerationMixin.generate] 的返回值是 [~utils.ModelOutput] 的子类实例——一种携带全部生成信息的数据结构,同时可以像元组或字典一样使用。以下面这段经典调用为例:
from transformers import GPT2Tokenizer, GPT2LMHeadModel
tokenizer = GPT2Tokenizer.from_pretrained("openai-community/gpt2")
model = GPT2LMHeadModel.from_pretrained("openai-community/gpt2")
inputs = tokenizer("Hello, my dog is cute and ", return_tensors="pt")
generation_output = model.generate(**inputs, return_dict_in_generate=True, output_scores=True)
由于使用了解码器专用(decoder-only)模型且未启用 beam search,generation_output 的类型是 [~generation.GenerateDecoderOnlyOutput],其核心属性包括:
sequences:生成的 token 序列;scores(可选):每个生成步骤语言模型头的已处理预测分数(SoftMax 之前的原始 logits);hidden_states(可选):每个生成步骤的模型隐藏状态;attentions(可选):每个生成步骤的注意力权重。
在上例中,因为传入了 output_scores=True,所以 scores 有值;而没有传 output_hidden_states=True 或 output_attentions=True,所以 hidden_states 与 attentions 均为 None。
2.2 属性、元组与字典三种访问方式
属性访问:像访问普通 Python 属性一样直接读取,未被返回的属性值为 None:
generation_output.scores # 所有生成步骤的 LM 头预测分数(元组)
generation_output.attentions # None(未开启 output_attentions)
元组访问:仅保留非 None 的属性。本例中只有 sequences 与 scores 两个非空字段,因此:
generation_output[:2] # 等价于 (generation_output.sequences, generation_output.scores)
字典访问:同样只保留非 None 属性,例如 {"sequences": ..., "scores": ...} 两个键。
这种"三元一体"的设计来自 ModelOutput 基类(定义于 src/transformers/utils/generic.py),它继承自 OrderedDict 并实现了 __getitem__ / __getattr__ 双通道访问,保证向后兼容早期 generate() 直接返回张量或元组的 API 习惯。
2.3 四种输出类型与字段差异
根据"是否 beam search × 是否 encoder-decoder"两个维度,仓库在 src/transformers/generation/utils.py 中定义了四个 dataclass,全部继承 ModelOutput:
① GenerateDecoderOnlyOutput(utils.py#L170-L202)——decoder-only 模型 + 非 beam 方法:
sequences:形状(batch_size, sequence_length),第二维等于max_length,若因eos_token_id提前结束则更短;scores/logits:均为每个生成 token 一个张量的元组,形状(batch_size, config.vocab_size),区别在于scores是已处理(如经 warper 修改)的分数,logits是未处理的原始分数;attentions、hidden_states:按层×按步的双层元组;past_key_values:use_cache=True时返回模型缓存(通常是 [~cache_utils.Cache] 实例,如DynamicCache/StaticCache),用于加速解码。
② GenerateEncoderDecoderOutput(utils.py#L206-L250)——encoder-decoder 模型 + 非 beam 方法。在 decoder-only 基础上额外包含 encoder_attentions、encoder_hidden_states、decoder_attentions、cross_attentions、decoder_hidden_states 等编码器侧信息,sequences 形状为 (batch_size*num_return_sequences, sequence_length)。
③ GenerateBeamDecoderOnlyOutput(utils.py#L254-L294)——decoder-only + beam 方法。额外字段:
sequences_scores:形状(batch_size*num_return_sequences,),每个生成序列的最终 beam 分数;scores:beam 转移分数(beam transition scores),即基于该 beam 此前已生成 token 的 log-softmax 条件对数概率,每个张量形状(batch_size*num_beams, config.vocab_size);beam_indices:形状(batch_size*num_return_sequences, sequence_length),记录每个生成步所用 beam 的索引,可用于回溯整条 beam 路径。
④ GenerateBeamEncoderDecoderOutput(utils.py#L298-L351)——encoder-decoder + beam 方法,为 ② 与 ③ 的并集。
源码中还定义了便捷类型别名(utils.py#L354-L357):
GenerateNonBeamOutput = GenerateDecoderOnlyOutput | GenerateEncoderDecoderOutput
GenerateBeamOutput = GenerateBeamDecoderOnlyOutput | GenerateBeamEncoderDecoderOutput
GenerateOutput = GenerateNonBeamOutput | GenerateBeamOutput
generate() 会根据 GenerationMode(GREEDY_SEARCH/SAMPLE/BEAM_SEARCH/BEAM_SAMPLE 等,映射表见 utils.py#L138-L149)自动选择对应的输出类型,无需手动指定。注意:只有显式传入 return_dict_in_generate=True 时才会返回上述结构化对象;否则 generate() 默认只返回 sequences 张量。
三、LogitsProcessor:在采词前修改语言模型头的预测分数
3.1 基类与处理器列表
[LogitsProcessor](logits_process.py#L49-L60)是抽象基类,约定统一的调用协议:输入 input_ids(形状 (batch_size, sequence_length))与 scores(形状 (batch_size, config.vocab_size),beam search 下为 log-softmax 分数),输出处理后的同形状分数张量:
class LogitsProcessor:
def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor) -> torch.FloatTensor:
raise NotImplementedError(...)
[LogitsProcessorList](logits_process.py#L63-L98)继承自 list,其 __call__ 会按列表顺序依次调用每个处理器,天然支持链式组合(这也是 generate() 内部 _get_logits_processor 组装默认处理器列表的方式):
for processor in self:
scores = processor(input_ids, scores, **kwargs) # 前一个的输出作为后一个的输入
return scores
3.2 完整处理器清单
以下为文档列出的全部 LogitsProcessor / Warper(它们均定义于 src/transformers/generation/logits_process.py,并有对应单元测试 tests/generation/test_logits_process.py):
采样 Warper(改变概率分布形状):
- [
TemperatureLogitsWarper]:温度缩放。temperature大于 1 增加随机性、小于 1 趋于贪心,需配合do_sample=True才生效; - [
TopKLogitsWarper]:仅保留分数最高的 K 个 token,其余置-inf; - [
TopPLogitsWarper](nucleus 采样):保留累计概率达到top_p的最小 token 集合; - [
TypicalLogitsWarper]:typical sampling,保留与条件熵最接近的 token; - [
EpsilonLogitsWarper]:按绝对概率阈值epsilon截断采样空间; - [
EtaLogitsWarper]:eta截断(min_tokens_to_keep保证最小保留数); - [
LogitNormalization]:对 logits 做 log-softmax 归一化。
长度与 token 约束(Processor):
- [
MinLengthLogitsProcessor]:在未达到min_length前,将 EOS token 分数置为-inf(注意 decoder-only 模型长度包含 prompt); - [
MinNewTokensLengthLogitsProcessor]:与上者相反,忽略 prompt,只约束新生成 token 数(min_new_tokens); - [
ForcedBOSTokenLogitsProcessor] / [ForcedEOSTokenLogitsProcessor]:强制首个/末尾 token 为指定 id; - [
SuppressTokensLogitsProcessor] / [SuppressTokensAtBeginLogitsProcessor]:抑制指定 token(如 Whisper 抑制非语言 token),后者仅作用于生成开始阶段; - [
NoRepeatNGramLogitsProcessor] / [EncoderNoRepeatNGramLogitsProcessor]:禁止出现 n-gram 重复,后者同时考虑编码器输入; - [
NoBadWordsLogitsProcessor]:禁止生成指定"坏词"序列; - [
PrefixConstrainedLogitsProcessor]:强制生成遵循给定前缀约束; - [
ExponentialDecayLengthPenalty]:对超过阈值后的长度施加指数衰减惩罚; - [
SequenceBiasLogitsProcessor]:对指定 token 序列施加正/负分数偏置; - [
RepetitionPenaltyLogitsProcessor] / [EncoderRepetitionPenaltyLogitsProcessor]:对已出现 token 的分数施加重复惩罚(如repetition_penalty参数)。
特殊场景处理器:
- [
InfNanRemoveLogitsProcessor]:过滤掉分数中的inf/nan; - [
ClassifierFreeGuidanceLogitsProcessor] / [UnbatchedClassifierFreeGuidanceLogitsProcessor]:无分类器引导(CFG),后者针对未分批输入优化; - [
WhisperTimeStampLogitsProcessor]:Whisper 时间戳 token 约束(强制单调递增、抑制不合法时间戳组合); - [
AlternatingCodebooksLogitsProcessor]:用于 EnCodec 类多码本(codebook)音频模型的交替解码。
实用建议:日常通过
generate()的字符串参数(temperature、top_k、top_p、repetition_penalty、no_repeat_ngram_size、min_new_tokens等)即可触发上述绝大多数处理器;只有需要自定义逻辑时才手动实例化并传入logits_processor=。
3.3 源码视角:三个代表性实现
以 [MinLengthLogitsProcessor](logits_process.py#L101-L161)为例,其核心逻辑是构造全词表索引,用 torch.isin 定位 EOS token 位置,再以 torch.where 将其分数替换为 -inf:
vocab_tensor = torch.arange(scores.shape[-1], device=scores.device)
eos_token_mask = torch.isin(vocab_tensor, self.eos_token_id)
if input_ids.shape[-1] < self.min_length:
scores_processed = torch.where(eos_token_mask, -math.inf, scores)
[MinNewTokensLengthLogitsProcessor](logits_process.py#L164-L235)则先计算 new_tokens_length = input_ids.shape[-1] - prompt_length_to_skip,再与 min_new_tokens 比较——两者恰好对应 generate(min_length=...) 与 generate(min_new_tokens=...) 的语义差异。注意源码中两者均标注 supports_continuous_batching = False,表明它们暂不支持 continuous batching 路径(该属性在 logits_process.py#L52-L54 定义,供连续批处理解码时做能力探测)。
此外,LogitsProcessorList.__call__ 通过 inspect.signature 检查每个处理器的参数数量,自动判断是否需要注入额外 kwargs,这使得自定义处理器可以声明额外参数而无需改动调用链(logits_process.py#L86-L96)。
四、StoppingCriteria:精确控制生成何时停止
4.1 基类与列表
[StoppingCriteria](stopping_criteria.py#L48-L59)是抽象基类(继承 ABC),注意文档明确说明:停止条件仅在 PyTorch 实现中可用。其调用协议为:输入 input_ids、scores(元组,每个生成步一个张量)与可选 kwargs,返回形状 (batch_size,) 的布尔张量,True 表示该样本应停止生成:
class StoppingCriteria(ABC):
def __call__(self, input_ids, scores=None, **kwargs) -> torch.BoolTensor:
raise NotImplementedError("StoppingCriteria needs to be subclassed")
[StoppingCriteriaList] 同样继承 list,__call__ 会对所有条件取或(任一满足即停止)。若自定义条件依赖 scores,必须同时给 generate() 传 return_dict_in_generate=True, output_scores=True,否则 scores 为 None(见 stopping_criteria.py#L49-L53 的类注释与 tests/generation/test_stopping_criteria.py 中的验证用例)。
4.2 内置条件详解
MaxLengthCriteria(stopping_criteria.py#L62-L90):当完整生成长度(decoder-only 模型含 prompt)超过 max_length 时停止;若同时传入 max_position_embeddings,超出模型预设最大长度时会发出 warning_once 提示:
cur_len = input_ids.shape[1]
is_done = cur_len >= self.max_length
return torch.full((input_ids.shape[0],), is_done, device=input_ids.device, dtype=torch.bool)
MaxTimeCriteria(stopping_criteria.py#L93-L115):以墙钟时间限制生成,max_time 单位为秒;计时默认从实例化时刻开始(initial_timestamp = time.time()),也可显式传入 initial_timestamp 覆盖。
StopStringCriteria(stopping_criteria.py#L118-L540):支持 stop_strings 字符串终止。其实现非常巧妙:为了让匹配过程可被 Torch/XLA 编译,它不依赖字符串操作,而是预计算整个词表中每个 token 与停止字符串的匹配位置(end-overlap、内部匹配位置、token 长度),打包成 embedding 张量,运行时只用 F.embedding + cumsum 等纯张量运算完成匹配(stopping_criteria.py#L429-L479)。匹配按"从序列尾部向前回溯"进行,因此 ["st", "opera"] 这类 token 边界跨越停止字符串的情况也能被正确捕获,而停止字符串不在末尾的情形(如 ["stop", "at"])不会误触发。模块级缓存 STOP_STRING_EMBEDDING_CACHE(LRU,上限 8 项)避免了重复计算(stopping_criteria.py#L17-L19)。
EosTokenCriteria(stopping_criteria.py#L543-L592):当最近生成的 token 命中 eos_token_id(默认取 model.generation_config.eos_token_id)时停止,支持通过 new_token_length 一次检查最近多个 token。
ConfidenceCriteria(stopping_criteria.py#L595):用于投机解码(assisted generation),当 assistant 模型对当前 token 的置信度低于 assistant_confidence_threshold 时提前结束本批投机步。
实用建议:
generate(max_new_tokens=..., max_time=..., stop_strings=["<|end|>"], tokenizer=tokenizer)即可同时组合上述条件;StoppingCriteriaList的或语义保证了"长度、时间、字符串"任一命中都会终止解码。
五、Streamers:把生成结果变成实时文本流
5.1 BaseStreamer 与 TextStreamer
所有流式类的基类是 [BaseStreamer](streamers.py#L28-L39),仅约定两个钩子:put(value)(generate() 每生成一批 token 时调用)与 end()(生成结束时调用)。解码主循环正是通过 StopCheck/DeferredStopCheck 在每步调用 streamer.put(tokens.cpu()) 完成 token 投递(utils.py#L371-L386)。
[TextStreamer](streamers.py#L42-L154)是一个简单的文本流式输出器:每当形成完整单词时立即把文本打印到 stdout。其打印策略(streamers.py#L93-L112)包含三条启发式规则:
- 文本以
\n结尾 → 立即 flush 并重置缓存; - 最后一个字符是 CJK 字符(中日韩统一表意文字区,见
_is_chinese_char)→ 立即打印该字符; - 否则打印到最后一个空格为止,避免输出半截单词(因为后续 token 可能改变单词拼写)。
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)
# 除了返回常规输出,还会把生成文本实时打印到 stdout
_ = model.generate(**inputs, streamer=streamer, max_new_tokens=20)
构造参数:skip_prompt(默认 False,设为 True 可跳过 prompt 的打印,适合聊天机器人场景)与 decode_kwargs(透传给 tokenizer 的 decode 方法)。注意 TextStreamer 仅支持 batch size 为 1,否则抛出 ValueError。
5.2 TextIteratorStreamer:非阻塞的迭代式流
[TextIteratorStreamer](streamers.py#L157-L223)继承 TextStreamer,把"打印"替换为"入队":内部使用 queue.Queue,重写 on_finalized_text 将成文文本放入队列,生成结束时放入 stop_signal;对外实现 __iter__ / __next__,让下游应用以迭代器方式非阻塞地消费文本。典型用法是把 generate() 放到独立线程中执行:
from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer
from threading import Thread
tok = AutoTokenizer.from_pretrained("openai-community/gpt2")
model = AutoModelForCausalLM.from_pretrained("openai-community/gpt2")
inputs = tok(["An increasing sequence: one,"], return_tensors="pt")
streamer = TextIteratorStreamer(tok)
generation_kwargs = dict(inputs, streamer=streamer, max_new_tokens=20)
thread = Thread(target=model.generate, kwargs=generation_kwargs)
thread.start()
generated_text = ""
for new_text in streamer:
generated_text += new_text
timeout 参数控制队列读写阻塞时间,若 generate() 在线程中异常退出,主线程也不会无限期挂起——这正是交互式 Gradio 等 UI 场景的标准接入方式。
5.3 仓库中的其他流式变体
- [
AsyncTextIteratorStreamer](streamers.py#L226-L311):异步迭代版本,基于asyncio.Queue与loop.call_soon_threadsafe实现线程安全投递,供async for消费;必须在协程内初始化(因为它需要asyncio.get_running_loop())。 - [
TextDiffusionStreamer](streamers.py#L314-L406):面向文本扩散模型(如 DiffusionGemma),支持put_draft()以黄色打印中间草稿并覆盖上一次草稿,确认文本到来后再固化输出。
以上类均有对应测试覆盖于 tests/generation/test_streamers.py,主循环整体行为则由 tests/generation/test_utils.py 验证。
六、实战组合:一次完整的可观测生成调用
综合前文,一个同时展示"结构化输出 + 自定义 logits 处理 + 多停止条件 + 流式输出"的完整示例:
import torch
from transformers import (
AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer,
LogitsProcessorList, MinNewTokensLengthLogitsProcessor,
StoppingCriteriaList, MaxLengthCriteria, MaxTimeCriteria,
)
from threading import Thread
tok = AutoTokenizer.from_pretrained("openai-community/gpt2")
model = AutoModelForCausalLM.from_pretrained("openai-community/gpt2")
inputs = tok("Hello, my dog is cute and", return_tensors="pt")
# 1) 链式 logits 处理器:最少生成 5 个新 token
logits_processor = LogitsProcessorList([
MinNewTokensLengthLogitsProcessor(
prompt_length_to_skip=inputs["input_ids"].shape[-1],
min_new_tokens=5,
eos_token_id=model.config.eos_token_id,
),
])
# 2) 多停止条件:长度 + 时间
stopping_criteria = StoppingCriteriaList([
MaxLengthCriteria(max_length=50),
MaxTimeCriteria(max_time=60.0),
])
# 3) 流式输出到迭代器
streamer = TextIteratorStreamer(tok, skip_prompt=True)
generation_kwargs = dict(
inputs,
max_new_tokens=30,
do_sample=True,
temperature=0.8,
top_p=0.9,
return_dict_in_generate=True,
output_scores=True,
logits_processor=logits_processor,
stopping_criteria=stopping_criteria,
streamer=streamer,
)
thread = Thread(target=model.generate, kwargs=generation_kwargs)
thread.start()
for new_text in streamer:
print(new_text, end="", flush=True)
要点回顾:logits_processor 与 stopping_criteria 均可显式注入覆盖默认行为(也可返回 None 让 generate() 使用内置默认处理器列表);output_scores=True 使停止条件能读取分数;return_dict_in_generate=True 让返回值为结构化 ModelOutput 子类,可按属性/元组/字典三种方式消费。
七、总结
transformers 的生成实用工具层由四个正交组件构成:输出类型(4 个 ModelOutput dataclass,按 beam × encoder-decoder 划分)、LogitsProcessor(采样 warper 与约束处理器,支持链式组合)、StoppingCriteria(长度、时间、字符串、EOS、置信度等停止条件,取或合并)、Streamers(同步打印、线程安全迭代、异步迭代与扩散草稿四种流式形态)。理解这套组件,即可在 GenerationMixin.generate 之上自由定制解码策略,并借助 logits_process.py、stopping_criteria.py、streamers.py 及其测试(tests/generation/)深入调试与扩展。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0631
MiniCPM5-2BMiniCPM5-2B 是一款面向端侧、本地部署和资源受限场景的 2B 稠密 Transformer,能够达到同尺寸开源模型 SOTA 水平。Markdown00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
HivisionIDPhotos⚡️HivisionIDPhotos: a lightweight and efficient AI ID photos tools. 一个轻量级的AI证件照制作算法。Python09
DragonOSDragonOS is an operating system developed from scratch using Rust, with Linux compatibility. It is designed for **Serverless** scenarios. 使用Rust从0自研内核,具有Linux兼容性的操作系统,面向云计算Serverless场景而设计。Rust00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00