diff --git a/scepter/modules/data/dataset/ms_dataset.py b/scepter/modules/data/dataset/ms_dataset.py index ab656fb..670ba7d 100644 --- a/scepter/modules/data/dataset/ms_dataset.py +++ b/scepter/modules/data/dataset/ms_dataset.py @@ -101,14 +101,13 @@ class ImageTextPairMSDataset(BaseDataset): if isinstance(self.output_size, numbers.Number): self.output_size = [self.output_size, self.output_size] # Use modelscope dataset - if not ms_dataset_name: raise ( 'Your must set MS_DATASET_NAME as modelscope dataset or your local dataset orignized ' 'as modelscope dataset.') if FS.exists(ms_dataset_name): ms_dataset_name = FS.get_dir_to_local_dir(ms_dataset_name) - ms_remap_path = ms_dataset_name + # ms_remap_path = ms_dataset_name try: self.data = MsDataset.load(str(ms_dataset_name), namespace=ms_dataset_namespace, @@ -151,8 +150,9 @@ class ImageTextPairMSDataset(BaseDataset): def _get(self, index: int): current_data = self.data[index % len(self.data)] # print(current_data.keys()) - image_path = current_data['Target:FILE'] - prompt = current_data['Prompt'] + image_path = current_data[ + 'Target:FILE'] if 'Target:FILE' in current_data else '' + prompt = current_data.get('Prompt', current_data.get('prompt', '')) style = current_data['Style'] if 'Style' in current_data else '' src_image_path = current_data[ 'Source:FILE'] if 'Source:FILE' in current_data else '' @@ -178,6 +178,9 @@ class ImageTextPairMSDataset(BaseDataset): } if self.output_size is not None: ret_item['meta']['image_size'] = self.output_size + for key in current_data: + if key not in ret_item['meta']: + ret_item['meta'][key] = current_data[key] return ret_item @staticmethod @@ -269,8 +272,9 @@ class ImageTextPairFolderDataset(BaseDataset): def _get(self, index: int): current_data = self.data[index % len(self.data)] # print(current_data.keys()) - image_path = current_data['Target:FILE'] - prompt = current_data['Prompt'] + image_path = current_data[ + 'Target:FILE'] if 'Target:FILE' in current_data else '' + prompt = current_data.get('Prompt', current_data.get('prompt', '')) style = current_data['Style'] if 'Style' in current_data else '' src_image_path = current_data[ 'Source:FILE'] if 'Source:FILE' in current_data else '' @@ -296,6 +300,9 @@ class ImageTextPairFolderDataset(BaseDataset): } if self.output_size is not None: ret_item['meta']['image_size'] = self.output_size + for key in current_data: + if key not in ret_item['meta']: + ret_item['meta'][key] = current_data[key] return ret_item @staticmethod diff --git a/scepter/modules/utils/file_clients/modelscope_fs.py b/scepter/modules/utils/file_clients/modelscope_fs.py index eb152e8..bf4fb26 100644 --- a/scepter/modules/utils/file_clients/modelscope_fs.py +++ b/scepter/modules/utils/file_clients/modelscope_fs.py @@ -48,6 +48,12 @@ class ModelscopeFs(BaseFs): from modelscope.hub.file_download import model_file_download key = osp.relpath(target_path, self.get_prefix()) + + if '#' in key: + key, token = key.split('#', 1) + else: + token = None + key, file_path = key.split('@', 1) if ':' in key: @@ -72,9 +78,14 @@ class ModelscopeFs(BaseFs): if not osp.exists(local_path): self._model_file_loaded.remove(key) else: + if token is not None: + cookies = self.get_modelscope_cookie(token) + else: + cookies = None local_path = model_file_download(model_id=key, revision=revision, file_path=file_path, + cookies=cookies, cache_dir=local_path) if osp.exists(local_path): break @@ -100,6 +111,12 @@ class ModelscopeFs(BaseFs): assert target_path.startswith(self.get_prefix()) key = osp.relpath(target_path, self.get_prefix()) + + if '#' in key: + key, token = key.split('#', 1) + else: + token = None + if '@' not in key: key, ret_folder = key.split('@', 1)[0], '' else: @@ -129,8 +146,13 @@ class ModelscopeFs(BaseFs): if not osp.exists(local_path): self._model_id_loaded.remove(key) else: + if token is not None: + cookies = self.get_modelscope_cookie(token) + else: + cookies = None local_path = snapshot_download(key, revision=revision, + cookies=cookies, cache_dir=local_path) if osp.exists(local_path): break @@ -147,6 +169,21 @@ class ModelscopeFs(BaseFs): local_path = os.path.join(local_path, ret_folder) return local_path + def get_modelscope_cookie(self, m_session_id: str): + from modelscope.hub.utils.utils import get_endpoint + from modelscope.hub.api import ModelScopeConfig + from modelscope.hub.errors import raise_on_error + import requests + path = f'{get_endpoint()}/api/v1/login' + r = requests.post( + path, + json={'AccessToken': m_session_id}, + headers={'user-agent': ModelScopeConfig.get_user_agent()}) + r.raise_for_status() + d = r.json() + raise_on_error(d) + return r.cookies + def get_object(self, target_path): try: local_data = open(self.get_object_to_local_file(target_path), diff --git a/tests/utils/test_fs.py b/tests/utils/test_fs.py index 06b7a5e..8da327e 100644 --- a/tests/utils/test_fs.py +++ b/tests/utils/test_fs.py @@ -52,6 +52,24 @@ class FSTest(unittest.TestCase): print(f'Download from {path} to {local_path}') self.assertTrue(os.path.exists(local_path)) + @unittest.skip('') + def test_modelscope_token(self): + fs_info = {'NAME': 'ModelscopeFs', 'TEMP_DIR': 'cache/data'} + config = Config(load=False, cfg_dict=fs_info) + FS.init_fs_client(config) + + # path = 'ms://group_name/model_id:revision@file_path' + path = 'ms://group_name/model_id:revision@file#token' + with FS.get_from(path, wait_finish=True) as local_path: + print(f'Download from {path} to {local_path}') + self.assertTrue(os.path.exists(local_path)) + + path = 'ms://group_name/model_id@file_dir#token' + # path = 'ms://group_name/model_id#token' + with FS.get_dir_to_local_dir(path, wait_finish=True) as local_path: + print(f'Download from {path} to {local_path}') + self.assertTrue(os.path.exists(local_path)) + @unittest.skip('') def test_huggingface(self): fs_info = {'NAME': 'HuggingfaceFs', 'TEMP_DIR': 'cache/data'} @@ -73,7 +91,7 @@ class FSTest(unittest.TestCase): print(f'Download from {path} to {local_path}') self.assertTrue(os.path.exists(local_path)) - # @unittest.skip('') + @unittest.skip('') def test_scedit(self): fs_info = {'NAME': 'ModelscopeFs', 'TEMP_DIR': 'cache/cache_data'} config = Config(load=False, cfg_dict=fs_info)