v1.0.0 update

This commit is contained in:
hanzhn
2024-05-31 14:05:54 +08:00
parent edff6352d5
commit aa7e959330
3 changed files with 69 additions and 7 deletions
+13 -6
View File
@@ -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),