首页
/ Transformers 模型输出体系全解:从 ModelOutput 基类到各任务通用输出类

Transformers 模型输出体系全解:从 ModelOutput 基类到各任务通用输出类

2026-09-09 20:10:29作者:毕习沙Eudora

导读

在 🤗 Transformers 中,无论你运行的是 BERT 文本分类、GPT 式自回归生成,还是语音/视觉/多模态模型,前向传播返回的结果都遵循同一套设计:每个模型返回的都是 ModelOutput 子类的实例——一种同时具备"属性访问、元组解包、字典查询"三种用法的数据容器。本文以 docs/source/ko/main_classes/output.md(对应英文版 docs/source/en/main_classes/output.md)为主线,结合源码 src/transformers/utils/generic.pysrc/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=Trueoutput_attentions=True,所以 hidden_statesattentions 不存在(为 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 的字段,并按字段定义顺序排列。本例中只有 losslogits 两个非空元素,因此:

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,但约定第一个字段(如 losslogits)承担主字段职责。这一约束保证了"非空字段即有效元素"的元组语义始终成立。

4.2 只读的"受限字典"

为了保证内部状态一致性,ModelOutput 显式禁用了若干可变字典操作,调用会直接抛异常(generic.py):

  • __delitem__del outputs["x"]
  • setdefault
  • pop
  • update

这一点在测试 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_nodegeneric.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_statehidden_statesattentions
BaseModelOutputWithPooling源码 BaseModelOutput 基础上增加 pooler_output(BERT 家族中为经过线性层与 tanh 处理的分类 token 表示)
BaseModelOutputWithCrossAttentions源码 增加 cross_attentions(解码器交叉注意力权重,需 add_cross_attention=True
BaseModelOutputWithPoolingAndCrossAttentions源码 上述两者合并
BaseModelOutputWithPast源码 增加 past_key_valuesCache 实例,加速自回归解码,受 use_cache=True 控制)
BaseModelOutputWithPastAndCrossAttentions源码 上述两者合并

需要留意:past_key_values 的类型在纯解码器模型中是 Cache,在编码器-解码器(seq2seq)模型中则是 EncoderDecoderCache。使用 past_key_values 时,last_hidden_state 通常只输出最后一个位置(形状 (batch_size, 1, hidden_size))。

5.2 语言模型输出

核心字段
CausalLMOutput源码 losslogitshidden_statesattentions
CausalLMOutputWithPast源码 增加 past_key_values
CausalLMOutputWithCrossAttentions源码 增加 cross_attentionspast_key_values
MaskedLMOutput源码 losslogitshidden_statesattentions
Seq2SeqLMOutput源码 losslogitspast_key_valuesdecoder_hidden_statesdecoder_attentionscross_attentionsencoder_last_hidden_stateencoder_hidden_statesencoder_attentions
NextSentencePredictorOutput 下一句预测任务的 losslogitshidden_statesattentions

其中 logits 的形状统一为 (batch_size, sequence_length, config.vocab_size)loss 仅在传入 labels 时出现。

5.3 任务头输出(分类 / 选择 / 问答 / 分词标注)

核心字段
SequenceClassifierOutput源码 losslogitshidden_statesattentions
Seq2SeqSequenceClassifierOutput losslogitspast_key_valuesdecoder_hidden_statesdecoder_attentionscross_attentionsencoder_last_hidden_stateencoder_hidden_statesencoder_attentions
MultipleChoiceModelOutput源码 losslogitshidden_statesattentions
TokenClassifierOutput源码 losslogitshidden_statesattentions
QuestionAnsweringModelOutput源码 lossstart_logitsend_logitshidden_statesattentions
Seq2SeqQuestionAnsweringModelOutput lossstart_logitsend_logitspast_key_valuesdecoder_hidden_statesdecoder_attentionscross_attentionsencoder_last_hidden_stateencoder_hidden_statesencoder_attentions

说明:分类类输出中 logits 形状为 (batch_size, config.num_labels);问答类输出用 start_logits/end_logits 表示答案区间起止位置的得分。

5.4 视觉与语音输出

核心字段
Seq2SeqSpectrogramOutput源码 谱图生成任务,字段含 lossspectrogram 及编解码器各隐藏状态
SemanticSegmenterOutput源码 losslogitspred_maskhidden_statesattentions
ImageClassifierOutput源码 losslogitshidden_statesattentions
ImageClassifierOutputWithNoAttention源码 无注意力权重版本:losslogitshidden_states
DepthEstimatorOutput源码 losspredicted_depthhidden_statesattentions
Wav2Vec2BaseModelOutput源码 last_hidden_stateextract_featureshidden_statesattentions
XVectorOutput源码 说话人嵌入任务:logitsembeddingshidden_statesattentions

5.5 时间序列(Time Series)输出

核心字段
Seq2SeqTSModelOutput last_hidden_statepast_key_valuesdecoder_hidden_statesdecoder_attentionscross_attentionsencoder_last_hidden_stateencoder_hidden_statesencoder_attentionsstatic_features
Seq2SeqTSPredictionOutput 在上一类基础上增加 prediction_lossprediction_sampleprediction_distribution 等预测相关字段
SampleTSPredictionOutput sequencesprediction_samples 等采样结果字段

这些类服务于库内的时间序列预测模型家族(如 Autoformer、PatchTST、TimeSformer 等),字段反映了"编码器-解码器 + 静态特征 + 分布采样"的输出结构。

六、MoE 模型与多模态扩展(源码中的延伸)

除原文列出的类外,modeling_outputs.py 中还定义了一批供 MoE(Mixture of Experts)模型与多模态模型使用的输出类,理解它们有助于把握输出体系的扩展方式:

  • MoEModelOutputMoeModelOutputWithPastMoeCausalLMOutputWithPastMoEModelOutputWithPastAndCrossAttentions:在基础字段之上增加 router_probs / router_logits(每层 MoE 路由器的概率或 logits,形状 (batch_size, sequence_length, num_experts)),用于计算辅助损失(auxiliary loss)与 z-loss;
  • Seq2SeqMoEModelOutput:在 seq2seq 基础上同时携带 decoder_router_logitsencoder_router_logits

这些类同样遵循"所有字段默认 None、非空字段构成元组/字典"的统一语义,印证了输出体系的可扩展设计。

七、行为验证与测试依据

输出容器的全部关键行为都有仓库测试背书,见 tests/utils/test_model_output.py

  • 属性访问与缺失字段:未赋值的字段返回 None,访问未定义属性抛 AttributeErrorL38-L44);
  • 整数/切片索引:索引与切片始终作用于非 None 字段(L46-L57);
  • 字符串键访问:只对存在的键有效,缺失键抛 KeyErrorL59-L70);
  • 字典协议keys()/values()/items() 只覆盖非 None 字段;updatedelpopsetdefault 均被禁止(L72-L98);
  • 属性与键联动x.a = 10 之后 x["a"] 同样为 10,反之亦然(L100-L110);
  • 从字典/可迭代对象实例化ModelOutputTest({"a": 30, "b": 10})ModelOutputTest([("a", 30), ("b", 10)]) 均能正确展开为字段(L112-L120)。

