update 1.4.0
This commit is contained in:
@@ -1,4 +1,21 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from typing import TYPE_CHECKING
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
|
||||
from scepter.modules.data import dataset, sampler
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from scepter.modules.data import dataset, sampler
|
||||
else:
|
||||
_import_structure = {
|
||||
'data': ['dataset', 'sampler']
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
@@ -1,12 +1,35 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from typing import TYPE_CHECKING
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
|
||||
from scepter.modules.data.dataset.base_dataset import BaseDataset
|
||||
from scepter.modules.data.dataset.dataset import (Image2ImageDataset,
|
||||
ImageClassifyPublicDataset,
|
||||
ImageTextPairDataset,
|
||||
Text2ImageDataset)
|
||||
from scepter.modules.data.dataset.ms_dataset import (
|
||||
ImageTextPairFolderDataset, ImageTextPairMSDataset)
|
||||
from scepter.modules.data.dataset.registry import DATASETS
|
||||
from scepter.modules.data.dataset.video_gen_dataset import VideoGenDataset
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from scepter.modules.data.dataset.base_dataset import BaseDataset
|
||||
from scepter.modules.data.dataset.dataset import (Image2ImageDataset,
|
||||
ImageClassifyPublicDataset,
|
||||
ImageTextPairDataset,
|
||||
Text2ImageDataset)
|
||||
from scepter.modules.data.dataset.ms_dataset import (
|
||||
ImageTextPairFolderDataset, ImageTextPairMSDataset)
|
||||
from scepter.modules.data.dataset.registry import DATASETS
|
||||
from scepter.modules.data.dataset.video_gen_dataset import VideoGenDataset
|
||||
else:
|
||||
_import_structure = {
|
||||
'base_dataset': ['BaseDataset'],
|
||||
'dataset': ['Image2ImageDataset', 'ImageClassifyPublicDataset',
|
||||
'ImageTextPairDataset', 'Text2ImageDataset'],
|
||||
'ms_dataset': ['ImageTextPairFolderDataset',
|
||||
'ImageTextPairMSDataset'],
|
||||
'registry': ['DATASETS'],
|
||||
'video_gen_dataset': ['VideoGenDataset']
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
@@ -82,7 +82,7 @@ class BaseDataset(Dataset, metaclass=ABCMeta):
|
||||
overwrite=False)
|
||||
self.worker_id = worker_id
|
||||
self.logger = self.worker_logger
|
||||
self.local_we["seed"] += (worker_id + we.rank)
|
||||
self.local_we["seed"] += (worker_id + self.local_we['rank'] * 1234)
|
||||
self.seed = self.local_we["seed"]
|
||||
we.set_env(self.local_we)
|
||||
|
||||
|
||||
@@ -4,6 +4,7 @@
|
||||
import numbers
|
||||
import os
|
||||
import sys
|
||||
import copy
|
||||
from collections.abc import Iterable
|
||||
|
||||
import numpy as np
|
||||
|
||||
@@ -386,6 +386,11 @@ class ImageTextPairMSDatasetForACE(BaseDataset):
|
||||
'description':
|
||||
'The keywords sign you want to add, which is like <{HIGHLIGHT_KEYWORDS}{KEYWORDS_SIGN}>'
|
||||
},
|
||||
'ALIGN_SIZE': {
|
||||
'value': False,
|
||||
'description':
|
||||
'Whether ensure the size align between the source image and target image.'
|
||||
},
|
||||
'OUTPUT_SIZE': {
|
||||
'value':
|
||||
None,
|
||||
@@ -414,6 +419,8 @@ class ImageTextPairMSDatasetForACE(BaseDataset):
|
||||
self.replace_keywords = cfg.get('HIGHLIGHT_KEYWORDS', '')
|
||||
self.keywords_sign = cfg.get('KEYWORDS_SIGN', '')
|
||||
self.add_indicator = cfg.get('ADD_INDICATOR', False)
|
||||
|
||||
self.align_size = cfg.get('ALIGN_SIZE', False)
|
||||
# Use modelscope dataset
|
||||
if not ms_dataset_name:
|
||||
raise ValueError(
|
||||
@@ -492,7 +499,7 @@ class ImageTextPairMSDatasetForACE(BaseDataset):
|
||||
tar_image_path,
|
||||
cvt_type='RGB')
|
||||
src_image = self.image_preprocess(src_image)
|
||||
tar_image = self.image_preprocess(tar_image)
|
||||
tar_image = self.image_preprocess(tar_image, size = src_image.shape[:2] if self.align_size else None)
|
||||
|
||||
tar_image = self.transforms(tar_image)
|
||||
src_image = self.transforms(src_image)
|
||||
@@ -501,13 +508,13 @@ class ImageTextPairMSDatasetForACE(BaseDataset):
|
||||
if self.add_indicator:
|
||||
if '{image}' not in prompt:
|
||||
prompt = '{image}, ' + prompt
|
||||
|
||||
return {
|
||||
'edit_image': [src_image],
|
||||
'edit_image_mask': [src_mask],
|
||||
'src_image_list': [src_image],
|
||||
'src_mask_list': [src_mask],
|
||||
'image': tar_image,
|
||||
'image_mask': tar_mask,
|
||||
'prompt': [prompt],
|
||||
'edit_id': [0]
|
||||
}
|
||||
|
||||
def load_image(self, prefix, img_path, cvt_type=None):
|
||||
|
||||
@@ -337,8 +337,12 @@ def build_dataset_config(cfg, registry, logger=None, *args, **kwargs):
|
||||
f'registry must be type Registry, got {type(registry)}')
|
||||
|
||||
cfg = deep_copy(cfg)
|
||||
|
||||
req_type = cfg.get('NAME')
|
||||
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
sig = (registry.name.upper(), req_type)
|
||||
LazyImportModule.import_module(sig)
|
||||
|
||||
if isinstance(req_type, str):
|
||||
req_type_entry = registry.get(req_type)
|
||||
if req_type_entry is None:
|
||||
|
||||
@@ -1,44 +1,46 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import io
|
||||
import os
|
||||
import random
|
||||
import sys
|
||||
import os
|
||||
import warnings
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.data.dataset import DATASETS, BaseDataset
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
try:
|
||||
import decord
|
||||
decord.bridge.set_bridge("torch")
|
||||
decord.bridge.set_bridge('torch')
|
||||
except ImportError:
|
||||
warnings.warn(
|
||||
"The `decord` package is required for loading the video dataset. Install with `pip install decord`"
|
||||
'The `decord` package is required for loading the video dataset. Install with `pip install decord`'
|
||||
)
|
||||
|
||||
|
||||
@DATASETS.register_class()
|
||||
class VideoGenDataset(BaseDataset):
|
||||
def __init__(self, cfg, logger = None):
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.prompt_prefix = cfg.get('PROMPT_PREFIX', '')
|
||||
self.path_prefix = cfg.get('PATH_PREFIX', '')
|
||||
self.p_zero = cfg.get('P_ZERO', 0.0)
|
||||
self.max_num_frames = cfg.get("NUM_FRAMES", 49)
|
||||
self.fps = cfg.get("FPS", 8)
|
||||
self.height = cfg.get("HEIGHT", 480)
|
||||
self.width = cfg.get("WIDTH", 720)
|
||||
self.skip_frames_start = cfg.get("SKIP_FRAMES_START", 0)
|
||||
self.skip_frames_end = cfg.get("SKIP_FRAMES_END", 0)
|
||||
self.max_num_frames = cfg.get('NUM_FRAMES', 49)
|
||||
self.fps = cfg.get('FPS', 8)
|
||||
self.height = cfg.get('HEIGHT', 480)
|
||||
self.width = cfg.get('WIDTH', 720)
|
||||
self.skip_frames_start = cfg.get('SKIP_FRAMES_START', 0)
|
||||
self.skip_frames_end = cfg.get('SKIP_FRAMES_END', 0)
|
||||
self.data_type = cfg.get('DATA_TYPE', 't2v')
|
||||
|
||||
def worker_init_fn(self, worker_id, num_workers=1):
|
||||
super().worker_init_fn(worker_id, num_workers=num_workers)
|
||||
randseed = np.random.randint(0, 2 ** 32 - num_workers - 1)
|
||||
randseed = np.random.randint(0, 2**32 - num_workers - 1)
|
||||
workerseed = randseed + worker_id
|
||||
random.seed(workerseed)
|
||||
np.random.seed(workerseed)
|
||||
@@ -46,7 +48,9 @@ class VideoGenDataset(BaseDataset):
|
||||
def _preprocess_video_data(self, video_path):
|
||||
|
||||
with FS.get_object(video_path) as video_data:
|
||||
video_reader = decord.VideoReader(io.BytesIO(video_data), width=self.width, height=self.height)
|
||||
video_reader = decord.VideoReader(io.BytesIO(video_data),
|
||||
width=self.width,
|
||||
height=self.height)
|
||||
video_num_frames = len(video_reader)
|
||||
|
||||
start_frame = min(self.skip_frames_start, video_num_frames)
|
||||
@@ -54,13 +58,16 @@ class VideoGenDataset(BaseDataset):
|
||||
if end_frame <= start_frame:
|
||||
frames = video_reader.get_batch([start_frame])
|
||||
elif end_frame - start_frame <= self.max_num_frames:
|
||||
frames = video_reader.get_batch(list(range(start_frame, end_frame)))
|
||||
frames = video_reader.get_batch(list(range(start_frame,
|
||||
end_frame)))
|
||||
else:
|
||||
indices = list(range(start_frame, end_frame, (end_frame - start_frame) // self.max_num_frames))
|
||||
indices = list(
|
||||
range(start_frame, end_frame,
|
||||
(end_frame - start_frame) // self.max_num_frames))
|
||||
frames = video_reader.get_batch(indices)
|
||||
|
||||
# Ensure that we don't go over the limit
|
||||
frames = frames[: self.max_num_frames]
|
||||
frames = frames[:self.max_num_frames]
|
||||
selected_num_frames = frames.shape[0]
|
||||
|
||||
# Choose first (4k + 1) frames as this is how many is required by the VAE
|
||||
@@ -73,14 +80,16 @@ class VideoGenDataset(BaseDataset):
|
||||
|
||||
# Training transforms
|
||||
frames = frames.float().div_(127.5).sub_(1.)
|
||||
frames = frames.permute(3, 0, 1, 2).contiguous() # [C, F, H, W]
|
||||
frames = frames.permute(3, 0, 1, 2).contiguous() # [C, F, H, W]
|
||||
return frames
|
||||
|
||||
def _parse_index(self, index):
|
||||
meta = dict()
|
||||
for key, value in zip(index[-1], index[:-1]):
|
||||
if key in ['oss_key', 'path', 'video_path']:
|
||||
if key in ['oss_key', 'path', 'video_path', 'target_video_path']:
|
||||
meta['video_path'] = value
|
||||
elif key in ['source_video_path', 'src_video_path']:
|
||||
meta['src_video_path'] = value
|
||||
elif key in ['prompt', 'caption', 'text']:
|
||||
meta['prompt'] = value
|
||||
elif key in ['width', 'height']:
|
||||
@@ -104,8 +113,13 @@ class VideoGenDataset(BaseDataset):
|
||||
'prompt': prompt,
|
||||
'meta': meta,
|
||||
}
|
||||
if self.data_type == 'i2v':
|
||||
if 'i2v' in self.data_type:
|
||||
item['image'] = item['video'][:, :1, :, :]
|
||||
if 'v2v' in self.data_type:
|
||||
src_video_path = os.path.join(self.path_prefix,
|
||||
meta.get('src_video_path', ''))
|
||||
src_video = self._preprocess_video_data(src_video_path)
|
||||
item['src_video'] = src_video
|
||||
return item
|
||||
|
||||
def __len__(self):
|
||||
@@ -122,7 +136,6 @@ class VideoGenDataset(BaseDataset):
|
||||
return collect
|
||||
|
||||
|
||||
|
||||
@DATASETS.register_class()
|
||||
class VideoGenDatasetOTF(VideoGenDataset):
|
||||
def __init__(self, cfg, logger=None):
|
||||
@@ -135,8 +148,11 @@ class VideoGenDatasetOTF(VideoGenDataset):
|
||||
from scepter.modules.model.registry import MODELS
|
||||
model_cfg = cfg.get('MODEL', None)
|
||||
if model_cfg is not None:
|
||||
self.model = MODELS.build(cfg.MODEL, logger=logger).eval().requires_grad_(False).to(we.device_id)
|
||||
self.items = self.parse_data(self.data_file, self.delimiter, self.fields)
|
||||
self.model = MODELS.build(
|
||||
cfg.MODEL,
|
||||
logger=logger).eval().requires_grad_(False).to(we.device_id)
|
||||
self.items = self.parse_data(self.data_file, self.delimiter,
|
||||
self.fields)
|
||||
if self.use_num and self.use_num > 0:
|
||||
self.items = self.items[:self.use_num]
|
||||
self.data = self.encode(self.items)
|
||||
@@ -169,11 +185,13 @@ class VideoGenDatasetOTF(VideoGenDataset):
|
||||
return items
|
||||
|
||||
def encode(self, items):
|
||||
self.logger.info("Start to encode video data [{}]!".format(len(items)))
|
||||
self.logger.info('Start to encode video data [{}]!'.format(len(items)))
|
||||
for item in tqdm(items):
|
||||
video_path = os.path.join(self.path_prefix, item.get('video_path', ''))
|
||||
video_path = os.path.join(self.path_prefix,
|
||||
item.get('video_path', ''))
|
||||
video = self._preprocess_video_data(video_path)
|
||||
latent = self.model.encode_first_stage(video.unsqueeze(0).to(we.device_id)).squeeze(0)
|
||||
latent = self.model.encode_first_stage(
|
||||
video.unsqueeze(0).to(we.device_id)).squeeze(0)
|
||||
item['video_latent'] = latent.detach().cpu()
|
||||
item['video'] = video
|
||||
if self.data_type == 'i2v':
|
||||
@@ -181,4 +199,4 @@ class VideoGenDatasetOTF(VideoGenDataset):
|
||||
return items
|
||||
|
||||
def _get(self, index):
|
||||
return self.data[index % self.real_number]
|
||||
return self.data[index % self.real_number]
|
||||
|
||||
@@ -1,9 +1,31 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from typing import TYPE_CHECKING
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
|
||||
from scepter.modules.data.sampler.base_sampler import BaseSampler
|
||||
from scepter.modules.data.sampler.registry import SAMPLERS
|
||||
from scepter.modules.data.sampler.sampler import (
|
||||
EvalDistributedSampler, LoopSampler, MixtureOfSamplers,
|
||||
MultiFoldDistributedSampler, MultiLevelBatchSampler,
|
||||
MultiLevelBatchSamplerMultiSource, ResolutionBatchSampler)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from scepter.modules.data.sampler.base_sampler import BaseSampler
|
||||
from scepter.modules.data.sampler.registry import SAMPLERS
|
||||
from scepter.modules.data.sampler.sampler import (
|
||||
EvalDistributedSampler, LoopSampler, MixtureOfSamplers,
|
||||
MultiFoldDistributedSampler, MultiLevelBatchSampler,
|
||||
MultiLevelBatchSamplerMultiSource, ResolutionBatchSampler)
|
||||
else:
|
||||
_import_structure = {
|
||||
'base_sampler': ['BaseSampler'],
|
||||
'registry': ['SAMPLERS'],
|
||||
'sampler': ['EvalDistributedSampler', 'LoopSampler',
|
||||
'MixtureOfSamplers', 'MultiFoldDistributedSampler',
|
||||
'MultiLevelBatchSampler', 'MultiLevelBatchSamplerMultiSource',
|
||||
'ResolutionBatchSampler']
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
@@ -35,8 +35,12 @@ def build_sampler_config(cfg, registry, logger=None, **kwargs):
|
||||
f'registry must be type Registry, got {type(registry)}')
|
||||
|
||||
cfg = deep_copy(cfg)
|
||||
|
||||
req_type = cfg.get('NAME')
|
||||
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
sig = (registry.name.upper(), req_type)
|
||||
LazyImportModule.import_module(sig)
|
||||
|
||||
if isinstance(req_type, str):
|
||||
req_type_entry = registry.get(req_type)
|
||||
if req_type_entry is None:
|
||||
|
||||
@@ -1,4 +1,21 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from typing import TYPE_CHECKING
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
|
||||
from scepter.modules.data.utils.data_bucket import BucketManager
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from scepter.modules.data.utils.data_bucket import BucketManager
|
||||
else:
|
||||
_import_structure = {
|
||||
'data_bucket': ['BucketManager']
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user