mirror of
https://github.com/modelscope/modelscope.git
synced 2026-09-01 19:49:03 +02:00
fix
This commit is contained in:
@@ -1,6 +1,4 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import inspect
|
||||
import os
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -71,32 +69,9 @@ if TYPE_CHECKING:
|
||||
|
||||
else:
|
||||
|
||||
class UnsupportedAutoClass:
|
||||
|
||||
def __init__(self, name: str):
|
||||
self.error_msg =\
|
||||
f'{name} is not supported with your installed Transformers version {transformers_version}. ' + \
|
||||
'Please update your Transformers by "pip install transformers -U".'
|
||||
|
||||
def from_pretrained(self, pretrained_model_name_or_path, *model_args,
|
||||
**kwargs):
|
||||
raise ImportError(self.error_msg)
|
||||
|
||||
def from_config(self, cls, config):
|
||||
raise ImportError(self.error_msg)
|
||||
|
||||
def user_agent(invoked_by=None):
|
||||
from modelscope.utils.constant import Invoke
|
||||
|
||||
if invoked_by is None:
|
||||
invoked_by = Invoke.PRETRAINED
|
||||
uagent = '%s/%s' % (Invoke.KEY, invoked_by)
|
||||
return uagent
|
||||
|
||||
from .patcher import get_all_imported_modules, _patch_pretrained_class
|
||||
|
||||
all_imported_modules = get_all_imported_modules()
|
||||
all_available_modules = _patch_pretrained_class(all_imported_modules, wrap=True)
|
||||
all_available_modules = _patch_pretrained_class(
|
||||
get_all_imported_modules(), wrap=True)
|
||||
|
||||
for module in all_available_modules:
|
||||
globals()[module.__name__] = module
|
||||
|
||||
@@ -13,8 +13,10 @@ from typing import BinaryIO, Dict, List, Optional, Union
|
||||
def get_all_imported_modules():
|
||||
"""Find all modules in transformers/peft/diffusers"""
|
||||
all_imported_modules = []
|
||||
transformers_include_names = ['Auto', 'T5', 'BitsAndBytes', 'GenerationConfig',
|
||||
'Quant', 'Awq', 'GPTQ', 'BatchFeature', 'Qwen2']
|
||||
transformers_include_names = [
|
||||
'Auto', 'T5', 'BitsAndBytes', 'GenerationConfig', 'Quant', 'Awq',
|
||||
'GPTQ', 'BatchFeature', 'Qwen2'
|
||||
]
|
||||
diffusers_include_names = ['Pipeline']
|
||||
if importlib.util.find_spec('transformers') is not None:
|
||||
import transformers
|
||||
@@ -48,7 +50,8 @@ def get_all_imported_modules():
|
||||
for key in _import_structure:
|
||||
values = _import_structure[key]
|
||||
for value in values:
|
||||
if any([name in value for name in diffusers_include_names]):
|
||||
if any([name in value
|
||||
for name in diffusers_include_names]):
|
||||
try:
|
||||
module = importlib.import_module(
|
||||
f'.{key}', diffusers.__name__)
|
||||
@@ -60,6 +63,14 @@ def get_all_imported_modules():
|
||||
|
||||
|
||||
def _patch_pretrained_class(all_imported_modules, wrap=False):
|
||||
"""Patch all class to download from modelscope
|
||||
|
||||
Args:
|
||||
wrap: Wrap the class or monkey patch the original class
|
||||
|
||||
Returns:
|
||||
The classes after patched
|
||||
"""
|
||||
|
||||
def get_model_dir(pretrained_model_name_or_path,
|
||||
ignore_file_pattern=None,
|
||||
@@ -67,39 +78,40 @@ def _patch_pretrained_class(all_imported_modules, wrap=False):
|
||||
**kwargs):
|
||||
from modelscope import snapshot_download
|
||||
if not os.path.exists(pretrained_model_name_or_path):
|
||||
revision = kwargs.pop('revision', None)
|
||||
model_dir = snapshot_download(
|
||||
pretrained_model_name_or_path,
|
||||
revision=revision,
|
||||
revision=kwargs.pop('revision', None),
|
||||
ignore_file_pattern=ignore_file_pattern,
|
||||
allow_file_pattern=allow_file_pattern)
|
||||
else:
|
||||
model_dir = pretrained_model_name_or_path
|
||||
return model_dir
|
||||
|
||||
ignore_file_pattern = [r'\w+\.bin', r'\w+\.safetensors', r'\w+\.pth', r'\w+\.pt', r'\w+\.h5']
|
||||
ignore_file_pattern = [
|
||||
r'\w+\.bin', r'\w+\.safetensors', r'\w+\.pth', r'\w+\.pt', r'\w+\.h5'
|
||||
]
|
||||
|
||||
def patch_pretrained_model_name_or_path(pretrained_model_name_or_path,
|
||||
*model_args, **kwargs):
|
||||
model_dir = get_model_dir(pretrained_model_name_or_path,
|
||||
kwargs.pop('ignore_file_pattern', None),
|
||||
**kwargs)
|
||||
"""Patch all from_pretrained/get_config_dict"""
|
||||
model_dir = get_model_dir(pretrained_model_name_or_path, **kwargs)
|
||||
return kwargs.pop('ori_func')(model_dir, *model_args, **kwargs)
|
||||
|
||||
def patch_peft_model_id(model, model_id, *model_args, **kwargs):
|
||||
model_dir = get_model_dir(model_id,
|
||||
kwargs.pop('ignore_file_pattern', None),
|
||||
**kwargs)
|
||||
"""Patch all peft.from_pretrained"""
|
||||
model_dir = get_model_dir(model_id, **kwargs)
|
||||
return kwargs.pop('ori_func')(model, model_dir, *model_args, **kwargs)
|
||||
|
||||
def _get_peft_type(model_id, **kwargs):
|
||||
model_dir = get_model_dir(model_id, ignore_file_pattern, **kwargs)
|
||||
"""Patch all _get_peft_type"""
|
||||
model_dir = get_model_dir(model_id, **kwargs)
|
||||
return kwargs.pop('ori_func')(model_dir, **kwargs)
|
||||
|
||||
def get_wrapped_class(module_class: 'PreTrainedModel',
|
||||
ignore_file_pattern: Optional[Union[str, List[str]]] = None,
|
||||
allow_file_pattern: Optional[Union[str, List[str]]] = None,
|
||||
**kwargs):
|
||||
def get_wrapped_class(
|
||||
module_class: 'PreTrainedModel',
|
||||
ignore_file_pattern: Optional[Union[str, List[str]]] = None,
|
||||
allow_file_pattern: Optional[Union[str, List[str]]] = None,
|
||||
**kwargs):
|
||||
"""Get a custom wrapper class for auto classes to download the models from the ModelScope hub
|
||||
Args:
|
||||
module_class (`PreTrainedModel`): The actual module class
|
||||
@@ -108,17 +120,19 @@ def _patch_pretrained_class(all_imported_modules, wrap=False):
|
||||
allow_file_pattern (`str` or `List`, *optional*, default to `None`):
|
||||
Any file pattern to be included, like exact file names or file extensions.
|
||||
Returns:
|
||||
The wrapper
|
||||
The wrapped class
|
||||
"""
|
||||
|
||||
def from_pretrained(model, model_id, *model_args, **kwargs):
|
||||
model_dir = get_model_dir(model_id,
|
||||
ignore_file_pattern=ignore_file_pattern,
|
||||
allow_file_pattern=allow_file_pattern,
|
||||
**kwargs)
|
||||
# model is an instance
|
||||
model_dir = get_model_dir(
|
||||
model_id,
|
||||
ignore_file_pattern=ignore_file_pattern,
|
||||
allow_file_pattern=allow_file_pattern,
|
||||
**kwargs)
|
||||
|
||||
module_obj = module_class.from_pretrained(
|
||||
model, model_dir, *model_args, **kwargs)
|
||||
module_obj = module_class.from_pretrained(model, model_dir,
|
||||
*model_args, **kwargs)
|
||||
|
||||
return module_obj
|
||||
|
||||
@@ -127,10 +141,11 @@ def _patch_pretrained_class(all_imported_modules, wrap=False):
|
||||
@classmethod
|
||||
def from_pretrained(cls, pretrained_model_name_or_path,
|
||||
*model_args, **kwargs):
|
||||
model_dir = get_model_dir(pretrained_model_name_or_path,
|
||||
ignore_file_pattern=ignore_file_pattern,
|
||||
allow_file_pattern=allow_file_pattern,
|
||||
**kwargs)
|
||||
model_dir = get_model_dir(
|
||||
pretrained_model_name_or_path,
|
||||
ignore_file_pattern=ignore_file_pattern,
|
||||
allow_file_pattern=allow_file_pattern,
|
||||
**kwargs)
|
||||
|
||||
module_obj = module_class.from_pretrained(
|
||||
model_dir, *model_args, **kwargs)
|
||||
@@ -141,18 +156,14 @@ def _patch_pretrained_class(all_imported_modules, wrap=False):
|
||||
|
||||
@classmethod
|
||||
def _get_peft_type(cls, model_id, **kwargs):
|
||||
model_dir = get_model_dir(model_id,
|
||||
kwargs.pop('ignore_file_pattern', None),
|
||||
**kwargs)
|
||||
|
||||
module_obj = module_class._get_peft_type(
|
||||
model_dir, **kwargs)
|
||||
model_dir = get_model_dir(model_id, **kwargs)
|
||||
module_obj = module_class._get_peft_type(model_dir, **kwargs)
|
||||
return module_obj
|
||||
|
||||
@classmethod
|
||||
def get_config_dict(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
def get_config_dict(cls, pretrained_model_name_or_path,
|
||||
*model_args, **kwargs):
|
||||
model_dir = get_model_dir(pretrained_model_name_or_path,
|
||||
kwargs.pop('ignore_file_pattern', None),
|
||||
**kwargs)
|
||||
|
||||
module_obj = module_class.get_config_dict(
|
||||
@@ -204,11 +215,13 @@ def _patch_pretrained_class(all_imported_modules, wrap=False):
|
||||
if not has_from_pretrained and not has_get_config_dict and not has_get_peft_type:
|
||||
all_available_modules.append(var)
|
||||
else:
|
||||
all_available_modules.append(get_wrapped_class(var, ignore_file_pattern))
|
||||
all_available_modules.append(
|
||||
get_wrapped_class(var, ignore_file_pattern))
|
||||
except Exception:
|
||||
all_available_modules.append(var)
|
||||
else:
|
||||
if has_from_pretrained and not hasattr(var, '_from_pretrained_origin'):
|
||||
if has_from_pretrained and not hasattr(var,
|
||||
'_from_pretrained_origin'):
|
||||
parameters = inspect.signature(var.from_pretrained).parameters
|
||||
# different argument names
|
||||
is_peft = 'model' in parameters and 'model_id' in parameters
|
||||
@@ -230,7 +243,8 @@ def _patch_pretrained_class(all_imported_modules, wrap=False):
|
||||
ori_func=var._get_peft_type_origin,
|
||||
**ignore_file_pattern_kwargs)
|
||||
|
||||
if has_get_config_dict and not hasattr(var, '_get_config_dict_origin'):
|
||||
if has_get_config_dict and not hasattr(var,
|
||||
'_get_config_dict_origin'):
|
||||
var._get_config_dict_origin = var.get_config_dict
|
||||
var.get_config_dict = partial(
|
||||
patch_pretrained_model_name_or_path,
|
||||
|
||||
@@ -152,8 +152,12 @@ class HFUtilTest(unittest.TestCase):
|
||||
|
||||
def test_patch_peft(self):
|
||||
with patch_context():
|
||||
from transformers import AutoModelForCausalLM
|
||||
from peft import PeftModel
|
||||
self.assertTrue(hasattr(PeftModel, '_from_pretrained_origin'))
|
||||
model = AutoModelForCausalLM.from_pretrained('OpenBMB/MiniCPM3-4B')
|
||||
model = PeftModel.from_pretrained(model,
|
||||
'OpenBMB/MiniCPM3-RAG-LoRA')
|
||||
self.assertTrue(model is not None)
|
||||
self.assertFalse(hasattr(PeftModel, '_from_pretrained_origin'))
|
||||
|
||||
def test_patch_file_exists(self):
|
||||
|
||||
Reference in New Issue
Block a user