This commit is contained in:
yuze.zyz
2025-01-26 16:53:50 +08:00
parent 4723e5c0ff
commit 1900b57450
3 changed files with 60 additions and 67 deletions

View File

@@ -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

View File

@@ -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,

View File

@@ -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):