Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
acfcbc8833 |
@@ -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())
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user