Compare commits

...
14 changed files with 15 additions and 15 deletions
+2 -2
View File
@@ -43,8 +43,8 @@ class PretrainedConfig(object):
- **is_composition** (:obj:`bool`): Whether the config class is composed of multiple sub-configs. In this case - **is_composition** (:obj:`bool`): Whether the config class is composed of multiple sub-configs. In this case
the config has to be initialized from two or more configs of type :class:`~transformers.PretrainedConfig` the config has to be initialized from two or more configs of type :class:`~transformers.PretrainedConfig`
like: :class:`~transformers.EncoderDecoderConfig` or :class:`~RagConfig`. like: :class:`~transformers.EncoderDecoderConfig` or :class:`~RagConfig`.
- **keys_to_ignore_at_inference** (:obj:`List[str]`): A list of keys to ignore by default when looking at - **output_keys_to_ignore_at_inference** (:obj:`List[str]`): A list of keys to ignore by default when looking
dictionary outputs of the model during inference. at dictionary outputs of the model during inference.
Args: Args:
name_or_path (:obj:`str`, `optional`, defaults to :obj:`""`): name_or_path (:obj:`str`, `optional`, defaults to :obj:`""`):
@@ -112,7 +112,7 @@ class BartConfig(PretrainedConfig):
Whether or not the model should return the last key/values attentions (not used by all models). Whether or not the model should return the last key/values attentions (not used by all models).
""" """
model_type = "bart" model_type = "bart"
keys_to_ignore_at_inference = ["past_key_values"] output_keys_to_ignore_at_inference = ["past_key_values"]
def __init__( def __init__(
self, self,
@@ -80,7 +80,7 @@ class CTRLConfig(PretrainedConfig):
""" """
model_type = "ctrl" model_type = "ctrl"
keys_to_ignore_at_inference = ["past_key_values"] output_keys_to_ignore_at_inference = ["past_key_values"]
def __init__( def __init__(
self, self,
@@ -122,7 +122,7 @@ class GPT2Config(PretrainedConfig):
""" """
model_type = "gpt2" model_type = "gpt2"
keys_to_ignore_at_inference = ["past_key_values"] output_keys_to_ignore_at_inference = ["past_key_values"]
def __init__( def __init__(
self, self,
@@ -97,4 +97,4 @@ class MarianConfig(BartConfig):
""" """
model_type = "marian" model_type = "marian"
keys_to_ignore_at_inference = ["past_key_values"] output_keys_to_ignore_at_inference = ["past_key_values"]
@@ -102,4 +102,4 @@ class MBartConfig(BartConfig):
""" """
model_type = "mbart" model_type = "mbart"
keys_to_ignore_at_inference = ["past_key_values"] output_keys_to_ignore_at_inference = ["past_key_values"]
@@ -64,7 +64,7 @@ class MT5Config(PretrainedConfig):
Whether or not the model should return the last key/values attentions (not used by all models). Whether or not the model should return the last key/values attentions (not used by all models).
""" """
model_type = "mt5" model_type = "mt5"
keys_to_ignore_at_inference = ["past_key_values"] output_keys_to_ignore_at_inference = ["past_key_values"]
def __init__( def __init__(
self, self,
@@ -141,5 +141,5 @@ class PegasusConfig(BartConfig):
""" """
model_type = "pegasus" model_type = "pegasus"
keys_to_ignore_at_inference = ["past_key_values"] output_keys_to_ignore_at_inference = ["past_key_values"]
# The implementation of the config object is in BartConfig # The implementation of the config object is in BartConfig
@@ -94,7 +94,7 @@ class ProphetNetConfig(PretrainedConfig):
Whether or not the model should return the last key/values attentions (not used by all models). Whether or not the model should return the last key/values attentions (not used by all models).
""" """
model_type = "prophetnet" model_type = "prophetnet"
keys_to_ignore_at_inference = ["past_key_values"] output_keys_to_ignore_at_inference = ["past_key_values"]
def __init__( def __init__(
self, self,
@@ -155,7 +155,7 @@ class ReformerConfig(PretrainedConfig):
>>> configuration = model.config >>> configuration = model.config
""" """
model_type = "reformer" model_type = "reformer"
keys_to_ignore_at_inference = ["past_buckets_states"] output_keys_to_ignore_at_inference = ["past_buckets_states"]
def __init__( def __init__(
self, self,
@@ -73,7 +73,7 @@ class T5Config(PretrainedConfig):
Whether or not the model should return the last key/values attentions (not used by all models). Whether or not the model should return the last key/values attentions (not used by all models).
""" """
model_type = "t5" model_type = "t5"
keys_to_ignore_at_inference = ["past_key_values"] output_keys_to_ignore_at_inference = ["past_key_values"]
def __init__( def __init__(
self, self,
@@ -105,7 +105,7 @@ class TransfoXLConfig(PretrainedConfig):
""" """
model_type = "transfo-xl" model_type = "transfo-xl"
keys_to_ignore_at_inference = ["mems"] output_keys_to_ignore_at_inference = ["mems"]
def __init__( def __init__(
self, self,
@@ -136,7 +136,7 @@ class XLNetConfig(PretrainedConfig):
""" """
model_type = "xlnet" model_type = "xlnet"
keys_to_ignore_at_inference = ["mems"] output_keys_to_ignore_at_inference = ["mems"]
def __init__( def __init__(
self, self,
+1 -1
View File
@@ -1463,7 +1463,7 @@ class Trainer:
inputs = self._prepare_inputs(inputs) inputs = self._prepare_inputs(inputs)
if ignore_keys is None: if ignore_keys is None:
if hasattr(self.model, "config"): if hasattr(self.model, "config"):
ignore_keys = getattr(self.model.config, "keys_to_ignore_at_inference", []) ignore_keys = getattr(self.model.config, "output_keys_to_ignore_at_inference", [])
else: else:
ignore_keys = [] ignore_keys = []