mirror of
https://github.com/modelscope/modelscope.git
synced 2026-08-29 10:08:40 +02:00
fix: forward download kwargs in HF patcher snapshot_download calls (#1762)
* fix: forward download kwargs in HF patcher snapshot_download calls Honor local_files_only, cache_dir, token, and revision/code_revision when resolving from_pretrained and cross-repo auto_map dynamic modules. Co-authored-by: Cursor <cursoragent@cursor.com> * style: fix yapf/flake8 formatting in hf util tests Co-authored-by: Cursor <cursoragent@cursor.com> * style: yapf format legacy cache reuse helpers Co-authored-by: Cursor <cursoragent@cursor.com> --------- Co-authored-by: suluyan <suluyan.sly@aliabab-inc.com> Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -251,8 +251,8 @@ def find_reusable_legacy_repo_dir(
|
|||||||
]
|
]
|
||||||
for path in legacy_candidates:
|
for path in legacy_candidates:
|
||||||
if _non_empty_dir(path):
|
if _non_empty_dir(path):
|
||||||
logger.info(
|
logger.info('Found legacy cache at %s for %s, reusing.', path,
|
||||||
'Found legacy cache at %s for %s, reusing.', path, repo_id)
|
repo_id)
|
||||||
return str(path)
|
return str(path)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|||||||
@@ -155,6 +155,31 @@ def _decide_allow_file_pattern(module_name, cls=None):
|
|||||||
return extra_allow_file_pattern
|
return extra_allow_file_pattern
|
||||||
|
|
||||||
|
|
||||||
|
def _ms_revision(revision):
|
||||||
|
"""Translate an HF revision string into one ModelScope accepts."""
|
||||||
|
return 'master' if revision in (None, 'main') else revision
|
||||||
|
|
||||||
|
|
||||||
|
def _ms_download_kwargs_from_hf(kwargs, revision=None):
|
||||||
|
"""Map transformers download kwargs onto ``snapshot_download`` arguments.
|
||||||
|
|
||||||
|
Forwards ``local_files_only``, ``cache_dir``, string ``token``, and an
|
||||||
|
optional revision (normalized via ``_ms_revision``).
|
||||||
|
"""
|
||||||
|
download_kwargs = {
|
||||||
|
'local_files_only': kwargs.get('local_files_only', False),
|
||||||
|
}
|
||||||
|
cache_dir = kwargs.get('cache_dir')
|
||||||
|
if cache_dir is not None:
|
||||||
|
download_kwargs['cache_dir'] = cache_dir
|
||||||
|
token = kwargs.get('token')
|
||||||
|
if isinstance(token, str):
|
||||||
|
download_kwargs['token'] = token
|
||||||
|
if revision is not None:
|
||||||
|
download_kwargs['revision'] = _ms_revision(revision)
|
||||||
|
return download_kwargs
|
||||||
|
|
||||||
|
|
||||||
def _get_class_from_dynamic_module(class_reference, *args, **kwargs):
|
def _get_class_from_dynamic_module(class_reference, *args, **kwargs):
|
||||||
"""Wrapper that redirects dynamic-module downloads to ModelScope.
|
"""Wrapper that redirects dynamic-module downloads to ModelScope.
|
||||||
|
|
||||||
@@ -189,7 +214,11 @@ def _get_class_from_dynamic_module(class_reference, *args, **kwargs):
|
|||||||
if (pretrained_model_name_or_path is not None
|
if (pretrained_model_name_or_path is not None
|
||||||
and not os.path.exists(pretrained_model_name_or_path)):
|
and not os.path.exists(pretrained_model_name_or_path)):
|
||||||
from modelscope import snapshot_download
|
from modelscope import snapshot_download
|
||||||
downloaded_path = snapshot_download(pretrained_model_name_or_path)
|
# Model weights/config: use ``revision`` (not ``code_revision``).
|
||||||
|
downloaded_path = snapshot_download(
|
||||||
|
pretrained_model_name_or_path,
|
||||||
|
**_ms_download_kwargs_from_hf(
|
||||||
|
kwargs, revision=kwargs.get('revision')))
|
||||||
if pretrained_in_kwargs:
|
if pretrained_in_kwargs:
|
||||||
kwargs['pretrained_model_name_or_path'] = downloaded_path
|
kwargs['pretrained_model_name_or_path'] = downloaded_path
|
||||||
else:
|
else:
|
||||||
@@ -199,7 +228,9 @@ def _get_class_from_dynamic_module(class_reference, *args, **kwargs):
|
|||||||
# Only the first ``--`` is the auto_map delimiter (repo vs module).
|
# Only the first ``--`` is the auto_map delimiter (repo vs module).
|
||||||
repo_id, class_reference = class_reference.split('--', 1)
|
repo_id, class_reference = class_reference.split('--', 1)
|
||||||
if not os.path.exists(repo_id):
|
if not os.path.exists(repo_id):
|
||||||
download_kwargs = {}
|
# Cross-repo code: transformers uses ``code_revision`` for this repo.
|
||||||
|
download_kwargs = _ms_download_kwargs_from_hf(
|
||||||
|
kwargs, revision=kwargs.get('code_revision'))
|
||||||
extra_allow_file_pattern = _decide_allow_file_pattern(
|
extra_allow_file_pattern = _decide_allow_file_pattern(
|
||||||
class_reference)
|
class_reference)
|
||||||
if extra_allow_file_pattern is not None:
|
if extra_allow_file_pattern is not None:
|
||||||
@@ -286,18 +317,16 @@ def _patch_pretrained_class(all_imported_modules, wrap=False):
|
|||||||
if subfolder:
|
if subfolder:
|
||||||
file_filter = f'{subfolder}/*'
|
file_filter = f'{subfolder}/*'
|
||||||
if not os.path.exists(pretrained_model_name_or_path):
|
if not os.path.exists(pretrained_model_name_or_path):
|
||||||
revision = kwargs.pop('revision', None)
|
revision = _ms_revision(kwargs.pop('revision', None))
|
||||||
if revision is None or revision == 'main':
|
|
||||||
revision = 'master'
|
|
||||||
if file_filter is not None:
|
if file_filter is not None:
|
||||||
allow_file_pattern = file_filter
|
allow_file_pattern = file_filter
|
||||||
local_files_only = kwargs.pop('local_files_only', False)
|
download_kwargs = _ms_download_kwargs_from_hf(
|
||||||
|
kwargs, revision=revision)
|
||||||
model_dir = snapshot_download(
|
model_dir = snapshot_download(
|
||||||
pretrained_model_name_or_path,
|
pretrained_model_name_or_path,
|
||||||
revision=revision,
|
|
||||||
local_files_only=local_files_only,
|
|
||||||
ignore_file_pattern=ignore_file_pattern,
|
ignore_file_pattern=ignore_file_pattern,
|
||||||
allow_file_pattern=allow_file_pattern)
|
allow_file_pattern=allow_file_pattern,
|
||||||
|
**download_kwargs)
|
||||||
if subfolder:
|
if subfolder:
|
||||||
model_dir = os.path.join(model_dir, subfolder)
|
model_dir = os.path.join(model_dir, subfolder)
|
||||||
else:
|
else:
|
||||||
@@ -630,11 +659,6 @@ def _unpatch_kernels():
|
|||||||
del kernels_utils._get_hf_api_origin
|
del kernels_utils._get_hf_api_origin
|
||||||
|
|
||||||
|
|
||||||
def _ms_revision(revision):
|
|
||||||
"""Translate an HF revision string into one ModelScope accepts."""
|
|
||||||
return 'master' if revision in (None, 'main') else revision
|
|
||||||
|
|
||||||
|
|
||||||
class _MsKernelApi:
|
class _MsKernelApi:
|
||||||
"""Minimal `HfApi` look-alike that forwards to ModelScope. Only the
|
"""Minimal `HfApi` look-alike that forwards to ModelScope. Only the
|
||||||
handful of methods that `kernels` actually calls are implemented.
|
handful of methods that `kernels` actually calls are implemented.
|
||||||
|
|||||||
@@ -75,7 +75,8 @@ class LegacyCacheReuseTest(unittest.TestCase):
|
|||||||
def test_uses_modelscope_cache_env(self):
|
def test_uses_modelscope_cache_env(self):
|
||||||
legacy = self.cache / self.owner / self.name
|
legacy = self.cache / self.owner / self.name
|
||||||
self._touch_model_dir(legacy)
|
self._touch_model_dir(legacy)
|
||||||
with mock.patch.dict(os.environ, {'MODELSCOPE_CACHE': str(self.cache)}):
|
with mock.patch.dict(os.environ,
|
||||||
|
{'MODELSCOPE_CACHE': str(self.cache)}):
|
||||||
found = find_reusable_legacy_repo_dir(self.model_id)
|
found = find_reusable_legacy_repo_dir(self.model_id)
|
||||||
self.assertEqual(found, str(legacy))
|
self.assertEqual(found, str(legacy))
|
||||||
|
|
||||||
|
|||||||
@@ -372,10 +372,131 @@ class HFUtilTest(unittest.TestCase):
|
|||||||
# item assignment.
|
# item assignment.
|
||||||
_get_class_from_dynamic_module('modeling.Foo', remote_id)
|
_get_class_from_dynamic_module('modeling.Foo', remote_id)
|
||||||
|
|
||||||
sd.assert_called_once_with(remote_id)
|
sd.assert_called_once_with(remote_id, local_files_only=False)
|
||||||
self.assertEqual(captured['class_reference'], 'modeling.Foo')
|
self.assertEqual(captured['class_reference'], 'modeling.Foo')
|
||||||
self.assertEqual(captured['pretrained'], downloaded)
|
self.assertEqual(captured['pretrained'], downloaded)
|
||||||
|
|
||||||
|
def test_dynamic_module_local_files_only_forwarded(self):
|
||||||
|
"""Download kwargs must be forwarded to both snapshot_download calls.
|
||||||
|
|
||||||
|
Cross-repo auto_map references previously omitted local_files_only /
|
||||||
|
cache_dir / token / code_revision, so offline and custom-cache loads
|
||||||
|
still hit the wrong download path for the referenced repo.
|
||||||
|
"""
|
||||||
|
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)
|
||||||
|
|
||||||
|
def fake_origin(class_reference, pretrained_model_name_or_path, *args,
|
||||||
|
**kwargs):
|
||||||
|
return type('DummyConfig', (), {})
|
||||||
|
|
||||||
|
remote_id = 'org/model-not-on-disk'
|
||||||
|
class_ref = 'org/other--configuration_foo.FooConfig'
|
||||||
|
call_kwargs = []
|
||||||
|
cache_dir = os.path.join(tmp, 'custom_cache')
|
||||||
|
token = 'ms-test-token'
|
||||||
|
|
||||||
|
def fake_download(repo_id, **kwargs):
|
||||||
|
call_kwargs.append((repo_id, dict(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):
|
||||||
|
_get_class_from_dynamic_module(
|
||||||
|
class_ref,
|
||||||
|
pretrained_model_name_or_path=remote_id,
|
||||||
|
local_files_only=True,
|
||||||
|
cache_dir=cache_dir,
|
||||||
|
token=token,
|
||||||
|
revision='model-rev',
|
||||||
|
code_revision='code-rev')
|
||||||
|
|
||||||
|
self.assertEqual(len(call_kwargs), 2)
|
||||||
|
self.assertEqual(call_kwargs[0][0], remote_id)
|
||||||
|
self.assertEqual(
|
||||||
|
call_kwargs[0][1], {
|
||||||
|
'local_files_only': True,
|
||||||
|
'cache_dir': cache_dir,
|
||||||
|
'token': token,
|
||||||
|
'revision': 'model-rev',
|
||||||
|
})
|
||||||
|
self.assertEqual(call_kwargs[1][0], 'org/other')
|
||||||
|
self.assertEqual(call_kwargs[1][1]['local_files_only'], True)
|
||||||
|
self.assertEqual(call_kwargs[1][1]['cache_dir'], cache_dir)
|
||||||
|
self.assertEqual(call_kwargs[1][1]['token'], token)
|
||||||
|
self.assertEqual(call_kwargs[1][1]['revision'], 'code-rev')
|
||||||
|
self.assertIn('ignore_file_pattern', call_kwargs[1][1])
|
||||||
|
|
||||||
|
def test_ms_download_kwargs_from_hf(self):
|
||||||
|
"""Shared HF→MS download kwargs mapping used by patcher download paths."""
|
||||||
|
from modelscope.utils.hf_util.patcher import _ms_download_kwargs_from_hf
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
_ms_download_kwargs_from_hf({}), {'local_files_only': False})
|
||||||
|
got = _ms_download_kwargs_from_hf(
|
||||||
|
{
|
||||||
|
'local_files_only': True,
|
||||||
|
'cache_dir': '/tmp/c',
|
||||||
|
'token': 'sekrit',
|
||||||
|
'token_ignored': True,
|
||||||
|
},
|
||||||
|
revision='main')
|
||||||
|
self.assertEqual(
|
||||||
|
got, {
|
||||||
|
'local_files_only': True,
|
||||||
|
'cache_dir': '/tmp/c',
|
||||||
|
'token': 'sekrit',
|
||||||
|
'revision': 'master',
|
||||||
|
})
|
||||||
|
# HF token=True means "default creds"; only string tokens are forwarded.
|
||||||
|
self.assertNotIn('token', _ms_download_kwargs_from_hf({'token': True}))
|
||||||
|
|
||||||
|
def test_get_model_dir_forwards_cache_dir_and_token(self):
|
||||||
|
"""from_pretrained download path must forward cache_dir and token."""
|
||||||
|
from unittest import mock
|
||||||
|
|
||||||
|
from modelscope import AutoConfig
|
||||||
|
|
||||||
|
tmp = tempfile.mkdtemp()
|
||||||
|
self.addCleanup(shutil.rmtree, tmp, ignore_errors=True)
|
||||||
|
with open(os.path.join(tmp, 'config.json'), 'w') as f:
|
||||||
|
f.write('{"model_type": "bert", "hidden_size": 8}')
|
||||||
|
|
||||||
|
cache_dir = os.path.join(tmp, 'custom_cache')
|
||||||
|
token = 'ms-test-token'
|
||||||
|
with mock.patch(
|
||||||
|
'modelscope.snapshot_download', return_value=tmp) as sd:
|
||||||
|
AutoConfig.from_pretrained(
|
||||||
|
'org/model-not-on-disk',
|
||||||
|
cache_dir=cache_dir,
|
||||||
|
token=token,
|
||||||
|
local_files_only=True)
|
||||||
|
|
||||||
|
sd.assert_called_once()
|
||||||
|
_, kwargs = sd.call_args
|
||||||
|
self.assertEqual(kwargs.get('cache_dir'), cache_dir)
|
||||||
|
self.assertEqual(kwargs.get('token'), token)
|
||||||
|
self.assertTrue(kwargs.get('local_files_only'))
|
||||||
|
|
||||||
def test_dynamic_module_pretrained_via_kwargs(self):
|
def test_dynamic_module_pretrained_via_kwargs(self):
|
||||||
"""pretrained_model_name_or_path may be passed as a keyword argument."""
|
"""pretrained_model_name_or_path may be passed as a keyword argument."""
|
||||||
from unittest import mock
|
from unittest import mock
|
||||||
|
|||||||
Reference in New Issue
Block a user