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:
suluyan
2026-07-20 16:48:39 +08:00
parent 3cd09acad5
commit b872cc83ac
2 changed files with 76 additions and 6 deletions

View File

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

View File

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