Transformers 模型输出体系全解:从 ModelOutput 基类到各任务通用输出类
导读
在 🤗 Transformers 中,无论你运行的是 BERT 文本分类、GPT 式自回归生成,还是语音/视觉/多模态模型,前向传播返回的结果都遵循同一套设计:每个模型返回的都是 ModelOutput 子类的实例——一种同时具备"属性访问、元组解包、字典查询"三种用法的数据容器。本文以 docs/source/ko/main_classes/output.md(对应英文版 docs/source/en/main_classes/output.md)为主线,结合源码 src/transformers/utils/generic.py 与 src/transformers/modeling_outputs.py 的实现细节,系统讲解输出对象的三种访问方式、None 字段的语义、各任务通用输出类的字段构成,以及背后的 dataclass 约束与实现原理,让你能熟练读懂并驾驭任何模型的返回值。
一、一切模型输出的共同基类:ModelOutput
所有模型(文本、视觉、音频、多模态)的输出都是 [~utils.ModelOutput] 子类的实例。这些输出是包含模型返回全部信息的数据结构,同时也可以被当作 tuple 或 dict 使用。这是整个库的统一约定:无论你调用哪个模型,返回值的"形状"都是可预测的。
源码层面,ModelOutput 定义于 src/transformers/utils/generic.py,其直接继承自 Python 标准库的 OrderedDict,因此天然具备字典的键值存储能力,同时通过自定义的 __getitem__ 提供了元组式索引能力。
二、快速上手:一个 BERT 序列分类示例
原文文档给出了最直观的入门示例,这里完整保留并补充注释:
from transformers import BertTokenizer, BertForSequenceClassification
import torch
tokenizer = BertTokenizer.from_pretrained("google-bert/bert-base-uncased")
model = BertForSequenceClassification.from_pretrained("google-bert/bert-base-uncased")
inputs = tokenizer("Hello, my dog is cute", return_tensors="pt")
labels = torch.tensor([1]).unsqueeze(0) # 批大小 1
outputs = model(**inputs, labels=labels)
运行后,outputs 对象是 [~modeling_outputs.SequenceClassifierOutput] 的实例。从该类文档可以看到它拥有 4 个字段:
| 字段 | 类型 | 说明 |
|---|---|---|
loss(可选) |
torch.FloatTensor |
分类(或 num_labels=1 时的回归)损失 |
logits |
torch.FloatTensor,形状 (batch_size, config.num_labels) |
分类得分(SoftMax 之前) |
hidden_states(可选) |
tuple(torch.FloatTensor) |
每层输出及初始嵌入输出 |
attentions(可选) |
tuple(torch.FloatTensor) |
每层注意力权重(softmax 之后) |
本例中,因为我们传入了 labels,所以 outputs.loss 存在;而我们没有传 output_hidden_states=True 或 output_attentions=True,所以 hidden_states 与 attentions 不存在(为 None)。
2.1 一个常见误区:hidden_states[-1] 与 last_hidden_state
文档中的 Tip 非常值得注意:当传入 output_hidden_states=True 时,你可能会预期 outputs.hidden_states[-1] 与 outputs.last_hidden_state 完全一致,但事实并非总是如此。部分模型在返回最后一层隐藏状态时会额外施加归一化(normalization)或后续处理。因此在依赖这两个值做等值比较或拼接时,务必核对具体模型的实现,而不是想当然地认为它们等价。
三、三种访问方式:属性、元组、字典
3.1 属性访问
你可以像访问普通 Python 对象属性一样访问每个字段;如果模型没有返回该字段,取到的值就是 None:
outputs.loss # 模型计算出的损失(本例存在)
outputs.attentions # None(因为没传 output_attentions=True)
3.2 当作 tuple 使用
把 outputs 当作元组时,只考虑值为非 None 的字段,并按字段定义顺序排列。本例中只有 loss 和 logits 两个非空元素,因此:
outputs[:2] # 等价于 (outputs.loss, outputs.logits)
更一般地,整数索引、切片都可用:outputs[0] 返回第一个非空字段,outputs[1:] 返回剩余非空字段组成的元组。这一行为在 tests/utils/test_model_output.py 中有专门测试:当 b 为空而 c 非空时,x[:2] 依然返回 (a, c),可见索引始终跳过 None。
3.3 当作字典使用
把 outputs 当作字典时,同样只保留非 None 的字段作为键。本例中有两个键:"loss" 与 "logits":
outputs["loss"] # 键访问
list(outputs.keys()) # ["loss", "logits"]
list(outputs.values()) # [损失张量, logits 张量]
3.4 重要限制:不能直接解包
ModelOutput 的类文档有一个醒目的警告:不能直接对 ModelOutput 实例做 *outputs 解包。如需解包,必须先调用 to_tuple() 方法:
loss, logits = outputs.to_tuple() # 正确做法
# loss, logits = outputs # 错误:ModelOutput 不可直接解包
四、源码级解析:ModelOutput 的实现原理
4.1 基于 OrderedDict 与 dataclass 的双重身份
ModelOutput 的类签名是 class ModelOutput(OrderedDict),同时其子类必须使用 @dataclass 装饰器定义。构造函数中(generic.py)会做强制校验:如果子类不是 dataclass,直接抛出 TypeError。
__post_init__(generic.py)执行两项一致性检查:
- 至少一个字段:dataclass 不能没有任何字段,否则抛
ValueError; - 最多一个必填字段:除第一个字段外,其余字段默认值必须为
None(即只能有一个必填字段)。例如SequenceClassifierOutput中只有loss的第一个字段loss默认None……实际上所有字段默认都是None,但约定第一个字段(如loss或logits)承担主字段职责。这一约束保证了"非空字段即有效元素"的元组语义始终成立。
4.2 只读的"受限字典"
为了保证内部状态一致性,ModelOutput 显式禁用了若干可变字典操作,调用会直接抛异常(generic.py):
__delitem__(del outputs["x"])setdefaultpopupdate
这一点在测试 tests/utils/test_model_output.py 中逐一验证。而赋值(outputs.loss = ... 或 outputs["loss"] = ...)是允许的,且属性与键会保持同步(见 generic.py 的 __setattr__/__setitem__ 联动)。
4.3 to_tuple():非 None 字段的元组化
def to_tuple(self) -> tuple:
"""将自身转换为包含所有非 None 属性/键的元组。"""
return tuple(self[k] for k in self.keys())
由于内部字典只存非 None 字段,to_tuple() 天然跳过 None,与"元组化时忽略 None"的语义完全一致(generic.py)。
4.4 与 Torch 生态的集成:pytree 注册
ModelOutput 的 __init_subclass__ 与 __init__ 都会调用 _register_model_output_pytree_node(generic.py),将每个输出类注册为 torch.utils._pytree 的节点。这在 torch.nn.parallel.DistributedDataParallel 配合 static_graph=True、以及 TorchDynamo 编译场景下,保证输出对象中的梯度能够正确同步与追踪。从源码结构看,这是库为分布式训练与编译优化预留的底层能力。
4.5 return_dict=False:退化为纯元组
[~utils.can_return_tuple] 装饰器(generic.py)是所有模型 forward 方法外围的关键逻辑:当 return_dict=False 被传入(或 config.return_dict=False)时,输出对象会被转换为 output.to_tuple(),从而返回一个普通的 Python 元组;反之则返回 ModelOutput 实例。这也是为什么你会看到模型 forward 签名里普遍带 return_dict 参数——它直接决定返回值形态。
五、通用模型输出类全景
原文文档指出:以下输出类被多种模型类型共用,因此单独成文归档;某个模型特有的输出类型则记录在对应模型页面。它们全部定义于 src/transformers/modeling_outputs.py,以下按族类整理(字段均可选,默认 None,除特别标注外与模型类型一一对应)。
5.1 基础主干输出(Base 系列)
| 类 | 核心字段 |
|---|---|
BaseModelOutput(源码) |
last_hidden_state、hidden_states、attentions |
BaseModelOutputWithPooling(源码) |
在 BaseModelOutput 基础上增加 pooler_output(BERT 家族中为经过线性层与 tanh 处理的分类 token 表示) |
BaseModelOutputWithCrossAttentions(源码) |
增加 cross_attentions(解码器交叉注意力权重,需 add_cross_attention=True) |
BaseModelOutputWithPoolingAndCrossAttentions(源码) |
上述两者合并 |
BaseModelOutputWithPast(源码) |
增加 past_key_values(Cache 实例,加速自回归解码,受 use_cache=True 控制) |
BaseModelOutputWithPastAndCrossAttentions(源码) |
上述两者合并 |
需要留意:past_key_values 的类型在纯解码器模型中是 Cache,在编码器-解码器(seq2seq)模型中则是 EncoderDecoderCache。使用 past_key_values 时,last_hidden_state 通常只输出最后一个位置(形状 (batch_size, 1, hidden_size))。
5.2 语言模型输出
| 类 | 核心字段 |
|---|---|
CausalLMOutput(源码) |
loss、logits、hidden_states、attentions |
CausalLMOutputWithPast(源码) |
增加 past_key_values |
CausalLMOutputWithCrossAttentions(源码) |
增加 cross_attentions 与 past_key_values |
MaskedLMOutput(源码) |
loss、logits、hidden_states、attentions |
Seq2SeqLMOutput(源码) |
loss、logits、past_key_values、decoder_hidden_states、decoder_attentions、cross_attentions、encoder_last_hidden_state、encoder_hidden_states、encoder_attentions |
NextSentencePredictorOutput |
下一句预测任务的 loss、logits、hidden_states、attentions |
其中 logits 的形状统一为 (batch_size, sequence_length, config.vocab_size);loss 仅在传入 labels 时出现。
5.3 任务头输出(分类 / 选择 / 问答 / 分词标注)
| 类 | 核心字段 |
|---|---|
SequenceClassifierOutput(源码) |
loss、logits、hidden_states、attentions |
Seq2SeqSequenceClassifierOutput |
loss、logits、past_key_values、decoder_hidden_states、decoder_attentions、cross_attentions、encoder_last_hidden_state、encoder_hidden_states、encoder_attentions |
MultipleChoiceModelOutput(源码) |
loss、logits、hidden_states、attentions |
TokenClassifierOutput(源码) |
loss、logits、hidden_states、attentions |
QuestionAnsweringModelOutput(源码) |
loss、start_logits、end_logits、hidden_states、attentions |
Seq2SeqQuestionAnsweringModelOutput |
loss、start_logits、end_logits、past_key_values、decoder_hidden_states、decoder_attentions、cross_attentions、encoder_last_hidden_state、encoder_hidden_states、encoder_attentions |
说明:分类类输出中 logits 形状为 (batch_size, config.num_labels);问答类输出用 start_logits/end_logits 表示答案区间起止位置的得分。
5.4 视觉与语音输出
| 类 | 核心字段 |
|---|---|
Seq2SeqSpectrogramOutput(源码) |
谱图生成任务,字段含 loss、spectrogram 及编解码器各隐藏状态 |
SemanticSegmenterOutput(源码) |
loss、logits、pred_mask、hidden_states、attentions |
ImageClassifierOutput(源码) |
loss、logits、hidden_states、attentions |
ImageClassifierOutputWithNoAttention(源码) |
无注意力权重版本:loss、logits、hidden_states |
DepthEstimatorOutput(源码) |
loss、predicted_depth、hidden_states、attentions |
Wav2Vec2BaseModelOutput(源码) |
last_hidden_state、extract_features、hidden_states、attentions |
XVectorOutput(源码) |
说话人嵌入任务:logits、embeddings、hidden_states、attentions |
5.5 时间序列(Time Series)输出
| 类 | 核心字段 |
|---|---|
Seq2SeqTSModelOutput |
last_hidden_state、past_key_values、decoder_hidden_states、decoder_attentions、cross_attentions、encoder_last_hidden_state、encoder_hidden_states、encoder_attentions、static_features |
Seq2SeqTSPredictionOutput |
在上一类基础上增加 prediction_loss、prediction_sample、prediction_distribution 等预测相关字段 |
SampleTSPredictionOutput |
sequences、prediction_samples 等采样结果字段 |
这些类服务于库内的时间序列预测模型家族(如 Autoformer、PatchTST、TimeSformer 等),字段反映了"编码器-解码器 + 静态特征 + 分布采样"的输出结构。
六、MoE 模型与多模态扩展(源码中的延伸)
除原文列出的类外,modeling_outputs.py 中还定义了一批供 MoE(Mixture of Experts)模型与多模态模型使用的输出类,理解它们有助于把握输出体系的扩展方式:
MoEModelOutput、MoeModelOutputWithPast、MoeCausalLMOutputWithPast、MoEModelOutputWithPastAndCrossAttentions:在基础字段之上增加router_probs/router_logits(每层 MoE 路由器的概率或 logits,形状(batch_size, sequence_length, num_experts)),用于计算辅助损失(auxiliary loss)与 z-loss;Seq2SeqMoEModelOutput:在 seq2seq 基础上同时携带decoder_router_logits与encoder_router_logits。
这些类同样遵循"所有字段默认 None、非空字段构成元组/字典"的统一语义,印证了输出体系的可扩展设计。
七、行为验证与测试依据
输出容器的全部关键行为都有仓库测试背书,见 tests/utils/test_model_output.py:
- 属性访问与缺失字段:未赋值的字段返回
None,访问未定义属性抛AttributeError(L38-L44); - 整数/切片索引:索引与切片始终作用于非
None字段(L46-L57); - 字符串键访问:只对存在的键有效,缺失键抛
KeyError(L59-L70); - 字典协议:
keys()/values()/items()只覆盖非None字段;update、del、pop、setdefault均被禁止(L72-L98); - 属性与键联动:
x.a = 10之后x["a"]同样为 10,反之亦然(L100-L110); - 从字典/可迭代对象实例化:
ModelOutputTest({"a": 30, "b": 10})与ModelOutputTest([("a", 30), ("b", 10)])均能正确展开为字段(L112-L120)。
八、实战建议与注意事项
- 判断字段是否存在用
is not None:由于未返回的字段统一为None,可以安全地写if outputs.attentions is not None:做分支处理,而无需先查hasattr。 - 元组/字典语义都忽略
None:这意味着len(outputs)随模型返回内容动态变化——传入output_hidden_states=True会改变元组的长度,跨调用比较长度或按下标硬编码取值时务必小心。 - 需要纯元组时使用
to_tuple():这是官方推荐的解包前置步骤,也适用于return_dict=False场景下的返回值。 - 不要修改输出容器:
update/pop/setdefault/del等操作被禁止,若需组合结果请构造新容器或直接使用其中的张量。 hidden_states[-1]不等于last_hidden_state是正常现象:部分模型对末层输出做额外归一化处理,比较两者时需查阅具体模型实现。- 自定义输出类必须用
@dataclass且字段默认值约束:如果你基于ModelOutput派生自定义输出(例如写新模型时),需保证除首个字段外其余字段默认值均为None,否则会触发ValueError。
九、小结
ModelOutput 是 🤗 Transformers 统一返回约定的基石:它既是 dataclass(属性访问)、又是 OrderedDict(字典访问)、还支持元组语义(索引/切片/to_tuple),并通过"非 None 字段才可见"这一核心规则把三种身份无缝统一。配合 modeling_outputs.py 中按任务族划分的数十个通用输出类(Base 主干、LM、分类/问答/标注、视觉/语音、时间序列、MoE),你可以对任何模型的返回值做到"不看文档也能猜到结构"。当你在文档或源码中看到 SequenceClassifierOutput、CausalLMOutputWithPast、BaseModelOutputWithPooling 等名字时,它们都只是这套统一体系在不同任务下的具体化而已。
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
video-shotcraftAI宣传片skill,使用 Remotion 制作电影级产品视频:提供106 张镜头配方卡和可复用的视频魔板。适用于 Claude Code 与 Codex以及所有其他智能体Markdown00
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