mirror of
https://github.com/modelscope/modelscope.git
synced 2026-09-01 19:49:03 +02:00
fix: handle pretrained_model_name_or_path passed via kwargs
When the arg is only in kwargs, downloading and then forcing a positional overwrite caused a duplicate-keyword TypeError. Resolve and update kwargs or args consistently for both the download and cross-repo branches. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -175,12 +175,26 @@ def _get_class_from_dynamic_module(class_reference, *args, **kwargs):
|
||||
has_pretrained_arg = (
|
||||
'pretrained_model_name_or_path'
|
||||
in inspect.signature(origin_get_class_from_dynamic_module).parameters)
|
||||
# Resolve pretrained_model_name_or_path from kwargs or positional args.
|
||||
# ``args`` is a tuple; never mutate it in place.
|
||||
if has_pretrained_arg and args:
|
||||
pretrained_model_name_or_path = args[0]
|
||||
if not os.path.exists(pretrained_model_name_or_path):
|
||||
from modelscope import snapshot_download
|
||||
args = (snapshot_download(pretrained_model_name_or_path), ) + args[1:]
|
||||
pretrained_in_kwargs = False
|
||||
pretrained_model_name_or_path = None
|
||||
if has_pretrained_arg:
|
||||
if 'pretrained_model_name_or_path' in kwargs:
|
||||
pretrained_model_name_or_path = kwargs[
|
||||
'pretrained_model_name_or_path']
|
||||
pretrained_in_kwargs = True
|
||||
elif args:
|
||||
pretrained_model_name_or_path = args[0]
|
||||
if (pretrained_model_name_or_path is not None
|
||||
and not os.path.exists(pretrained_model_name_or_path)):
|
||||
from modelscope import snapshot_download
|
||||
downloaded_path = snapshot_download(pretrained_model_name_or_path)
|
||||
if pretrained_in_kwargs:
|
||||
kwargs['pretrained_model_name_or_path'] = downloaded_path
|
||||
else:
|
||||
args = (downloaded_path, ) + args[1:]
|
||||
pretrained_model_name_or_path = downloaded_path
|
||||
if '--' in class_reference:
|
||||
# Only the first ``--`` is the auto_map delimiter (repo vs module).
|
||||
repo_id, class_reference = class_reference.split('--', 1)
|
||||
@@ -197,7 +211,11 @@ def _get_class_from_dynamic_module(class_reference, *args, **kwargs):
|
||||
repo_id = snapshot_download(repo_id, **download_kwargs)
|
||||
if has_pretrained_arg:
|
||||
# Local path + bare class name; do not rejoin with ``--``.
|
||||
args = (repo_id, ) + args[1:]
|
||||
# Keep kwargs/positional form consistent with the original call.
|
||||
if pretrained_in_kwargs:
|
||||
kwargs['pretrained_model_name_or_path'] = repo_id
|
||||
else:
|
||||
args = (repo_id, ) + args[1:]
|
||||
else:
|
||||
# Legacy transformers without pretrained_model_name_or_path.
|
||||
# Unsafe if repo_id (local cache) contains '--'; modern
|
||||
|
||||
@@ -380,6 +380,58 @@ class HFUtilTest(unittest.TestCase):
|
||||
self.assertEqual(captured['class_reference'], 'modeling.Foo')
|
||||
self.assertEqual(captured['pretrained'], downloaded)
|
||||
|
||||
def test_dynamic_module_pretrained_via_kwargs(self):
|
||||
"""pretrained_model_name_or_path may be passed as a keyword argument."""
|
||||
from unittest import mock
|
||||
|
||||
from modelscope.utils.hf_util.patcher import \
|
||||
_get_class_from_dynamic_module
|
||||
|
||||
tmp = tempfile.mkdtemp()
|
||||
self.addCleanup(shutil.rmtree, tmp, ignore_errors=True)
|
||||
downloaded = os.path.join(tmp, 'models', 'org--model', 'snapshots',
|
||||
'rev')
|
||||
cross_repo = os.path.join(tmp, 'models', 'org--other', 'snapshots',
|
||||
'rev')
|
||||
os.makedirs(downloaded)
|
||||
os.makedirs(cross_repo)
|
||||
|
||||
captured = {}
|
||||
|
||||
def fake_origin(class_reference,
|
||||
pretrained_model_name_or_path,
|
||||
*args,
|
||||
**kwargs):
|
||||
captured['class_reference'] = class_reference
|
||||
captured['pretrained'] = pretrained_model_name_or_path
|
||||
captured['kwargs'] = kwargs
|
||||
return type('DummyConfig', (), {})
|
||||
|
||||
remote_id = 'org/model-not-on-disk'
|
||||
class_ref = 'org/other--configuration_foo.FooConfig'
|
||||
|
||||
def fake_download(repo_id, **kwargs):
|
||||
if repo_id == remote_id:
|
||||
return downloaded
|
||||
if repo_id == 'org/other':
|
||||
return cross_repo
|
||||
raise AssertionError(f'unexpected download: {repo_id}')
|
||||
|
||||
with mock.patch(
|
||||
'transformers.dynamic_module_utils.origin_get_class_from_dynamic_module',
|
||||
new=fake_origin,
|
||||
create=True):
|
||||
with mock.patch(
|
||||
'modelscope.snapshot_download', side_effect=fake_download):
|
||||
# Keyword form: must download and not pass duplicate positional.
|
||||
_get_class_from_dynamic_module(
|
||||
class_ref, pretrained_model_name_or_path=remote_id)
|
||||
|
||||
self.assertEqual(captured['class_reference'],
|
||||
'configuration_foo.FooConfig')
|
||||
self.assertEqual(captured['pretrained'], cross_repo)
|
||||
self.assertNotIn('pretrained_model_name_or_path', captured['kwargs'])
|
||||
|
||||
def test_import_not_pollute_dynamic_module(self):
|
||||
"""Importing from modelscope must not globally patch
|
||||
transformers' get_class_from_dynamic_module (issue #1751).
|
||||
|
||||
Reference in New Issue
Block a user