Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d53b78680d |
@@ -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,
|
||||||
|
|||||||
@@ -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 = []
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user