v1.0.0 update
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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),
|
||||
|
||||
+19
-1
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user