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
the config has to be initialized from two or more configs of type :class:`~transformers.PretrainedConfig`
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
dictionary outputs of the model during inference.
- **output_keys_to_ignore_at_inference** (:obj:`List[str]`): A list of keys to ignore by default when looking
at dictionary outputs of the model during inference.
Args:
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).
"""
model_type = "bart"
keys_to_ignore_at_inference = ["past_key_values"]
output_keys_to_ignore_at_inference = ["past_key_values"]
def __init__(
self,
@@ -80,7 +80,7 @@ class CTRLConfig(PretrainedConfig):
"""
model_type = "ctrl"
keys_to_ignore_at_inference = ["past_key_values"]
output_keys_to_ignore_at_inference = ["past_key_values"]
def __init__(
self,
@@ -122,7 +122,7 @@ class GPT2Config(PretrainedConfig):
"""
model_type = "gpt2"
keys_to_ignore_at_inference = ["past_key_values"]
output_keys_to_ignore_at_inference = ["past_key_values"]
def __init__(
self,
@@ -97,4 +97,4 @@ class MarianConfig(BartConfig):
"""
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"
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).
"""
model_type = "mt5"
keys_to_ignore_at_inference = ["past_key_values"]
output_keys_to_ignore_at_inference = ["past_key_values"]
def __init__(
self,
@@ -141,5 +141,5 @@ class PegasusConfig(BartConfig):
"""
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
@@ -94,7 +94,7 @@ class ProphetNetConfig(PretrainedConfig):
Whether or not the model should return the last key/values attentions (not used by all models).
"""
model_type = "prophetnet"
keys_to_ignore_at_inference = ["past_key_values"]
output_keys_to_ignore_at_inference = ["past_key_values"]
def __init__(
self,
@@ -155,7 +155,7 @@ class ReformerConfig(PretrainedConfig):
>>> configuration = model.config
"""
model_type = "reformer"
keys_to_ignore_at_inference = ["past_buckets_states"]
output_keys_to_ignore_at_inference = ["past_buckets_states"]
def __init__(
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).
"""
model_type = "t5"
keys_to_ignore_at_inference = ["past_key_values"]
output_keys_to_ignore_at_inference = ["past_key_values"]
def __init__(
self,
@@ -105,7 +105,7 @@ class TransfoXLConfig(PretrainedConfig):
"""
model_type = "transfo-xl"
keys_to_ignore_at_inference = ["mems"]
output_keys_to_ignore_at_inference = ["mems"]
def __init__(
self,
@@ -136,7 +136,7 @@ class XLNetConfig(PretrainedConfig):
"""
model_type = "xlnet"
keys_to_ignore_at_inference = ["mems"]
output_keys_to_ignore_at_inference = ["mems"]
def __init__(
self,
+1 -1
View File
@@ -1463,7 +1463,7 @@ class Trainer:
inputs = self._prepare_inputs(inputs)
if ignore_keys is None:
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:
ignore_keys = []