v1.0.0 update

This commit is contained in:
hanzhn
2024-05-27 13:15:48 +08:00
parent 8076aae7da
commit c70ef0fc47
186 changed files with 7505 additions and 3117 deletions
+1 -2
View File
@@ -3,14 +3,13 @@
from abc import ABCMeta, abstractmethod
from torch.utils.data import Dataset
from scepter.modules.transform.registry import TRANSFORMS, build_pipeline
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.logger import get_logger
from scepter.modules.utils.registry import old_python_version
from torch.utils.data import Dataset
class BaseDataset(Dataset, metaclass=ABCMeta):
+5 -1
View File
@@ -8,7 +8,6 @@ from collections.abc import Iterable
import numpy as np
import torchvision
from scepter.modules.data.dataset.base_dataset import BaseDataset
from scepter.modules.data.dataset.registry import DATASETS
from scepter.modules.utils.config import dict_to_yaml
@@ -272,6 +271,11 @@ class Text2ImageDataset(BaseDataset):
item['prompt'] = prompt_prefix + value
elif key in ['oss_key', 'path', 'img_path', 'target_img_path']:
item['meta']['img_path'] = os.path.join(path_prefix, value)
elif key in [
'src_oss_key', 'src_path', 'src_img_path',
'src_target_img_path'
]:
item['meta']['src_path'] = os.path.join(path_prefix, value)
elif key in ['width', 'height']:
item['meta'][key] = int(value)
else:
+18 -11
View File
@@ -133,9 +133,11 @@ class ImageTextPairMSDataset(BaseDataset):
if ms_remap_path:
def map_func(example):
example['Target:FILE'] = os.path.join(ms_remap_path,
example['Target:FILE'])
return example
return {
k: os.path.join(ms_remap_path, v)
if k.endswith(':FILE') else v
for k, v in example.items()
}
self.data = self.data.ds_instance.map(map_func)
self.real_number = len(self.data)
@@ -152,6 +154,8 @@ class ImageTextPairMSDataset(BaseDataset):
image_path = current_data['Target:FILE']
prompt = current_data['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 ''
# print(prompt, style)
if self.replace_style and not style == '':
prompt = prompt.replace(style, f'<{self.keywords_sign}>')
@@ -166,6 +170,7 @@ class ImageTextPairMSDataset(BaseDataset):
ret_item = {
'meta': {
'img_path': image_path,
'src_path': src_image_path,
'data_key': style,
'data_num': self.real_number
},
@@ -241,19 +246,18 @@ class ImageTextPairFolderDataset(BaseDataset):
data_folder = FS.get_dir_to_local_dir(data_folder)
all_lines = open(os.path.join(data_folder, 'train.csv'),
'r').read().split('\n')
assert all_lines[0] == 'Target:FILE,Prompt'
header = all_lines[0].split(',')
self.data = []
for line in all_lines[1:]:
line = line.strip()
if line == '':
continue
self.data.append({
'Target:FILE':
os.path.join(data_folder,
line.split(',', 1)[0]),
'Prompt':
line.split(',', 1)[1]
})
record = dict(zip(header, line.split(',', len(header) - 1)))
record = {
k: os.path.join(data_folder, v) if k.endswith(':FILE') else v
for k, v in record.items()
}
self.data.append(record)
self.real_number = len(self.data)
def __len__(self):
@@ -268,6 +272,8 @@ class ImageTextPairFolderDataset(BaseDataset):
image_path = current_data['Target:FILE']
prompt = current_data['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 ''
# print(prompt, style)
if self.replace_style and not style == '':
prompt = prompt.replace(style, f'<{self.keywords_sign}>')
@@ -282,6 +288,7 @@ class ImageTextPairFolderDataset(BaseDataset):
ret_item = {
'meta': {
'img_path': image_path,
'src_path': src_image_path,
'data_key': style,
'data_num': self.real_number
},
+1 -2
View File
@@ -7,13 +7,12 @@ import re
from functools import partial
import torch
from torch.utils.data import DataLoader, DistributedSampler
from scepter.modules.data.sampler import (SAMPLERS, MixtureOfSamplers,
MultiFoldDistributedSampler,
MultiLevelBatchSampler)
from scepter.modules.utils.registry import (Registry, deep_copy,
old_python_version)
from torch.utils.data import DataLoader, DistributedSampler
string_classes = (str, bytes)
int_classes = (int, )
+1 -2
View File
@@ -1,10 +1,9 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from torch.utils.data.sampler import Sampler
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from torch.utils.data.sampler import Sampler
class BaseSampler(Sampler):
-1
View File
@@ -12,7 +12,6 @@ from typing import List, Optional
import numpy as np
import torch
import torch.distributed as dist
from scepter.modules.data.sampler.base_sampler import BaseSampler
from scepter.modules.data.sampler.registry import SAMPLERS
from scepter.modules.data.utils.data_bucket import (BucketBatchIndex,