Compare commits

...
Author SHA1 Message Date
Lysandre 513ba3b6a0 Better warning 2020-06-09 19:08:27 -04:00
Lysandre 6342e803a9 warn with FutureWarning when using output_attentions in the configuration 2020-06-09 19:06:49 -04:00
+15
View File
@@ -20,6 +20,7 @@ import copy
import json
import logging
import os
import warnings
from typing import Dict, Tuple
from .file_utils import CONFIG_NAME, cached_path, hf_bucket_url, is_remote_url
@@ -54,6 +55,13 @@ class PretrainedConfig(object):
def __init__(self, **kwargs):
# Attributes with defaults
self.output_hidden_states = kwargs.pop("output_hidden_states", False)
if "output_attentions" in kwargs:
warnings.warn(
"The `output_attentions` in the configuration is deprecated and will be removed in a later version. "
"Please use `output_attention` as an argument to your model's foward method instead.",
FutureWarning,
)
self.output_attentions = kwargs.pop("output_attentions", False)
self.use_cache = kwargs.pop("use_cache", True) # Not used by all models
self.torchscript = kwargs.pop("torchscript", False) # Only used by PyTorch models
@@ -287,6 +295,13 @@ class PretrainedConfig(object):
if hasattr(config, "pruned_heads"):
config.pruned_heads = dict((int(key), value) for key, value in config.pruned_heads.items())
if "output_attentions" in kwargs:
warnings.warn(
"The `output_attentions` in the configuration is deprecated and will be removed in a later version. "
"Please use `output_attention` as an argument to your model's forward method instead.",
FutureWarning,
)
# Update config with kwargs if needed
to_remove = []
for key, value in kwargs.items():