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