HuggingFace中BaseModelOutput对象为何支持下标访问?
BaseModelOutput支持下标访问的实现原理
你的代码里BaseModelOutput对象能像列表一样用下标访问(比如encoder_outputs[0]),但打印时显示为自身类型,核心原因在于它的继承体系和类实现:
1. 继承自增强版字典类ModelOutput
BaseModelOutput是HuggingFace Transformers库中封装模型输出的基础类,它继承自ModelOutput。而ModelOutput本身继承自Python内置的dict,同时做了扩展:
- 实现了
__getitem__方法,既支持字典的键名访问(如encoder_outputs["last_hidden_state"]),也支持下标索引访问; - 约定了输出字段的默认顺序,对于
BaseModelOutput,下标对应的字段依次是:- 0:
last_hidden_state(也就是你代码里的sequence_output) - 1:
hidden_states(可选,需开启output_hidden_states=True) - 2:
attentions(可选,需开启output_attentions=True)
- 0:
2. 保留类类型的同时兼容字典特性
因为ModelOutput在继承dict的基础上,自定义了类的显示逻辑,所以打印时会显示BaseModelOutput的类类型,但底层依然具备字典的所有功能,包括下标访问、键值对遍历等。
举个对应实现的简化逻辑(模仿Transformers库的核心思路):
class ModelOutput(dict): def __getitem__(self, key): # 如果是整数下标,按预定义顺序返回对应字段 if isinstance(key, int): return list(self.values())[key] # 否则按字典键名返回 return super().__getitem__(key) class BaseModelOutput(ModelOutput): # 预定义字段顺序 _fields = ["last_hidden_state", "hidden_states", "attentions"]
这样,当你调用encoder_outputs[0]时,就会按顺序取出第一个字段last_hidden_state,而打印对象时,因为它是BaseModelOutput类的实例,所以会显示对应的类型信息。
内容的提问来源于stack exchange,提问作者Naren Dhyani
相关产品推荐
相关产品推荐

