update 1.4.0

This commit is contained in:
jiangzeyinzi
2025-02-03 13:36:44 +08:00
parent d7dbdc5292
commit 043222de49
130 changed files with 5065 additions and 704 deletions
+18 -1
View File
@@ -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={},
)
+32 -9
View File
@@ -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={},
)
+1 -1
View File
@@ -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)
+1
View File
@@ -4,6 +4,7 @@
import numbers
import os
import sys
import copy
from collections.abc import Iterable
import numpy as np
+11 -4
View File
@@ -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):
+5 -1
View File
@@ -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]
+28 -6
View File
@@ -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={},
)
+5 -1
View File
@@ -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:
+18 -1
View File
@@ -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={},
)