diff --git a/modelscope/utils/hf_util/patcher.py b/modelscope/utils/hf_util/patcher.py index 2404518c..790eb02c 100644 --- a/modelscope/utils/hf_util/patcher.py +++ b/modelscope/utils/hf_util/patcher.py @@ -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 diff --git a/tests/utils/test_hf_util.py b/tests/utils/test_hf_util.py index bf21317c..33e04ceb 100644 --- a/tests/utils/test_hf_util.py +++ b/tests/utils/test_hf_util.py @@ -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).