diff --git a/modelscope/hub/api.py b/modelscope/hub/api.py index 3a249f26..16440235 100644 --- a/modelscope/hub/api.py +++ b/modelscope/hub/api.py @@ -23,6 +23,7 @@ import json import requests from requests import Session from requests.adapters import HTTPAdapter, Retry +from requests.exceptions import HTTPError from tqdm.auto import tqdm from modelscope.hub.constants import (API_HTTP_CLIENT_MAX_RETRIES, @@ -34,10 +35,12 @@ from modelscope.hub.constants import (API_HTTP_CLIENT_MAX_RETRIES, API_RESPONSE_FIELD_USERNAME, DEFAULT_CREDENTIALS_PATH, DEFAULT_MAX_WORKERS, - DEFAULT_MODELSCOPE_DOMAIN, MODELSCOPE_CLOUD_ENVIRONMENT, MODELSCOPE_CLOUD_USERNAME, - MODELSCOPE_REQUEST_ID, ONE_YEAR_SECONDS, + MODELSCOPE_DOMAIN, + MODELSCOPE_PREFER_AI_SITE, + MODELSCOPE_REQUEST_ID, + MODELSCOPE_URL_SCHEME, ONE_YEAR_SECONDS, REQUESTS_API_HTTP_METHOD, TEMPORARY_FOLDER_NAME, DatasetVisibility, Licenses, ModelVisibility, Visibility, @@ -50,9 +53,9 @@ from modelscope.hub.errors import (InvalidParameter, NotExistError, raise_for_http_status, raise_on_error) from modelscope.hub.git import GitCommandWrapper from modelscope.hub.repository import Repository -from modelscope.hub.utils.utils import (add_content_to_file, get_endpoint, - get_readable_folder_size, - get_release_datetime, +from modelscope.hub.utils.utils import (add_content_to_file, get_domain, + get_endpoint, get_readable_folder_size, + get_release_datetime, is_env_true, model_id_to_group_owner_name) from modelscope.utils.constant import (DEFAULT_DATASET_REVISION, DEFAULT_MODEL_REVISION, @@ -118,14 +121,14 @@ class HubApi: jar = RequestsCookieJar() jar.set('m_session_id', access_token, - domain=os.getenv('MODELSCOPE_DOMAIN', - DEFAULT_MODELSCOPE_DOMAIN), + domain=get_domain(), path='/') return jar def login( self, - access_token: Optional[str] = None + access_token: Optional[str] = None, + endpoint: Optional[str] = None ): """Login with your SDK access token, which can be obtained from https://www.modelscope.cn user center. @@ -133,6 +136,7 @@ class HubApi: Args: access_token (str): user access token on modelscope, set this argument or set `MODELSCOPE_API_TOKEN`. If neither of the tokens exist, login will directly return. + endpoint: the endpoint to use, default to None to use endpoint specified in the class Returns: cookies: to authenticate yourself to ModelScope open-api @@ -145,7 +149,9 @@ class HubApi: access_token = os.environ.get('MODELSCOPE_API_TOKEN') if not access_token: return None, None - path = f'{self.endpoint}/api/v1/login' + if not endpoint: + endpoint = self.endpoint + path = f'{endpoint}/api/v1/login' r = self.session.post( path, json={'AccessToken': access_token}, @@ -172,7 +178,8 @@ class HubApi: visibility: Optional[int] = ModelVisibility.PUBLIC, license: Optional[str] = Licenses.APACHE_V2, chinese_name: Optional[str] = None, - original_model_id: Optional[str] = '') -> str: + original_model_id: Optional[str] = '', + endpoint: Optional[str] = None) -> str: """Create model repo at ModelScope Hub. Args: @@ -181,6 +188,7 @@ class HubApi: license (str, optional): license of the model, default none. chinese_name (str, optional): chinese name of the model. original_model_id (str, optional): the base model id which this model is trained from + endpoint: the endpoint to use, default to None to use endpoint specified in the class Returns: Name of the model created @@ -197,8 +205,9 @@ class HubApi: cookies = ModelScopeConfig.get_cookies() if cookies is None: raise ValueError('Token does not exist, please login first.') - - path = f'{self.endpoint}/api/v1/models' + if not endpoint: + endpoint = self.endpoint + path = f'{endpoint}/api/v1/models' owner_or_group, name = model_id_to_group_owner_name(model_id) body = { 'Path': owner_or_group, @@ -216,14 +225,15 @@ class HubApi: headers=self.builder_headers(self.headers)) handle_http_post_error(r, path, body) raise_on_error(r.json()) - model_repo_url = f'{self.endpoint}/{model_id}' + model_repo_url = f'{endpoint}/{model_id}' return model_repo_url - def delete_model(self, model_id: str): + def delete_model(self, model_id: str, endpoint: Optional[str] = None): """Delete model_id from ModelScope. Args: model_id (str): The model id. + endpoint: the endpoint to use, default to None to use endpoint specified in the class Raises: ValueError: If not login. @@ -232,9 +242,11 @@ class HubApi: model_id = {owner}/{name} """ cookies = ModelScopeConfig.get_cookies() + if not endpoint: + endpoint = self.endpoint if cookies is None: raise ValueError('Token does not exist, please login first.') - path = f'{self.endpoint}/api/v1/models/{model_id}' + path = f'{endpoint}/api/v1/models/{model_id}' r = self.session.delete(path, cookies=cookies, @@ -242,19 +254,23 @@ class HubApi: raise_for_http_status(r) raise_on_error(r.json()) - def get_model_url(self, model_id: str): - return f'{self.endpoint}/api/v1/models/{model_id}.git' + def get_model_url(self, model_id: str, endpoint: Optional[str] = None): + if not endpoint: + endpoint = self.endpoint + return f'{endpoint}/api/v1/models/{model_id}.git' def get_model( self, model_id: str, revision: Optional[str] = DEFAULT_MODEL_REVISION, + endpoint: Optional[str] = None ) -> str: """Get model information at ModelScope Args: model_id (str): The model id. revision (str optional): revision of model. + endpoint: the endpoint to use, default to None to use endpoint specified in the class Returns: The model detail information. @@ -267,10 +283,13 @@ class HubApi: """ cookies = ModelScopeConfig.get_cookies() owner_or_group, name = model_id_to_group_owner_name(model_id) + if not endpoint: + endpoint = self.endpoint + if revision: - path = f'{self.endpoint}/api/v1/models/{owner_or_group}/{name}?Revision={revision}' + path = f'{endpoint}/api/v1/models/{owner_or_group}/{name}?Revision={revision}' else: - path = f'{self.endpoint}/api/v1/models/{owner_or_group}/{name}' + path = f'{endpoint}/api/v1/models/{owner_or_group}/{name}' r = self.session.get(path, cookies=cookies, headers=self.builder_headers(self.headers)) @@ -283,11 +302,55 @@ class HubApi: else: raise_for_http_status(r) + def get_endpoint_for_read(self, + repo_id: str, + *, + repo_type: Optional[str] = None) -> str: + """Get proper endpoint for read operation (such as download, list etc.) + 1. If user has set MODELSCOPE_DOMAIN, construct endpoint with user-specified domain. + If the repo does not exist on that endpoint, throw 404 error, otherwise return the endpoint. + 2. If domain is not set, check existence of repo in cn-site and ai-site (intl version) respectively. + Checking order is determined by MODELSCOPE_PREFER_AI_SITE. + a. if MODELSCOPE_PREFER_AI_SITE is not set ,check cn-site first before ai-site (intl version) + b. otherwise check ai-site before cn-site + return the endpoint with which the given repo_id exists. + if neither exists, throw 404 error + """ + s = os.environ.get(MODELSCOPE_DOMAIN) + if s is not None and s.strip() != '': + endpoint = MODELSCOPE_URL_SCHEME + s + try: + self.repo_exists(repo_id=repo_id, repo_type=repo_type, endpoint=endpoint, re_raise=True) + except Exception: + logger.error(f'Repo {repo_id} does not exist on {endpoint}.') + raise + return endpoint + + check_cn_first = not is_env_true(MODELSCOPE_PREFER_AI_SITE) + prefer_endpoint = get_endpoint(cn_site=check_cn_first) + if not self.repo_exists( + repo_id, repo_type=repo_type, endpoint=prefer_endpoint): + alternative_endpoint = get_endpoint(cn_site=(not check_cn_first)) + logger.warning(f'Repo {repo_id} not exists on {prefer_endpoint}, ' + f'will try on alternative endpoint {alternative_endpoint}.') + try: + self.repo_exists( + repo_id, repo_type=repo_type, endpoint=alternative_endpoint, re_raise=True) + except Exception: + logger.error(f'Repo {repo_id} not exists on either {prefer_endpoint} or {alternative_endpoint}') + raise + else: + return alternative_endpoint + else: + return prefer_endpoint + def repo_exists( self, repo_id: str, *, repo_type: Optional[str] = None, + endpoint: Optional[str] = None, + re_raise: Optional[bool] = False ) -> bool: """ Checks if a repository exists on ModelScope @@ -299,18 +362,27 @@ class HubApi: repo_type (`str`, *optional*): `None` or `"model"` if getting repository info from a model. Default is `None`. TODO: support dataset and studio - + endpoint(`str`): + None or specific endpoint to use, when None, use the default endpoint + set in HubApi class (self.endpoint) + re_raise(`bool`): + raise exception when error Returns: True if the repository exists, False otherwise. """ - if (repo_type is not None) and repo_type.lower() != REPO_TYPE_MODEL: + if endpoint is None: + endpoint = self.endpoint + if (repo_type is not None) and repo_type.lower() not in REPO_TYPE_SUPPORT: raise Exception('Not support repo-type: %s' % repo_type) if (repo_id is None) or repo_id.count('/') != 1: raise Exception('Invalid repo_id: %s, must be of format namespace/name' % repo_type) cookies = ModelScopeConfig.get_cookies() owner_or_group, name = model_id_to_group_owner_name(repo_id) - path = f'{self.endpoint}/api/v1/models/{owner_or_group}/{name}' + if (repo_type is not None) and repo_type.lower() == REPO_TYPE_DATASET: + path = f'{endpoint}/api/v1/datasets/{owner_or_group}/{name}' + else: + path = f'{endpoint}/api/v1/models/{owner_or_group}/{name}' r = self.session.get(path, cookies=cookies, headers=self.builder_headers(self.headers)) @@ -318,7 +390,10 @@ class HubApi: if code == 200: return True elif code == 404: - return False + if re_raise: + raise HTTPError(r) + else: + return False else: logger.warn(f'Check repo_exists return status code {code}.') raise Exception( @@ -476,13 +551,15 @@ class HubApi: def list_models(self, owner_or_group: str, page_number: Optional[int] = 1, - page_size: Optional[int] = 10) -> dict: + page_size: Optional[int] = 10, + endpoint: Optional[str] = None) -> dict: """List models in owner or group. Args: owner_or_group(str): owner or group. page_number(int, optional): The page number, default: 1 page_size(int, optional): The page size, default: 10 + endpoint: the endpoint to use, default to None to use endpoint specified in the class Raises: RequestError: The request error. @@ -491,7 +568,9 @@ class HubApi: dict: {"models": "list of models", "TotalCount": total_number_of_models_in_owner_or_group} """ cookies = ModelScopeConfig.get_cookies() - path = f'{self.endpoint}/api/v1/models/' + if not endpoint: + endpoint = self.endpoint + path = f'{endpoint}/api/v1/models/' r = self.session.put( path, data='{"Path":"%s", "PageNumber":%s, "PageSize": %s}' % @@ -547,7 +626,8 @@ class HubApi: self, model_id: str, cutoff_timestamp: Optional[int] = None, - use_cookies: Union[bool, CookieJar] = False) -> List[str]: + use_cookies: Union[bool, CookieJar] = False, + endpoint: Optional[str] = None) -> List[str]: """Get model branch and tags. Args: @@ -556,6 +636,7 @@ class HubApi: The timestamp is represented by the seconds elapsed from the epoch time. use_cookies (Union[bool, CookieJar], optional): If is cookieJar, we will use this cookie, if True, will load cookie from local. Defaults to False. + endpoint: the endpoint to use, default to None to use endpoint specified in the class Returns: Tuple[List[str], List[str]]: Return list of branch name and tags @@ -563,7 +644,9 @@ class HubApi: cookies = self._check_cookie(use_cookies) if cutoff_timestamp is None: cutoff_timestamp = get_release_datetime() - path = f'{self.endpoint}/api/v1/models/{model_id}/revisions?EndTime=%s' % cutoff_timestamp + if not endpoint: + endpoint = self.endpoint + path = f'{endpoint}/api/v1/models/{model_id}/revisions?EndTime=%s' % cutoff_timestamp r = self.session.get(path, cookies=cookies, headers=self.builder_headers(self.headers)) handle_http_response(r, logger, cookies, model_id) @@ -582,21 +665,24 @@ class HubApi: def get_valid_revision_detail(self, model_id: str, revision=None, - cookies: Optional[CookieJar] = None): + cookies: Optional[CookieJar] = None, + endpoint: Optional[str] = None): + if not endpoint: + endpoint = self.endpoint release_timestamp = get_release_datetime() current_timestamp = int(round(datetime.datetime.now().timestamp())) # for active development in library codes (non-release-branches), release_timestamp # is set to be a far-away-time-in-the-future, to ensure that we shall # get the master-HEAD version from model repo by default (when no revision is provided) all_branches_detail, all_tags_detail = self.get_model_branches_and_tags_details( - model_id, use_cookies=False if cookies is None else cookies) + model_id, use_cookies=False if cookies is None else cookies, endpoint=endpoint) all_branches = [x['Revision'] for x in all_branches_detail] if all_branches_detail else [] all_tags = [x['Revision'] for x in all_tags_detail] if all_tags_detail else [] if release_timestamp > current_timestamp + ONE_YEAR_SECONDS: if revision is None: revision = MASTER_MODEL_BRANCH logger.info( - 'Model revision not specified, using default: [%s] version.' + 'Model revision not specified, using default [%s] version.' % revision) if revision not in all_branches and revision not in all_tags: raise NotExistError('The model: %s has no revision : %s .' % (model_id, revision)) @@ -649,15 +735,18 @@ class HubApi: def get_valid_revision(self, model_id: str, revision=None, - cookies: Optional[CookieJar] = None): + cookies: Optional[CookieJar] = None, + endpoint: Optional[str] = None): return self.get_valid_revision_detail(model_id=model_id, revision=revision, - cookies=cookies)['Revision'] + cookies=cookies, + endpoint=endpoint)['Revision'] def get_model_branches_and_tags_details( self, model_id: str, use_cookies: Union[bool, CookieJar] = False, + endpoint: Optional[str] = None ) -> Tuple[List[str], List[str]]: """Get model branch and tags. @@ -665,13 +754,15 @@ class HubApi: model_id (str): The model id use_cookies (Union[bool, CookieJar], optional): If is cookieJar, we will use this cookie, if True, will load cookie from local. Defaults to False. + endpoint: the endpoint to use, default to None to use endpoint specified in the class Returns: Tuple[List[str], List[str]]: Return list of branch name and tags """ cookies = self._check_cookie(use_cookies) - - path = f'{self.endpoint}/api/v1/models/{model_id}/revisions' + if not endpoint: + endpoint = self.endpoint + path = f'{endpoint}/api/v1/models/{model_id}/revisions' r = self.session.get(path, cookies=cookies, headers=self.builder_headers(self.headers)) handle_http_response(r, logger, cookies, model_id) @@ -709,7 +800,8 @@ class HubApi: root: Optional[str] = None, recursive: Optional[str] = False, use_cookies: Union[bool, CookieJar] = False, - headers: Optional[dict] = {}) -> List[dict]: + headers: Optional[dict] = {}, + endpoint: Optional[str] = None) -> List[dict]: """List the models files. Args: @@ -720,16 +812,19 @@ class HubApi: use_cookies (Union[bool, CookieJar], optional): If is cookieJar, we will use this cookie, if True, will load cookie from local. Defaults to False. headers: request headers + endpoint: the endpoint to use, default to None to use endpoint specified in the class Returns: List[dict]: Model file list. """ + if not endpoint: + endpoint = self.endpoint if revision: path = '%s/api/v1/models/%s/repo/files?Revision=%s&Recursive=%s' % ( - self.endpoint, model_id, revision, recursive) + endpoint, model_id, revision, recursive) else: path = '%s/api/v1/models/%s/repo/files?Recursive=%s' % ( - self.endpoint, model_id, recursive) + endpoint, model_id, recursive) cookies = self._check_cookie(use_cookies) if root is not None: path = path + f'&Root={root}' @@ -777,7 +872,8 @@ class HubApi: chinese_name: Optional[str] = '', license: Optional[str] = Licenses.APACHE_V2, visibility: Optional[int] = DatasetVisibility.PUBLIC, - description: Optional[str] = '') -> str: + description: Optional[str] = '', + endpoint: Optional[str] = None, ) -> str: if dataset_name is None or namespace is None: raise InvalidParameter('dataset_name and namespace are required!') @@ -785,8 +881,9 @@ class HubApi: cookies = ModelScopeConfig.get_cookies() if cookies is None: raise ValueError('Token does not exist, please login first.') - - path = f'{self.endpoint}/api/v1/datasets' + if not endpoint: + endpoint = self.endpoint + path = f'{endpoint}/api/v1/datasets' files = { 'Name': (None, dataset_name), 'ChineseName': (None, chinese_name), @@ -805,12 +902,14 @@ class HubApi: handle_http_post_error(r, path, files) raise_on_error(r.json()) - dataset_repo_url = f'{self.endpoint}/datasets/{namespace}/{dataset_name}' + dataset_repo_url = f'{endpoint}/datasets/{namespace}/{dataset_name}' logger.info(f'Create dataset success: {dataset_repo_url}') return dataset_repo_url - def list_datasets(self): - path = f'{self.endpoint}/api/v1/datasets' + def list_datasets(self, endpoint: Optional[str] = None): + if not endpoint: + endpoint = self.endpoint + path = f'{endpoint}/api/v1/datasets' params = {} r = self.session.get(path, params=params, headers=self.builder_headers(self.headers)) @@ -818,9 +917,11 @@ class HubApi: dataset_list = r.json()[API_RESPONSE_FIELD_DATA] return [x['Name'] for x in dataset_list] - def get_dataset_id_and_type(self, dataset_name: str, namespace: str): + def get_dataset_id_and_type(self, dataset_name: str, namespace: str, endpoint: Optional[str] = None): """ Get the dataset id and type. """ - datahub_url = f'{self.endpoint}/api/v1/datasets/{namespace}/{dataset_name}' + if not endpoint: + endpoint = self.endpoint + datahub_url = f'{endpoint}/api/v1/datasets/{namespace}/{dataset_name}' cookies = ModelScopeConfig.get_cookies() r = self.session.get(datahub_url, cookies=cookies) resp = r.json() @@ -834,11 +935,14 @@ class HubApi: revision: str, files_metadata: bool = False, timeout: float = 100, - recursive: str = 'True'): + recursive: str = 'True', + endpoint: Optional[str] = None): """ Get dataset infos. """ - datahub_url = f'{self.endpoint}/api/v1/datasets/{dataset_hub_id}/repo/tree' + if not endpoint: + endpoint = self.endpoint + datahub_url = f'{endpoint}/api/v1/datasets/{dataset_hub_id}/repo/tree' params = {'Revision': revision, 'Root': None, 'Recursive': recursive} cookies = ModelScopeConfig.get_cookies() if files_metadata: @@ -856,13 +960,16 @@ class HubApi: root_path: str, recursive: bool = True, page_number: int = 1, - page_size: int = 100): + page_size: int = 100, + endpoint: Optional[str] = None): dataset_hub_id, dataset_type = self.get_dataset_id_and_type( - dataset_name=dataset_name, namespace=namespace) + dataset_name=dataset_name, namespace=namespace, endpoint=endpoint) recursive = 'True' if recursive else 'False' - datahub_url = f'{self.endpoint}/api/v1/datasets/{dataset_hub_id}/repo/tree' + if not endpoint: + endpoint = self.endpoint + datahub_url = f'{endpoint}/api/v1/datasets/{dataset_hub_id}/repo/tree' params = {'Revision': revision if revision else 'master', 'Root': root_path if root_path else '/', 'Recursive': recursive, 'PageNumber': page_number, 'PageSize': page_size} @@ -874,9 +981,12 @@ class HubApi: return resp - def get_dataset_meta_file_list(self, dataset_name: str, namespace: str, dataset_id: str, revision: str): + def get_dataset_meta_file_list(self, dataset_name: str, namespace: str, + dataset_id: str, revision: str, endpoint: Optional[str] = None): """ Get the meta file-list of the dataset. """ - datahub_url = f'{self.endpoint}/api/v1/datasets/{dataset_id}/repo/tree?Revision={revision}' + if not endpoint: + endpoint = self.endpoint + datahub_url = f'{endpoint}/api/v1/datasets/{dataset_id}/repo/tree?Revision={revision}' cookies = ModelScopeConfig.get_cookies() r = self.session.get(datahub_url, cookies=cookies, @@ -908,7 +1018,8 @@ class HubApi: def get_dataset_meta_files_local_paths(self, dataset_name: str, namespace: str, revision: str, - meta_cache_dir: str, dataset_type: int, file_list: list): + meta_cache_dir: str, dataset_type: int, file_list: list, + endpoint: Optional[str] = None): local_paths = defaultdict(list) dataset_formation = DatasetFormations(dataset_type) dataset_meta_format = DatasetMetaFormats[dataset_formation] @@ -916,12 +1027,13 @@ class HubApi: # Dump the data_type as a local file HubApi.dump_datatype_file(dataset_type=dataset_type, meta_cache_dir=meta_cache_dir) - + if not endpoint: + endpoint = self.endpoint for file_info in file_list: file_path = file_info['Path'] extension = os.path.splitext(file_path)[-1] if extension in dataset_meta_format: - datahub_url = f'{self.endpoint}/api/v1/datasets/{namespace}/{dataset_name}/repo?' \ + datahub_url = f'{endpoint}/api/v1/datasets/{namespace}/{dataset_name}/repo?' \ f'Revision={revision}&FilePath={file_path}' r = self.session.get(datahub_url, cookies=cookies) raise_for_http_status(r) @@ -1001,7 +1113,8 @@ class HubApi: namespace: str, revision: Optional[str] = DEFAULT_DATASET_REVISION, view: Optional[bool] = False, - extension_filter: Optional[bool] = True): + extension_filter: Optional[bool] = True, + endpoint: Optional[str] = None): if not file_name or not dataset_name or not namespace: raise ValueError('Args (file_name, dataset_name, namespace) cannot be empty!') @@ -1009,7 +1122,9 @@ class HubApi: # Note: make sure the FilePath is the last parameter in the url params: dict = {'Source': 'SDK', 'Revision': revision, 'FilePath': file_name, 'View': view} params: str = urlencode(params) - file_url = f'{self.endpoint}/api/v1/datasets/{namespace}/{dataset_name}/repo?{params}' + if not endpoint: + endpoint = self.endpoint + file_url = f'{endpoint}/api/v1/datasets/{namespace}/{dataset_name}/repo?{params}' return file_url @@ -1028,9 +1143,12 @@ class HubApi: file_name: str, dataset_name: str, namespace: str, - revision: Optional[str] = DEFAULT_DATASET_REVISION): + revision: Optional[str] = DEFAULT_DATASET_REVISION, + endpoint: Optional[str] = None): + if not endpoint: + endpoint = self.endpoint if file_name and os.path.splitext(file_name)[-1] in META_FILES_FORMAT: - file_name = f'{self.endpoint}/api/v1/datasets/{namespace}/{dataset_name}/repo?' \ + file_name = f'{endpoint}/api/v1/datasets/{namespace}/{dataset_name}/repo?' \ f'Revision={revision}&FilePath={file_name}' return file_name @@ -1038,8 +1156,11 @@ class HubApi: self, dataset_name: str, namespace: str, - revision: Optional[str] = DEFAULT_DATASET_REVISION): - datahub_url = f'{self.endpoint}/api/v1/datasets/{namespace}/{dataset_name}/' \ + revision: Optional[str] = DEFAULT_DATASET_REVISION, + endpoint: Optional[str] = None): + if not endpoint: + endpoint = self.endpoint + datahub_url = f'{endpoint}/api/v1/datasets/{namespace}/{dataset_name}/' \ f'ststoken?Revision={revision}' return self.datahub_remote_call(datahub_url) @@ -1048,9 +1169,12 @@ class HubApi: dataset_name: str, namespace: str, check_cookie: bool, - revision: Optional[str] = DEFAULT_DATASET_REVISION): + revision: Optional[str] = DEFAULT_DATASET_REVISION, + endpoint: Optional[str] = None): - datahub_url = f'{self.endpoint}/api/v1/datasets/{namespace}/{dataset_name}/' \ + if not endpoint: + endpoint = self.endpoint + datahub_url = f'{endpoint}/api/v1/datasets/{namespace}/{dataset_name}/' \ f'ststoken?Revision={revision}' if check_cookie: cookies = self._check_cookie(use_cookies=True) @@ -1098,8 +1222,11 @@ class HubApi: dataset_name: str, namespace: str, revision: str, - zip_file_name: str): - datahub_url = f'{self.endpoint}/api/v1/datasets/{namespace}/{dataset_name}' + zip_file_name: str, + endpoint: Optional[str] = None): + if not endpoint: + endpoint = self.endpoint + datahub_url = f'{endpoint}/api/v1/datasets/{namespace}/{dataset_name}' cookies = ModelScopeConfig.get_cookies() r = self.session.get(url=datahub_url, cookies=cookies, headers=self.builder_headers(self.headers)) @@ -1120,8 +1247,10 @@ class HubApi: return data_sts def list_oss_dataset_objects(self, dataset_name, namespace, max_limit, - is_recursive, is_filter_dir, revision): - url = f'{self.endpoint}/api/v1/datasets/{namespace}/{dataset_name}/oss/tree/?' \ + is_recursive, is_filter_dir, revision, endpoint: Optional[str] = None): + if not endpoint: + endpoint = self.endpoint + url = f'{endpoint}/api/v1/datasets/{namespace}/{dataset_name}/oss/tree/?' \ f'MaxLimit={max_limit}&Revision={revision}&Recursive={is_recursive}&FilterDir={is_filter_dir}' cookies = ModelScopeConfig.get_cookies() @@ -1132,11 +1261,12 @@ class HubApi: return resp def delete_oss_dataset_object(self, object_name: str, dataset_name: str, - namespace: str, revision: str) -> str: + namespace: str, revision: str, endpoint: Optional[str] = None) -> str: if not object_name or not dataset_name or not namespace or not revision: raise ValueError('Args cannot be empty!') - - url = f'{self.endpoint}/api/v1/datasets/{namespace}/{dataset_name}/oss?Path={object_name}&Revision={revision}' + if not endpoint: + endpoint = self.endpoint + url = f'{endpoint}/api/v1/datasets/{namespace}/{dataset_name}/oss?Path={object_name}&Revision={revision}' cookies = ModelScopeConfig.get_cookies() resp = self.session.delete(url=url, cookies=cookies) @@ -1146,11 +1276,12 @@ class HubApi: return resp def delete_oss_dataset_dir(self, object_name: str, dataset_name: str, - namespace: str, revision: str) -> str: + namespace: str, revision: str, endpoint: Optional[str] = None) -> str: if not object_name or not dataset_name or not namespace or not revision: raise ValueError('Args cannot be empty!') - - url = f'{self.endpoint}/api/v1/datasets/{namespace}/{dataset_name}/oss/prefix?Prefix={object_name}/' \ + if not endpoint: + endpoint = self.endpoint + url = f'{endpoint}/api/v1/datasets/{namespace}/{dataset_name}/oss/prefix?Prefix={object_name}/' \ f'&Revision={revision}' cookies = ModelScopeConfig.get_cookies() @@ -1170,14 +1301,17 @@ class HubApi: datahub_raise_on_error(url, resp, r) return resp['Data'] - def dataset_download_statistics(self, dataset_name: str, namespace: str, use_streaming: bool = False) -> None: + def dataset_download_statistics(self, dataset_name: str, namespace: str, + use_streaming: bool = False, endpoint: Optional[str] = None) -> None: is_ci_test = os.getenv('CI_TEST') == 'True' + if not endpoint: + endpoint = self.endpoint if dataset_name and namespace and not is_ci_test and not use_streaming: try: cookies = ModelScopeConfig.get_cookies() # Download count - download_count_url = f'{self.endpoint}/api/v1/datasets/{namespace}/{dataset_name}/download/increase' + download_count_url = f'{endpoint}/api/v1/datasets/{namespace}/{dataset_name}/download/increase' download_count_resp = self.session.post(download_count_url, cookies=cookies, headers=self.builder_headers(self.headers)) raise_for_http_status(download_count_resp) @@ -1189,7 +1323,7 @@ class HubApi: channel = os.environ[MODELSCOPE_CLOUD_ENVIRONMENT] if MODELSCOPE_CLOUD_USERNAME in os.environ: user_name = os.environ[MODELSCOPE_CLOUD_USERNAME] - download_uv_url = f'{self.endpoint}/api/v1/datasets/{namespace}/{dataset_name}/download/uv/' \ + download_uv_url = f'{endpoint}/api/v1/datasets/{namespace}/{dataset_name}/download/uv/' \ f'{channel}?user={user_name}' download_uv_resp = self.session.post(download_uv_url, cookies=cookies, headers=self.builder_headers(self.headers)) @@ -1203,9 +1337,11 @@ class HubApi: return {MODELSCOPE_REQUEST_ID: str(uuid.uuid4().hex), **headers} - def get_file_base_path(self, repo_id: str) -> str: + def get_file_base_path(self, repo_id: str, endpoint: Optional[str] = None) -> str: _namespace, _dataset_name = repo_id.split('/') - return f'{self.endpoint}/api/v1/datasets/{_namespace}/{_dataset_name}/repo?' + if not endpoint: + endpoint = self.endpoint + return f'{endpoint}/api/v1/datasets/{_namespace}/{_dataset_name}/repo?' # return f'{endpoint}/api/v1/datasets/{namespace}/{dataset_name}/repo?Revision={revision}&FilePath=' def create_repo( @@ -1217,13 +1353,15 @@ class HubApi: repo_type: Optional[str] = REPO_TYPE_MODEL, chinese_name: Optional[str] = '', license: Optional[str] = Licenses.APACHE_V2, + endpoint: Optional[str] = None, **kwargs, ) -> str: # TODO: exist_ok if not repo_id: raise ValueError('Repo id cannot be empty!') - + if not endpoint: + endpoint = self.endpoint self.login(access_token=token) repo_id_list = repo_id.split('/') @@ -1261,7 +1399,7 @@ class HubApi: 'configuration.json', [json.dumps(config)], ignore_push_error=True) else: - repo_url = f'{self.endpoint}/{repo_id}' + repo_url = f'{endpoint}/{repo_id}' elif repo_type == REPO_TYPE_DATASET: visibilities = {k: v for k, v in DatasetVisibility.__dict__.items() if not k.startswith('__')} @@ -1278,7 +1416,7 @@ class HubApi: visibility=visibility, ) else: - repo_url = f'{self.endpoint}/datasets/{namespace}/{repo_name}' + repo_url = f'{endpoint}/datasets/{namespace}/{repo_name}' else: raise ValueError(f'Invalid repo type: {repo_type}, supported repos: {REPO_TYPE_SUPPORT}') @@ -1295,9 +1433,12 @@ class HubApi: token: str = None, repo_type: Optional[str] = None, revision: Optional[str] = DEFAULT_REPOSITORY_REVISION, + endpoint: Optional[str] = None ) -> CommitInfo: - url = f'{self.endpoint}/api/v1/repos/{repo_type}s/{repo_id}/commit/{revision}' + if not endpoint: + endpoint = self.endpoint + url = f'{endpoint}/api/v1/repos/{repo_type}s/{repo_id}/commit/{revision}' commit_message = commit_message or f'Commit to {repo_id}' commit_description = commit_description or '' @@ -1640,6 +1781,7 @@ class HubApi: repo_id: str, repo_type: str, objects: List[Dict[str, Any]], + endpoint: Optional[str] = None ) -> List[Dict[str, Any]]: """ Check the blob has already uploaded. @@ -1651,13 +1793,16 @@ class HubApi: objects (List[Dict[str, Any]]): The objects to check. oid (str): The sha256 hash value. size (int): The size of the blob. + endpoint: the endpoint to use, default to None to use endpoint specified in the class Returns: List[Dict[str, Any]]: The result of the check. """ # construct URL - url = f'{self.endpoint}/api/v1/repos/{repo_type}s/{repo_id}/info/lfs/objects/batch' + if not endpoint: + endpoint = self.endpoint + url = f'{endpoint}/api/v1/repos/{repo_type}s/{repo_id}/info/lfs/objects/batch' # build payload payload = { @@ -1839,8 +1984,8 @@ class ModelScopeConfig: if cookie.name == 'm_session_id' and cookie.is_expired() and \ not ModelScopeConfig.cookie_expired_warning: ModelScopeConfig.cookie_expired_warning = True - logger.warning('Authentication has expired, ' - 'please re-login for uploading or accessing controlled entities.') + logger.info('Not logged-in, you can login for uploading' + 'or accessing controlled entities.') return None return cookies return None diff --git a/modelscope/hub/constants.py b/modelscope/hub/constants.py index 64b517c0..ea086be7 100644 --- a/modelscope/hub/constants.py +++ b/modelscope/hub/constants.py @@ -4,7 +4,9 @@ from pathlib import Path MODELSCOPE_URL_SCHEME = 'https://' DEFAULT_MODELSCOPE_DOMAIN = 'www.modelscope.cn' +DEFAULT_MODELSCOPE_INTL_DOMAIN = 'www.modelscope.ai' DEFAULT_MODELSCOPE_DATA_ENDPOINT = MODELSCOPE_URL_SCHEME + DEFAULT_MODELSCOPE_DOMAIN +DEFAULT_MODELSCOPE_INTL_DATA_ENDPOINT = MODELSCOPE_URL_SCHEME + DEFAULT_MODELSCOPE_INTL_DOMAIN MODELSCOPE_PARALLEL_DOWNLOAD_THRESHOLD_MB = int( os.environ.get('MODELSCOPE_PARALLEL_DOWNLOAD_THRESHOLD_MB', 500)) MODELSCOPE_DOWNLOAD_PARALLELS = int( @@ -28,6 +30,8 @@ API_RESPONSE_FIELD_MESSAGE = 'Message' MODELSCOPE_CLOUD_ENVIRONMENT = 'MODELSCOPE_ENVIRONMENT' MODELSCOPE_CLOUD_USERNAME = 'MODELSCOPE_USERNAME' MODELSCOPE_SDK_DEBUG = 'MODELSCOPE_SDK_DEBUG' +MODELSCOPE_PREFER_AI_SITE = 'MODELSCOPE_PREFER_AI_SITE' +MODELSCOPE_DOMAIN = 'MODELSCOPE_DOMAIN' MODELSCOPE_ENABLE_DEFAULT_HASH_VALIDATION = 'MODELSCOPE_ENABLE_DEFAULT_HASH_VALIDATION' ONE_YEAR_SECONDS = 24 * 365 * 60 * 60 MODELSCOPE_REQUEST_ID = 'X-Request-ID' diff --git a/modelscope/hub/file_download.py b/modelscope/hub/file_download.py index ee0f5d89..322c55bc 100644 --- a/modelscope/hub/file_download.py +++ b/modelscope/hub/file_download.py @@ -199,17 +199,19 @@ def _repo_file_download( if cookies is None: cookies = ModelScopeConfig.get_cookies() repo_files = [] + endpoint = _api.get_endpoint_for_read(repo_id=repo_id, repo_type=repo_type) file_to_download_meta = None if repo_type == REPO_TYPE_MODEL: revision = _api.get_valid_revision( - repo_id, revision=revision, cookies=cookies) + repo_id, revision=revision, cookies=cookies, endpoint=endpoint) # we need to confirm the version is up-to-date # we need to get the file list to check if the latest version is cached, if so return, otherwise download repo_files = _api.get_model_files( model_id=repo_id, revision=revision, recursive=True, - use_cookies=False if cookies is None else cookies) + use_cookies=False if cookies is None else cookies, + endpoint=endpoint) for repo_file in repo_files: if repo_file['Type'] == 'tree': continue @@ -238,7 +240,8 @@ def _repo_file_download( root_path='/', recursive=True, page_number=page_number, - page_size=page_size) + page_size=page_size, + endpoint=endpoint) if not ('Code' in files_list_tree and files_list_tree['Code'] == 200): print( @@ -273,13 +276,15 @@ def _repo_file_download( # we need to download again if repo_type == REPO_TYPE_MODEL: - url_to_download = get_file_download_url(repo_id, file_path, revision) + url_to_download = get_file_download_url(repo_id, file_path, revision, + endpoint) elif repo_type == REPO_TYPE_DATASET: url_to_download = _api.get_dataset_file_url( file_name=file_to_download_meta['Path'], dataset_name=name, namespace=group_or_owner, - revision=revision) + revision=revision, + endpoint=endpoint) else: raise ValueError(f'Invalid repo type {repo_type}') @@ -354,7 +359,10 @@ def create_temporary_directory_and_cache(model_id: str, return temporary_cache_dir, cache -def get_file_download_url(model_id: str, file_path: str, revision: str): +def get_file_download_url(model_id: str, + file_path: str, + revision: str, + endpoint: Optional[str] = None): """Format file download url according to `model_id`, `revision` and `file_path`. e.g., Given `model_id=john/bert`, `revision=master`, `file_path=README.md`, the resulted download url is: https://modelscope.cn/api/v1/models/john/bert/repo?Revision=master&FilePath=README.md @@ -363,6 +371,7 @@ def get_file_download_url(model_id: str, file_path: str, revision: str): model_id (str): The model_id. file_path (str): File path revision (str): File revision. + endpoint (str): The remote endpoint Returns: str: The file url. @@ -370,8 +379,10 @@ def get_file_download_url(model_id: str, file_path: str, revision: str): file_path = urllib.parse.quote_plus(file_path) revision = urllib.parse.quote_plus(revision) download_url_template = '{endpoint}/api/v1/models/{model_id}/repo?Revision={revision}&FilePath={file_path}' + if not endpoint: + endpoint = get_endpoint() return download_url_template.format( - endpoint=get_endpoint(), + endpoint=endpoint, model_id=model_id, revision=revision, file_path=file_path, @@ -420,15 +431,14 @@ def download_part_with_retry(params): retry.sleep() -def parallel_download( - url: str, - local_dir: str, - file_name: str, - cookies: CookieJar, - headers: Optional[Dict[str, str]] = None, - file_size: int = None, - disable_tqdm: bool = False, -): +def parallel_download(url: str, + local_dir: str, + file_name: str, + cookies: CookieJar, + headers: Optional[Dict[str, str]] = None, + file_size: int = None, + disable_tqdm: bool = False, + endpoint: str = None): # create temp file with tqdm( unit='B', diff --git a/modelscope/hub/snapshot_download.py b/modelscope/hub/snapshot_download.py index 8923e9e3..3585b5e5 100644 --- a/modelscope/hub/snapshot_download.py +++ b/modelscope/hub/snapshot_download.py @@ -241,6 +241,8 @@ def _snapshot_download( 'snapshot-identifier': str(uuid.uuid4()), } _api = HubApi() + endpoint = _api.get_endpoint_for_read( + repo_id=repo_id, repo_type=repo_type) if cookies is None: cookies = ModelScopeConfig.get_cookies() if repo_type == REPO_TYPE_MODEL: @@ -251,9 +253,10 @@ def _snapshot_download( else: directory = os.path.join(system_cache, 'models', *repo_id.split('/')) - print(f'Downloading Model to directory: {directory}') + print( + f'Downloading Model from {endpoint} to directory: {directory}') revision_detail = _api.get_valid_revision_detail( - repo_id, revision=revision, cookies=cookies) + repo_id, revision=revision, cookies=cookies, endpoint=endpoint) revision = revision_detail['Revision'] # Add snapshot-ci-test for counting the ci test download @@ -272,7 +275,7 @@ def _snapshot_download( recursive=True, use_cookies=False if cookies is None else cookies, headers=snapshot_header, - ) + endpoint=endpoint) _download_file_lists( repo_files, cache, @@ -289,7 +292,9 @@ def _snapshot_download( allow_file_pattern=allow_file_pattern, ignore_patterns=ignore_patterns, allow_patterns=allow_patterns, - max_workers=max_workers) + max_workers=max_workers, + endpoint=endpoint, + ) if '.' in repo_id: masked_directory = get_model_masked_directory( directory, repo_id) @@ -322,7 +327,7 @@ def _snapshot_download( logger.info('Fetching dataset repo file list...') repo_files = fetch_repo_files(_api, name, group_or_owner, - revision_detail) + revision_detail, endpoint) if repo_files is None: logger.error( @@ -345,14 +350,16 @@ def _snapshot_download( allow_file_pattern=allow_file_pattern, ignore_patterns=ignore_patterns, allow_patterns=allow_patterns, - max_workers=max_workers) + max_workers=max_workers, + endpoint=endpoint, + ) cache.save_model_version(revision_info=revision_detail) cache_root_path = cache.get_root_location() return cache_root_path -def fetch_repo_files(_api, name, group_or_owner, revision): +def fetch_repo_files(_api, name, group_or_owner, revision, endpoint): page_number = 1 page_size = 150 repo_files = [] @@ -365,7 +372,8 @@ def fetch_repo_files(_api, name, group_or_owner, revision): root_path='/', recursive=True, page_number=page_number, - page_size=page_size) + page_size=page_size, + endpoint=endpoint) if not ('Code' in files_list_tree and files_list_tree['Code'] == 200): logger.error(f'Get dataset file list failed, request_id: \ @@ -414,22 +422,24 @@ def _get_valid_regex_pattern(patterns: List[str]): def _download_file_lists( - repo_files: List[str], - cache: ModelFileSystemCache, - temporary_cache_dir: str, - repo_id: str, - api: HubApi, - name: str, - group_or_owner: str, - headers, - repo_type: Optional[str] = None, - revision: Optional[str] = DEFAULT_MODEL_REVISION, - cookies: Optional[CookieJar] = None, - ignore_file_pattern: Optional[Union[str, List[str]]] = None, - allow_file_pattern: Optional[Union[str, List[str]]] = None, - allow_patterns: Optional[Union[List[str], str]] = None, - ignore_patterns: Optional[Union[List[str], str]] = None, - max_workers: int = 8): + repo_files: List[str], + cache: ModelFileSystemCache, + temporary_cache_dir: str, + repo_id: str, + api: HubApi, + name: str, + group_or_owner: str, + headers, + repo_type: Optional[str] = None, + revision: Optional[str] = DEFAULT_MODEL_REVISION, + cookies: Optional[CookieJar] = None, + ignore_file_pattern: Optional[Union[str, List[str]]] = None, + allow_file_pattern: Optional[Union[str, List[str]]] = None, + allow_patterns: Optional[Union[List[str], str]] = None, + ignore_patterns: Optional[Union[List[str], str]] = None, + max_workers: int = 8, + endpoint: Optional[str] = None, +): ignore_patterns = _normalize_patterns(ignore_patterns) allow_patterns = _normalize_patterns(allow_patterns) ignore_file_pattern = _normalize_patterns(ignore_file_pattern) @@ -490,13 +500,15 @@ def _download_file_lists( url = get_file_download_url( model_id=repo_id, file_path=repo_file['Path'], - revision=revision) + revision=revision, + endpoint=endpoint) elif repo_type == REPO_TYPE_DATASET: url = api.get_dataset_file_url( file_name=repo_file['Path'], dataset_name=name, namespace=group_or_owner, - revision=revision) + revision=revision, + endpoint=endpoint) else: raise InvalidParameter( f'Invalid repo type: {repo_type}, supported types: {REPO_TYPE_SUPPORT}' diff --git a/modelscope/hub/utils/utils.py b/modelscope/hub/utils/utils.py index 0fd078b0..2cb1bb33 100644 --- a/modelscope/hub/utils/utils.py +++ b/modelscope/hub/utils/utils.py @@ -8,7 +8,9 @@ from typing import List, Optional, Union from modelscope.hub.constants import (DEFAULT_MODELSCOPE_DOMAIN, DEFAULT_MODELSCOPE_GROUP, - MODEL_ID_SEPARATOR, MODELSCOPE_SDK_DEBUG, + DEFAULT_MODELSCOPE_INTL_DOMAIN, + MODEL_ID_SEPARATOR, MODELSCOPE_DOMAIN, + MODELSCOPE_SDK_DEBUG, MODELSCOPE_URL_SCHEME) from modelscope.hub.errors import FileIntegrityError from modelscope.utils.logger import get_logger @@ -26,6 +28,20 @@ def model_id_to_group_owner_name(model_id): return group_or_owner, name +def is_env_true(var_name): + value = os.environ.get(var_name, '').strip().lower() + return value == 'true' + + +def get_domain(cn_site=True): + if MODELSCOPE_DOMAIN in os.environ and os.getenv(MODELSCOPE_DOMAIN): + return os.getenv(MODELSCOPE_DOMAIN) + if cn_site: + return DEFAULT_MODELSCOPE_DOMAIN + else: + return DEFAULT_MODELSCOPE_INTL_DOMAIN + + def convert_patterns(raw_input: Union[str, List[str]]): output = None if isinstance(raw_input, str): @@ -105,11 +121,8 @@ def get_release_datetime(): return rt -def get_endpoint(): - modelscope_domain = os.getenv( - 'MODELSCOPE_DOMAIN', - DEFAULT_MODELSCOPE_DOMAIN) or DEFAULT_MODELSCOPE_DOMAIN - return MODELSCOPE_URL_SCHEME + modelscope_domain +def get_endpoint(cn_site=True): + return MODELSCOPE_URL_SCHEME + get_domain(cn_site) def compute_hash(file_path): diff --git a/modelscope/models/cv/image_try_on/landmark.py b/modelscope/models/cv/image_try_on/landmark.py index 489e59c3..e74d53c9 100644 --- a/modelscope/models/cv/image_try_on/landmark.py +++ b/modelscope/models/cv/image_try_on/landmark.py @@ -369,7 +369,7 @@ class VTONLandmark(nn.Module): 'SHIFT_HEATMAP': True }, 'DEBUG': { - 'DEBUG': True, + 'DEBUG': False, 'SAVE_BATCH_IMAGES_GT': True, 'SAVE_BATCH_IMAGES_PRED': True, 'SAVE_HEATMAPS_GT': True, diff --git a/modelscope/models/cv/video_depth_estimation/utils/load.py b/modelscope/models/cv/video_depth_estimation/utils/load.py index 8c2b326c..06a12e6d 100644 --- a/modelscope/models/cv/video_depth_estimation/utils/load.py +++ b/modelscope/models/cv/video_depth_estimation/utils/load.py @@ -1,9 +1,6 @@ # Part of the implementation is borrowed and modified from PackNet-SfM, # made publicly available under the MIT License at https://github.com/TRI-ML/packnet-sfm import importlib -import logging -import os -import warnings from collections import OrderedDict from inspect import signature @@ -16,23 +13,6 @@ from modelscope.models.cv.video_depth_estimation.utils.misc import (make_list, from modelscope.models.cv.video_depth_estimation.utils.types import is_str -def set_debug(debug): - """ - Enable or disable debug terminal logging - - Parameters - ---------- - debug : bool - Debugging flag (True to enable) - """ - # Disable logging if requested - if not debug: - os.environ['NCCL_DEBUG'] = '' - os.environ['WANDB_SILENT'] = 'false' - warnings.filterwarnings('ignore') - logging.disable(logging.CRITICAL) - - def filter_args(func, keys): """ Filters a dictionary so it only contains keys that are arguments of a function diff --git a/modelscope/msdatasets/meta/data_meta_manager.py b/modelscope/msdatasets/meta/data_meta_manager.py index e5a57f02..afef97b0 100644 --- a/modelscope/msdatasets/meta/data_meta_manager.py +++ b/modelscope/msdatasets/meta/data_meta_manager.py @@ -13,8 +13,8 @@ from modelscope.msdatasets.context.dataset_context_config import \ from modelscope.msdatasets.meta.data_meta_config import DataMetaConfig from modelscope.msdatasets.utils.dataset_utils import ( get_dataset_files, get_target_dataset_structure) -from modelscope.utils.constant import (DatasetFormations, DatasetPathName, - DownloadMode) +from modelscope.utils.constant import (REPO_TYPE_DATASET, DatasetFormations, + DatasetPathName, DownloadMode) class DataMetaManager(object): @@ -177,9 +177,13 @@ class DataMetaManager(object): def _fetch_meta_from_hub(self, dataset_name: str, namespace: str, revision: str, meta_cache_dir: str): + _api = HubApi() + endpoint = _api.get_endpoint_for_read( + repo_id=namespace + '/' + dataset_name, + repo_type=REPO_TYPE_DATASET) # Fetch id and type of dataset dataset_id, dataset_type = self.api.get_dataset_id_and_type( - dataset_name, namespace) + dataset_name, namespace, endpoint) # Fetch meta file-list of dataset file_list = self.api.get_dataset_meta_file_list( diff --git a/modelscope/msdatasets/ms_dataset.py b/modelscope/msdatasets/ms_dataset.py index 21599a1b..ec7028f7 100644 --- a/modelscope/msdatasets/ms_dataset.py +++ b/modelscope/msdatasets/ms_dataset.py @@ -28,7 +28,8 @@ from modelscope.preprocessors import build_preprocessor from modelscope.utils.config import Config, ConfigDict from modelscope.utils.config_ds import MS_DATASETS_CACHE from modelscope.utils.constant import (DEFAULT_DATASET_NAMESPACE, - DEFAULT_DATASET_REVISION, ConfigFields, + DEFAULT_DATASET_REVISION, + REPO_TYPE_DATASET, ConfigFields, DatasetFormations, DownloadMode, Hubs, ModeKeys, Tasks, UploadMode) from modelscope.utils.import_utils import is_tf_available, is_torch_available @@ -290,12 +291,16 @@ class MsDataset: # Load from the modelscope hub elif hub == Hubs.modelscope: - # Get dataset type from ModelScope Hub; dataset_type->4: General Dataset from modelscope.hub.api import HubApi _api = HubApi() + endpoint = _api.get_endpoint_for_read( + repo_id=namespace + '/' + dataset_name, + repo_type=REPO_TYPE_DATASET) dataset_id_on_hub, dataset_type = _api.get_dataset_id_and_type( - dataset_name=dataset_name, namespace=namespace) + dataset_name=dataset_name, + namespace=namespace, + endpoint=endpoint) # Load from the ModelScope Hub for type=4 (general) if str(dataset_type) == str(DatasetFormations.general.value): diff --git a/modelscope/msdatasets/utils/hf_datasets_util.py b/modelscope/msdatasets/utils/hf_datasets_util.py index fea304f6..031ab63f 100644 --- a/modelscope/msdatasets/utils/hf_datasets_util.py +++ b/modelscope/msdatasets/utils/hf_datasets_util.py @@ -62,7 +62,7 @@ from modelscope import HubApi from modelscope.hub.utils.utils import get_endpoint from modelscope.msdatasets.utils.hf_file_utils import get_from_cache_ms from modelscope.utils.config_ds import MS_DATASETS_CACHE -from modelscope.utils.constant import DEFAULT_DATASET_NAMESPACE, DEFAULT_DATASET_REVISION +from modelscope.utils.constant import DEFAULT_DATASET_NAMESPACE, DEFAULT_DATASET_REVISION, REPO_TYPE_DATASET from modelscope.utils.import_utils import has_attr_in_class from modelscope.utils.logger import get_logger @@ -160,14 +160,17 @@ def _dataset_info( """ _api = HubApi() _namespace, _dataset_name = repo_id.split('/') + endpoint = _api.get_endpoint_for_read( + repo_id=repo_id, repo_type=REPO_TYPE_DATASET) dataset_hub_id, dataset_type = _api.get_dataset_id_and_type( - dataset_name=_dataset_name, namespace=_namespace) + dataset_name=_dataset_name, namespace=_namespace, endpoint=endpoint) revision: str = revision or DEFAULT_DATASET_REVISION data = _api.get_dataset_infos(dataset_hub_id=dataset_hub_id, revision=revision, files_metadata=files_metadata, - timeout=timeout) + timeout=timeout, + endpoint=endpoint) # Parse data data_d: dict = data['Data'] @@ -220,7 +223,8 @@ def _list_repo_tree( ) -> Iterable[Union[RepoFile, RepoFolder]]: _api = HubApi(timeout=3 * 60, max_retries=3) - + endpoint = _api.get_endpoint_for_read( + repo_id=repo_id, repo_type=REPO_TYPE_DATASET) if is_relative_path(repo_id) and repo_id.count('/') == 1: _namespace, _dataset_name = repo_id.split('/') elif is_relative_path(repo_id) and repo_id.count('/') == 0: @@ -240,6 +244,7 @@ def _list_repo_tree( recursive=True, page_number=page_number, page_size=page_size, + endpoint=endpoint ) if not ('Code' in data and data['Code'] == 200): logger.error(f'Get dataset: {repo_id} file list failed, message: {data["Message"]}') @@ -275,8 +280,10 @@ def _get_paths_info( _api = HubApi() _namespace, _dataset_name = repo_id.split('/') + endpoint = _api.get_endpoint_for_read( + repo_id=repo_id, repo_type=REPO_TYPE_DATASET) dataset_hub_id, dataset_type = _api.get_dataset_id_and_type( - dataset_name=_dataset_name, namespace=_namespace) + dataset_name=_dataset_name, namespace=_namespace, endpoint=endpoint) revision: str = revision or DEFAULT_DATASET_REVISION data = _api.get_dataset_infos(dataset_hub_id=dataset_hub_id, @@ -300,7 +307,8 @@ def _get_paths_info( def _download_repo_file(repo_id: str, path_in_repo: str, download_config: DownloadConfig, revision: str): _api = HubApi() _namespace, _dataset_name = repo_id.split('/') - + endpoint = _api.get_endpoint_for_read( + repo_id=repo_id, repo_type=REPO_TYPE_DATASET) if download_config and download_config.download_desc is None: download_config.download_desc = f'Downloading [{path_in_repo}]' try: @@ -310,6 +318,7 @@ def _download_repo_file(repo_id: str, path_in_repo: str, download_config: Downlo namespace=_namespace, revision=revision, extension_filter=False, + endpoint=endpoint ) repo_file_path = cached_path( url_or_filename=url_or_filename, download_config=download_config) @@ -656,10 +665,14 @@ def get_module_without_script(self) -> DatasetModule: ) ] default_config_name = None + _api = HubApi() + endpoint = _api.get_endpoint_for_read( + repo_id=repo_id, repo_type=REPO_TYPE_DATASET) + builder_kwargs = { # "base_path": hf_hub_url(self.name, "", revision=revision).rstrip("/"), 'base_path': - HubApi().get_file_base_path(repo_id=repo_id), + HubApi().get_file_base_path(repo_id=repo_id, endpoint=endpoint), 'repo_id': self.name, 'dataset_name': @@ -1021,9 +1034,12 @@ class DatasetsWrapperHF: try: _api = HubApi() + if is_relative_path(path) and path.count('/') == 1: _namespace, _dataset_name = path.split('/') - _api.dataset_download_statistics(dataset_name=_dataset_name, namespace=_namespace) + endpoint = _api.get_endpoint_for_read( + repo_id=path, repo_type=REPO_TYPE_DATASET) + _api.dataset_download_statistics(dataset_name=_dataset_name, namespace=_namespace, endpoint=endpoint) except Exception as e: logger.warning(f'Could not record download statistics: {e}') diff --git a/modelscope/utils/logger.py b/modelscope/utils/logger.py index bc471044..80f98730 100644 --- a/modelscope/utils/logger.py +++ b/modelscope/utils/logger.py @@ -11,6 +11,7 @@ formatter = logging.Formatter( '%(asctime)s - %(name)s - %(levelname)s - %(message)s') default_log_level = int(os.getenv('MODELSCOPE_LOG_LEVEL', str(logging.INFO))) +logging.getLogger('numba').setLevel(logging.INFO) def get_logger(log_file: Optional[str] = None, diff --git a/tests/hub/test_ai_site.py b/tests/hub/test_ai_site.py new file mode 100644 index 00000000..57aaa88a --- /dev/null +++ b/tests/hub/test_ai_site.py @@ -0,0 +1,67 @@ +# Copyright (c) Alibaba, Inc. and its affiliates. +import os +import unittest + +from requests import HTTPError + +from modelscope import MsDataset, snapshot_download +from modelscope.hub.constants import MODELSCOPE_PREFER_AI_SITE + + +class HubAiSiteTest(unittest.TestCase): + + def setUp(self): + ... + + # test download from an ai-site only model, it should + # work as expected since we shall fall back to ai-site + # when the model is not found on cn-site. + def test_default_download_from_ai_site(self): + model_id = 'ModelScope_Developer/ai_only' + model_dir = snapshot_download(model_id) + contents = os.listdir(model_dir) + assert len(contents) > 0 + + # test download from a cn-site only model, it should + # work as expected as it is found directly on cn-site. + def test_default_download_from_cn_site(self): + model_id = 'ModelScope_Developer/cn_only' + model_dir = snapshot_download(model_id) + contents = os.listdir(model_dir) + assert len(contents) > 0 + + # test download a model that exists on both cn and ai site + # when prefer-ai-site is set, we should found the version from + # on ai-site, not cn-site + def test_prefer_ai_site_and_download_from_ai_site(self): + os.environ[MODELSCOPE_PREFER_AI_SITE] = 'True' + model_id = 'ModelScope_Developer/same_name' + model_dir = snapshot_download(model_id) + cn_site_only_file = os.path.join(model_dir, 'on_ai_site') + assert os.path.exists(cn_site_only_file) + + # test download a model that exists on both cn and ai site + # when prefer-ai-site is NOT set, we should found the version from + # on cn-site, not ai-site + def test_prefer_cn_site_and_download_from_cn_site(self): + os.environ[MODELSCOPE_PREFER_AI_SITE] = 'False' + model_id = 'ModelScope_Developer/same_name' + model_dir = snapshot_download(model_id) + cn_site_only_file = os.path.join(model_dir, 'on_cn_site') + assert os.path.exists(cn_site_only_file) + + def test_download_non_exist_model(self): + with self.assertRaises(HTTPError): + model_id = 'ModelScope_Developer/not_exist_model' + snapshot_download(model_id) + + # test download dataset from ai site + def test_download_dataset_from_ai_site(self): + os.environ[MODELSCOPE_PREFER_AI_SITE] = 'True' + dataset_id = 'ModelScope_Developer/ai_only_dataset' + dataset = MsDataset.load(dataset_id) + assert dataset + + +if __name__ == '__main__': + unittest.main() diff --git a/tests/hub/test_file_exists.py b/tests/hub/test_file_exists.py index 453541ce..e5ffe68c 100644 --- a/tests/hub/test_file_exists.py +++ b/tests/hub/test_file_exists.py @@ -5,7 +5,6 @@ from modelscope.hub.api import HubApi from modelscope.utils.logger import get_logger logger = get_logger() -logger.setLevel('DEBUG') DEFAULT_GIT_PATH = 'git' download_model_file_name = 'test.bin' diff --git a/tests/hub/test_hub_repository.py b/tests/hub/test_hub_repository.py index 7631f5db..92d89e74 100644 --- a/tests/hub/test_hub_repository.py +++ b/tests/hub/test_hub_repository.py @@ -21,7 +21,6 @@ from modelscope.utils.test_utils import (TEST_ACCESS_TOKEN1, TEST_MODEL_ORG, delete_credential) logger = get_logger() -logger.setLevel('DEBUG') DEFAULT_GIT_PATH = 'git' download_model_file_name = 'test.bin' diff --git a/tests/hub/test_hub_revision.py b/tests/hub/test_hub_revision.py index 642742bc..9a1e9f8a 100644 --- a/tests/hub/test_hub_revision.py +++ b/tests/hub/test_hub_revision.py @@ -18,7 +18,6 @@ from modelscope.utils.test_utils import (TEST_ACCESS_TOKEN1, TEST_MODEL_ORG) logger = get_logger() -logger.setLevel('DEBUG') download_model_file_name = 'test.bin' download_model_file_name2 = 'test2.bin' diff --git a/tests/hub/test_hub_revision_release_mode.py b/tests/hub/test_hub_revision_release_mode.py index 823e1d5d..74a48527 100644 --- a/tests/hub/test_hub_revision_release_mode.py +++ b/tests/hub/test_hub_revision_release_mode.py @@ -21,7 +21,6 @@ from modelscope.utils.test_utils import (TEST_ACCESS_TOKEN1, TEST_MODEL_ORG) logger = get_logger() -logger.setLevel('DEBUG') download_model_file_name = 'test.bin' download_model_file_name2 = 'test2.bin' diff --git a/tests/hub/test_hub_upload.py b/tests/hub/test_hub_upload.py index 8a67a9de..640ddf69 100644 --- a/tests/hub/test_hub_upload.py +++ b/tests/hub/test_hub_upload.py @@ -11,7 +11,7 @@ from modelscope.hub.constants import Licenses, ModelVisibility from modelscope.hub.errors import GitError, HTTPError, NotLoginException from modelscope.hub.push_to_hub import push_to_hub, push_to_hub_async from modelscope.hub.repository import Repository -from modelscope.utils.constant import ModelFile +from modelscope.utils.constant import REPO_TYPE_DATASET, ModelFile from modelscope.utils.logger import get_logger from modelscope.utils.test_utils import (TEST_ACCESS_TOKEN1, TEST_MODEL_ORG, delete_credential, test_level) @@ -52,6 +52,12 @@ class HubUploadTest(unittest.TestCase): self.assertTrue(res) res = self.api.repo_exists('Qwen/not-a-repo') self.assertFalse(res) + res = self.api.repo_exists( + 'Qwen/ProcessBench', repo_type=REPO_TYPE_DATASET) + self.assertTrue(res) + res = self.api.repo_exists( + 'Qwen/not-a-repo', repo_type=REPO_TYPE_DATASET) + self.assertFalse(res) def test_upload_exits_repo_master(self): logger.info('basic test for upload!')