158 lines
6.2 KiB
Python
158 lines
6.2 KiB
Python
# Copyright (c) 2024 PaddlePaddle Authors. All Rights Reserved.
|
|
# Copyright 2018 Google AI, Google Brain and the HuggingFace Inc. team.
|
|
#
|
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
# you may not use this file except in compliance with the License.
|
|
# You may obtain a copy of the License at
|
|
#
|
|
# http://www.apache.org/licenses/LICENSE-2.0
|
|
#
|
|
# Unless required by applicable law or agreed to in writing, software
|
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
# See the License for the specific language governing permissions and
|
|
# limitations under the License.
|
|
|
|
import importlib
|
|
from collections import OrderedDict
|
|
|
|
from paddlenlp.transformers.auto.configuration import model_type_to_module_name
|
|
|
|
|
|
def getattribute_from_module(module, attr):
|
|
if attr is None:
|
|
return None
|
|
if isinstance(attr, tuple):
|
|
return tuple(getattribute_from_module(module, a) for a in attr)
|
|
if hasattr(module, attr):
|
|
return getattr(module, attr)
|
|
# Some of the mappings have entries model_type -> object of another model type. In that case we try to grab the
|
|
# object at the top level.
|
|
paddlenlp_module = importlib.import_module("paddlenlp")
|
|
|
|
if module != paddlenlp_module:
|
|
try:
|
|
return getattribute_from_module(paddlenlp_module, attr)
|
|
except ValueError:
|
|
raise ValueError(f"Could not find {attr} neither in {module} nor in {paddlenlp_module}!")
|
|
else:
|
|
raise ValueError(f"Could not find {attr} in {paddlenlp_module}!")
|
|
|
|
|
|
class _LazyAutoMapping(OrderedDict):
|
|
"""
|
|
" A mapping config to object (model or tokenizer for instance) that will load keys and values when it is accessed.
|
|
|
|
Args:
|
|
- config_mapping: The map model type to config class
|
|
- model_mapping: The map model type to model (or tokenizer) class
|
|
"""
|
|
|
|
def __init__(self, config_mapping, model_mapping):
|
|
self._config_mapping = config_mapping
|
|
self._reverse_config_mapping = {v: k for k, v in config_mapping.items()}
|
|
self._model_mapping = model_mapping
|
|
self._model_mapping._model_mapping = self
|
|
self._extra_content = {}
|
|
self._modules = {}
|
|
|
|
def __len__(self):
|
|
common_keys = set(self._config_mapping.keys()).intersection(self._model_mapping.keys())
|
|
return len(common_keys) + len(self._extra_content)
|
|
|
|
def __getitem__(self, key):
|
|
if key in self._extra_content:
|
|
return self._extra_content[key]
|
|
model_type = self._reverse_config_mapping[key.__name__]
|
|
if model_type in self._model_mapping:
|
|
model_name = self._model_mapping[model_type]
|
|
return self._load_attr_from_module(model_type, model_name)
|
|
|
|
# Maybe there was several model types associated with this config.
|
|
model_types = [k for k, v in self._config_mapping.items() if v == key.__name__]
|
|
for mtype in model_types:
|
|
if mtype in self._model_mapping:
|
|
model_name = self._model_mapping[mtype]
|
|
return self._load_attr_from_module(mtype, model_name)
|
|
raise KeyError(key)
|
|
|
|
def _load_attr_from_module(self, model_type, attr):
|
|
module_name = model_type_to_module_name(model_type)
|
|
if module_name not in self._modules:
|
|
if any(["Tokenizer" in name for name in [model_type, attr]]):
|
|
try:
|
|
self._modules[module_name] = importlib.import_module(
|
|
f".{module_name}.tokenizer", "paddlenlp.transformers"
|
|
)
|
|
except ImportError:
|
|
pass
|
|
if module_name not in self._modules:
|
|
if any(["Config" in name for name in [model_type, attr]]):
|
|
try:
|
|
self._modules[module_name] = importlib.import_module(
|
|
f".{module_name}.configuration", "paddlenlp.transformers"
|
|
)
|
|
except ImportError:
|
|
pass
|
|
if module_name not in self._modules:
|
|
self._modules[module_name] = importlib.import_module(f".{module_name}", "paddlenlp.transformers")
|
|
return getattribute_from_module(self._modules[module_name], attr)
|
|
|
|
def keys(self):
|
|
mapping_keys = [
|
|
self._load_attr_from_module(key, name)
|
|
for key, name in self._config_mapping.items()
|
|
if key in self._model_mapping.keys()
|
|
]
|
|
return mapping_keys + list(self._extra_content.keys())
|
|
|
|
def get(self, key, default):
|
|
try:
|
|
return self.__getitem__(key)
|
|
except KeyError:
|
|
return default
|
|
|
|
def __bool__(self):
|
|
return bool(self.keys())
|
|
|
|
def values(self):
|
|
mapping_values = [
|
|
self._load_attr_from_module(key, name)
|
|
for key, name in self._model_mapping.items()
|
|
if key in self._config_mapping.keys()
|
|
]
|
|
return mapping_values + list(self._extra_content.values())
|
|
|
|
def items(self):
|
|
mapping_items = [
|
|
(
|
|
self._load_attr_from_module(key, self._config_mapping[key]),
|
|
self._load_attr_from_module(key, self._model_mapping[key]),
|
|
)
|
|
for key in self._model_mapping.keys()
|
|
if key in self._config_mapping.keys()
|
|
]
|
|
return mapping_items + list(self._extra_content.items())
|
|
|
|
def __iter__(self):
|
|
return iter(self.keys())
|
|
|
|
def __contains__(self, item):
|
|
if item in self._extra_content:
|
|
return True
|
|
if not hasattr(item, "__name__") or item.__name__ not in self._reverse_config_mapping:
|
|
return False
|
|
model_type = self._reverse_config_mapping[item.__name__]
|
|
return model_type in self._model_mapping
|
|
|
|
def register(self, key, value, exist_ok=False):
|
|
"""
|
|
Register a new model in this mapping.
|
|
"""
|
|
if hasattr(key, "__name__") and key.__name__ in self._reverse_config_mapping:
|
|
model_type = self._reverse_config_mapping[key.__name__]
|
|
if model_type in self._model_mapping.keys() and not exist_ok:
|
|
raise ValueError(f"'{key}' is already used by a Transformers model.")
|
|
|
|
self._extra_content[key] = value
|