Compare commits

...
Author SHA1 Message Date
Thomas Wolf acfcbc8833 fix #5693 2020-07-14 12:01:45 +02:00
11 changed files with 145 additions and 67 deletions
+82 -4
View File
@@ -14,6 +14,7 @@ import sys
import tarfile
import tempfile
from contextlib import contextmanager
from dataclasses import fields
from functools import partial, wraps
from hashlib import sha256
from pathlib import Path
@@ -856,12 +857,83 @@ def tf_required(func):
return wrapper
class ModelOutput:
"""
Base class for all model outputs as dataclass. Has a ``__getitem__`` that allows indexing by integer or slice (like
a tuple) or strings (like a dictionnary) that will ignore the ``None`` attributes.
class ModelOutput(dict):
""" Base class for all model outputs as dataclass.
The sub-classes of ``ModelOutput`` must have **at most one** required field (the first field)
(see ``__post_init__``docstring for details)
``ModelOutput`` has a ``__getitem__`` method that allows indexing by:
- integer and slice (like a tuple) or
- strings (like a dictionnary).
``__getitem__`` will ignore attributes containing ``None`` values when indexing with integers.
Apart from accepting integers as key, ``ModelOutput``mostly behaves like a dictionnary for compatiblity with torch.DataParallel:
- Sub-class of ``dict``
- providing an iterator of (keys, values) to the first argument without any other argument will set the associated attributes
- ``__setitem__``, ``__delitem__``, ``setdefault``, ``pop``, `ùpdate`` will raise errors
- ``__iter__`` iterate over the keys
"""
def __post_init__(self):
""" Does a few safety checks.
If the first field is a list/tuple, spread its values in the other fields.
This is currently necessary for compatibility with torch.DataParallel which only handles dict/list/tuples
cf. https://github.com/huggingface/transformers/issues/5693
and https://github.com/pytorch/pytorch/issues/41327
"""
class_fields = fields(self)
# Safety and consistency checks
assert len(class_fields), f"{self.__class__.__name__} has no fields."
assert all(
field.default is None for field in class_fields[1:]
), f"{self.__class__.__name__} should not have more than one required field."
# Check if we should spread the first field on the other fields in a dict-mapping fashion
first_field = getattr(self, class_fields[0].name)
other_fields_are_none = all(getattr(self, field.name) is None for field in class_fields[1:])
if other_fields_are_none:
try:
iterator = iter(first_field)
first_field_iterator = True
except TypeError:
first_field_iterator = False
# if we provided an iterator as first field and the iterator is a (key, value) iterator
# set the associated fields
if first_field_iterator:
for element in iterator:
if (
not isinstance(element, (list, tuple))
or not len(element) == 2
or not isinstance(element[0], str)
):
break
setattr(self, element[0], element[1])
def __delitem__(self, *args, **kwargs):
raise Exception(f"You cannot use ``__delitem__`` on a {self.__class__.__name__} instance.")
def setdefault(self, *args, **kwargs):
raise Exception(f"You cannot use ``setdefault`` on a {self.__class__.__name__} instance.")
def pop(self, *args, **kwargs):
raise Exception(f"You cannot use ``pop`` on a {self.__class__.__name__} instance.")
def update(self, *args, **kwargs):
raise Exception(f"You cannot use ``update`` on a {self.__class__.__name__} instance.")
def __iter__(self):
""" Will return the attributes names (dictionnary-like) """
for f in self.__dataclass_fields__.keys():
if getattr(self, f, None) is not None:
yield f
def to_tuple(self):
"""
Converts :obj:`self` to a tuple.
@@ -881,5 +953,11 @@ class ModelOutput:
def __getitem__(self, i):
return self.to_dict()[i] if isinstance(i, str) else self.to_tuple()[i]
def __setitem__(self, key, value):
if isinstance(key, str):
setattr(self, key, value)
else:
raise Exception(f"Key {key} must be a string but was given a {type(key)}.")
def __len__(self):
return len(self.to_tuple())
+3 -3
View File
@@ -430,9 +430,9 @@ class AlbertForPretrainingOutput(ModelOutput):
heads.
"""
loss: Optional[torch.FloatTensor]
prediction_logits: torch.FloatTensor
sop_logits: torch.FloatTensor
loss: torch.FloatTensor = None
prediction_logits: torch.FloatTensor = None
sop_logits: torch.FloatTensor = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
+3 -3
View File
@@ -605,9 +605,9 @@ class BertForPretrainingOutput(ModelOutput):
heads.
"""
loss: Optional[torch.FloatTensor]
prediction_logits: torch.FloatTensor
seq_relationship_logits: torch.FloatTensor
loss: torch.FloatTensor = None
prediction_logits: torch.FloatTensor = None
seq_relationship_logits: torch.FloatTensor = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
+5 -5
View File
@@ -73,7 +73,7 @@ class DPRContextEncoderOutput(ModelOutput):
heads.
"""
pooler_output: torch.FloatTensor
pooler_output: torch.FloatTensor = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -102,7 +102,7 @@ class DPRQuestionEncoderOutput(ModelOutput):
heads.
"""
pooler_output: torch.FloatTensor
pooler_output: torch.FloatTensor = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -133,9 +133,9 @@ class DPRReaderOutput(ModelOutput):
heads.
"""
start_logits: torch.FloatTensor
end_logits: torch.FloatTensor
relevance_logits: torch.FloatTensor
start_logits: torch.FloatTensor = None
end_logits: torch.FloatTensor = None
relevance_logits: torch.FloatTensor = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
+2 -2
View File
@@ -208,8 +208,8 @@ class ElectraForPretrainingOutput(ModelOutput):
heads.
"""
loss: Optional[torch.FloatTensor]
logits: torch.FloatTensor
loss: torch.FloatTensor = None
logits: torch.FloatTensor = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
+4 -4
View File
@@ -323,10 +323,10 @@ class GPT2DoubleHeadsModelOutput(ModelOutput):
heads.
"""
lm_loss: Optional[torch.FloatTensor]
mc_loss: Optional[torch.FloatTensor]
lm_logits: torch.FloatTensor
mc_logits: torch.FloatTensor
lm_loss: torch.FloatTensor = None
mc_loss: torch.FloatTensor = None
lm_logits: torch.FloatTensor = None
mc_logits: torch.FloatTensor = None
past_key_values: Optional[List[torch.FloatTensor]] = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
+3 -3
View File
@@ -705,9 +705,9 @@ class MobileBertForPretrainingOutput(ModelOutput):
heads.
"""
loss: Optional[torch.FloatTensor]
prediction_logits: torch.FloatTensor
seq_relationship_logits: torch.FloatTensor
loss: torch.FloatTensor = None
prediction_logits: torch.FloatTensor = None
seq_relationship_logits: torch.FloatTensor = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
+4 -4
View File
@@ -314,10 +314,10 @@ class OpenAIGPTDoubleHeadsModelOutput(ModelOutput):
heads.
"""
lm_loss: Optional[torch.FloatTensor]
mc_loss: Optional[torch.FloatTensor]
lm_logits: torch.FloatTensor
mc_logits: torch.FloatTensor
lm_loss: torch.FloatTensor = None
mc_loss: torch.FloatTensor = None
lm_logits: torch.FloatTensor = None
mc_logits: torch.FloatTensor = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
+25 -25
View File
@@ -63,7 +63,7 @@ class BaseModelOutputWithPooling(ModelOutput):
"""
last_hidden_state: torch.FloatTensor
pooler_output: torch.FloatTensor
pooler_output: torch.FloatTensor = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -178,8 +178,8 @@ class CausalLMOutput(ModelOutput):
heads.
"""
loss: Optional[torch.FloatTensor]
logits: torch.FloatTensor
loss: torch.FloatTensor = None
logits: torch.FloatTensor = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -213,8 +213,8 @@ class CausalLMOutputWithPast(ModelOutput):
heads.
"""
loss: Optional[torch.FloatTensor]
logits: torch.FloatTensor
loss: torch.FloatTensor = None
logits: torch.FloatTensor = None
past_key_values: Optional[List[torch.FloatTensor]] = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -243,8 +243,8 @@ class MaskedLMOutput(ModelOutput):
heads.
"""
loss: Optional[torch.FloatTensor]
logits: torch.FloatTensor
loss: torch.FloatTensor = None
logits: torch.FloatTensor = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -291,8 +291,8 @@ class Seq2SeqLMOutput(ModelOutput):
self-attention heads.
"""
loss: Optional[torch.FloatTensor]
logits: torch.FloatTensor
loss: torch.FloatTensor = None
logits: torch.FloatTensor = None
decoder_past_key_values: Optional[List[torch.FloatTensor]] = None
decoder_hidden_states: Optional[Tuple[torch.FloatTensor]] = None
decoder_attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -324,8 +324,8 @@ class NextSentencePredictorOutput(ModelOutput):
heads.
"""
loss: Optional[torch.FloatTensor]
logits: torch.FloatTensor
loss: torch.FloatTensor = None
logits: torch.FloatTensor = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -353,8 +353,8 @@ class SequenceClassifierOutput(ModelOutput):
heads.
"""
loss: Optional[torch.FloatTensor]
logits: torch.FloatTensor
loss: torch.FloatTensor = None
logits: torch.FloatTensor = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -401,8 +401,8 @@ class Seq2SeqSequenceClassifierOutput(ModelOutput):
self-attention heads.
"""
loss: Optional[torch.FloatTensor]
logits: torch.FloatTensor
loss: torch.FloatTensor = None
logits: torch.FloatTensor = None
decoder_past_key_values: Optional[List[torch.FloatTensor]] = None
decoder_hidden_states: Optional[Tuple[torch.FloatTensor]] = None
decoder_attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -436,8 +436,8 @@ class MultipleChoiceModelOutput(ModelOutput):
heads.
"""
loss: Optional[torch.FloatTensor]
logits: torch.FloatTensor
loss: torch.FloatTensor = None
logits: torch.FloatTensor = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -465,8 +465,8 @@ class TokenClassifierOutput(ModelOutput):
heads.
"""
loss: Optional[torch.FloatTensor]
logits: torch.FloatTensor
loss: torch.FloatTensor = None
logits: torch.FloatTensor = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -496,9 +496,9 @@ class QuestionAnsweringModelOutput(ModelOutput):
heads.
"""
loss: Optional[torch.FloatTensor]
start_logits: torch.FloatTensor
end_logits: torch.FloatTensor
loss: torch.FloatTensor = None
start_logits: torch.FloatTensor = None
end_logits: torch.FloatTensor = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -547,9 +547,9 @@ class Seq2SeqQuestionAnsweringModelOutput(ModelOutput):
self-attention heads.
"""
loss: Optional[torch.FloatTensor]
start_logits: torch.FloatTensor
end_logits: torch.FloatTensor
loss: torch.FloatTensor = None
start_logits: torch.FloatTensor = None
end_logits: torch.FloatTensor = None
decoder_past_key_values: Optional[List[torch.FloatTensor]] = None
decoder_hidden_states: Optional[Tuple[torch.FloatTensor]] = None
decoder_attentions: Optional[Tuple[torch.FloatTensor]] = None
+3 -3
View File
@@ -618,7 +618,7 @@ class TransfoXLModelOutput(ModelOutput):
"""
last_hidden_state: torch.FloatTensor
mems: List[torch.FloatTensor]
mems: Optional[List[torch.FloatTensor]] = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -653,8 +653,8 @@ class TransfoXLLMHeadModelOutput(ModelOutput):
"""
losses: Optional[torch.FloatTensor]
prediction_scores: torch.FloatTensor
mems: List[torch.FloatTensor]
prediction_scores: torch.FloatTensor = None
mems: Optional[List[torch.FloatTensor]] = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
+11 -11
View File
@@ -627,8 +627,8 @@ class XLNetLMHeadModelOutput(ModelOutput):
heads.
"""
loss: Optional[torch.FloatTensor]
logits: torch.FloatTensor
loss: torch.FloatTensor = None
logits: torch.FloatTensor = None
mems: Optional[List[torch.FloatTensor]] = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -661,8 +661,8 @@ class XLNetForSequenceClassificationOutput(ModelOutput):
heads.
"""
loss: Optional[torch.FloatTensor]
logits: torch.FloatTensor
loss: torch.FloatTensor = None
logits: torch.FloatTensor = None
mems: Optional[List[torch.FloatTensor]] = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -695,8 +695,8 @@ class XLNetForTokenClassificationOutput(ModelOutput):
heads.
"""
loss: Optional[torch.FloatTensor]
logits: torch.FloatTensor
loss: torch.FloatTensor = None
logits: torch.FloatTensor = None
mems: Optional[List[torch.FloatTensor]] = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -731,8 +731,8 @@ class XLNetForMultipleChoiceOutput(ModelOutput):
heads.
"""
loss: Optional[torch.FloatTensor]
logits: torch.FloatTensor
loss: torch.FloatTensor = None
logits: torch.FloatTensor = None
mems: Optional[List[torch.FloatTensor]] = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
@@ -767,9 +767,9 @@ class XLNetForQuestionAnsweringSimpleOutput(ModelOutput):
heads.
"""
loss: Optional[torch.FloatTensor]
start_logits: torch.FloatTensor
end_logits: torch.FloatTensor
loss: torch.FloatTensor = None
start_logits: torch.FloatTensor = None
end_logits: torch.FloatTensor = None
mems: Optional[List[torch.FloatTensor]] = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None