首页
/ Transformers 生成实用工具全解析:generate() 输出、Logits 处理、停止条件与流式输出实战指南

Transformers 生成实用工具全解析:generate() 输出、Logits 处理、停止条件与流式输出实战指南

2026-09-09 12:09:57作者:齐冠琰

本篇技术指南以 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() 主循环与输出数据结构 GenerateDecoderOnlyOutputGenerateEncoderDecoderOutputGenerateBeamDecoderOnlyOutputGenerateBeamEncoderDecoderOutput
logits_process.py 修改/重塑语言模型头的预测分数 LogitsProcessorLogitsProcessorList 及 28+ 个具体实现
stopping_criteria.py 决定何时终止生成(除 EOS 之外) StoppingCriteriaStoppingCriteriaListMaxLengthCriteriaMaxTimeCriteria
streamers.py 逐 token 流式产出文本 TextStreamerTextIteratorStreamerAsyncTextIteratorStreamer

这四类组件恰好构成了解码流水线的四个可插拔环节:采词前(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=Trueoutput_attentions=True,所以 hidden_statesattentions 均为 None

2.2 属性、元组与字典三种访问方式

属性访问:像访问普通 Python 属性一样直接读取,未被返回的属性值为 None

generation_output.scores      # 所有生成步骤的 LM 头预测分数(元组)
generation_output.attentions  # None(未开启 output_attentions)

元组访问:仅保留非 None 的属性。本例中只有 sequencesscores 两个非空字段,因此:

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

① GenerateDecoderOnlyOutpututils.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未处理的原始分数;
  • attentionshidden_states:按层×按步的双层元组;
  • past_key_valuesuse_cache=True 时返回模型缓存(通常是 [~cache_utils.Cache] 实例,如 DynamicCache/StaticCache),用于加速解码。

② GenerateEncoderDecoderOutpututils.py#L206-L250)——encoder-decoder 模型 + 非 beam 方法。在 decoder-only 基础上额外包含 encoder_attentionsencoder_hidden_statesdecoder_attentionscross_attentionsdecoder_hidden_states 等编码器侧信息,sequences 形状为 (batch_size*num_return_sequences, sequence_length)

③ GenerateBeamDecoderOnlyOutpututils.py#L254-L294)——decoder-only + beam 方法。额外字段:

  • sequences_scores:形状 (batch_size*num_return_sequences,),每个生成序列的最终 beam 分数;
  • scoresbeam 转移分数(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 路径。

④ GenerateBeamEncoderDecoderOutpututils.py#L298-L351)——encoder-decoder + beam 方法,为 ② 与 ③ 的并集。

源码中还定义了便捷类型别名(utils.py#L354-L357):

GenerateNonBeamOutput = GenerateDecoderOnlyOutput | GenerateEncoderDecoderOutput
GenerateBeamOutput = GenerateBeamDecoderOnlyOutput | GenerateBeamEncoderDecoderOutput
GenerateOutput = GenerateNonBeamOutput | GenerateBeamOutput

generate() 会根据 GenerationModeGREEDY_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() 的字符串参数(temperaturetop_ktop_prepetition_penaltyno_repeat_ngram_sizemin_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_idsscores(元组,每个生成步一个张量)与可选 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,否则 scoresNone(见 stopping_criteria.py#L49-L53 的类注释与 tests/generation/test_stopping_criteria.py 中的验证用例)。

4.2 内置条件详解

MaxLengthCriteriastopping_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)

MaxTimeCriteriastopping_criteria.py#L93-L115):以墙钟时间限制生成,max_time 单位为秒;计时默认从实例化时刻开始(initial_timestamp = time.time()),也可显式传入 initial_timestamp 覆盖。

StopStringCriteriastopping_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)。

EosTokenCriteriastopping_criteria.py#L543-L592):当最近生成的 token 命中 eos_token_id(默认取 model.generation_config.eos_token_id)时停止,支持通过 new_token_length 一次检查最近多个 token。

ConfidenceCriteriastopping_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.Queueloop.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_processorstopping_criteria 均可显式注入覆盖默认行为(也可返回 Nonegenerate() 使用内置默认处理器列表);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.pystopping_criteria.pystreamers.py 及其测试(tests/generation/)深入调试与扩展。

登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.14 K
2.76 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
858
1.35 K
docsdocs
暂无描述
Markdown
899
5.82 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
923
1.85 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.83 K
1.02 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
532
596
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.03 K
524
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.37 K
1.46 K
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
548
393