八、实战建议与注意事项

  1. 判断字段是否存在用 is not None:由于未返回的字段统一为 None,可以安全地写 if outputs.attentions is not None: 做分支处理,而无需先查 hasattr
  2. 元组/字典语义都忽略 None:这意味着 len(outputs) 随模型返回内容动态变化——传入 output_hidden_states=True 会改变元组的长度,跨调用比较长度或按下标硬编码取值时务必小心。
  3. 需要纯元组时使用 to_tuple():这是官方推荐的解包前置步骤,也适用于 return_dict=False 场景下的返回值。
  4. 不要修改输出容器update/pop/setdefault/del 等操作被禁止,若需组合结果请构造新容器或直接使用其中的张量。
  5. hidden_states[-1] 不等于 last_hidden_state 是正常现象:部分模型对末层输出做额外归一化处理,比较两者时需查阅具体模型实现。
  6. 自定义输出类必须用 @dataclass 且字段默认值约束:如果你基于 ModelOutput 派生自定义输出(例如写新模型时),需保证除首个字段外其余字段默认值均为 None,否则会触发 ValueError

九、小结

ModelOutput 是 🤗 Transformers 统一返回约定的基石:它既是 dataclass(属性访问)、又是 OrderedDict(字典访问)、还支持元组语义(索引/切片/to_tuple),并通过"非 None 字段才可见"这一核心规则把三种身份无缝统一。配合 modeling_outputs.py 中按任务族划分的数十个通用输出类(Base 主干、LM、分类/问答/标注、视觉/语音、时间序列、MoE),你可以对任何模型的返回值做到"不看文档也能猜到结构"。当你在文档或源码中看到 SequenceClassifierOutputCausalLMOutputWithPastBaseModelOutputWithPooling 等名字时,它们都只是这套统一体系在不同任务下的具体化而已。

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

项目优选

收起
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
860
1.35 K
docsdocs
暂无描述
Markdown
899
5.83 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
925
1.85 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.84 K
1.02 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
533
601
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.03 K
525
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
395