upgrade from 1.2.0 to 1.3.0

This commit is contained in:
maochaojie
2024-11-19 19:20:02 +08:00
parent 4ec0492897
commit a683061c6f
87 changed files with 10400 additions and 1117 deletions
+2 -2
View File
@@ -7,6 +7,6 @@ from scepter.modules.data.dataset.dataset import (Image2ImageDataset,
ImageTextPairDataset,
Text2ImageDataset)
from scepter.modules.data.dataset.ms_dataset import (
ImageTextPairFolderDataset, ImageTextPairMSDataset,
ImageTextPairMSDatasetForACE)
ImageTextPairFolderDataset, ImageTextPairMSDataset)
from scepter.modules.data.dataset.registry import DATASETS
from scepter.modules.data.dataset.video_gen_dataset import VideoGenDataset
+8 -1
View File
@@ -242,6 +242,8 @@ class Text2ImageDataset(BaseDataset):
prompt_prefix = cfg.get('PROMPT_PREFIX', '')
path_prefix = cfg.get('PATH_PREFIX', '')
use_num = cfg.get('USE_NUM', -1)
meta_cfg = cfg.get('META_CFG', None)
meta_cfg = meta_cfg.get_lowercase_dict() if meta_cfg is not None else None
image_size = cfg.get('IMAGE_SIZE', 1024)
if isinstance(image_size, numbers.Number):
@@ -264,7 +266,12 @@ class Text2ImageDataset(BaseDataset):
self.items = list()
for i, row in enumerate(rows):
item = {'index': i, 'meta': {'image_size': image_size}}
if meta_cfg is not None:
meta_cfg_copy = copy.deepcopy(meta_cfg)
meta_cfg_copy['image_size'] = image_size
item = {'index': i, 'meta': meta_cfg_copy}
else:
item = {'index': i, 'meta': {'image_size': image_size}}
for key, value in zip(fields, row):
if key in ['prompt', 'caption', 'text']:
item['ori_prompt'] = value
@@ -0,0 +1,184 @@
import io
import random
import sys
import os
import warnings
import torch
import numpy as np
from tqdm import tqdm
from scepter.modules.utils.distribute import we
from scepter.modules.data.dataset import DATASETS, BaseDataset
from scepter.modules.utils.file_system import FS
try:
import decord
decord.bridge.set_bridge("torch")
except ImportError:
warnings.warn(
"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):
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.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)
workerseed = randseed + worker_id
random.seed(workerseed)
np.random.seed(workerseed)
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_num_frames = len(video_reader)
start_frame = min(self.skip_frames_start, video_num_frames)
end_frame = max(0, video_num_frames - self.skip_frames_end)
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)))
else:
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]
selected_num_frames = frames.shape[0]
# Choose first (4k + 1) frames as this is how many is required by the VAE
remainder = (3 + (selected_num_frames % 4)) % 4
if remainder != 0:
frames = frames[:-remainder]
selected_num_frames = frames.shape[0]
assert (selected_num_frames - 1) % 4 == 0
# Training transforms
frames = frames.float().div_(127.5).sub_(1.)
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']:
meta['video_path'] = value
elif key in ['prompt', 'caption', 'text']:
meta['prompt'] = value
elif key in ['width', 'height']:
meta[key] = int(value)
else:
meta[key] = value
return meta
def _get(self, index):
meta = self._parse_index(index)
video_path = os.path.join(self.path_prefix, meta.get('video_path', ''))
video = self._preprocess_video_data(video_path)
prompt = self.prompt_prefix + meta.get('prompt', '')
if self.mode == 'train' and np.random.uniform() < self.p_zero:
prompt = ''
item = {
'video': video,
'prompt': prompt,
'meta': meta,
}
if self.data_type == 'i2v':
item['image'] = item['video'][:, :1, :, :]
return item
def __len__(self):
return sys.maxsize
@staticmethod
def collate_fn(batch):
collect = {}
for sample in batch:
for k, v in sample.items():
if k not in collect:
collect[k] = []
collect[k].append(v)
return collect
@DATASETS.register_class()
class VideoGenDatasetOTF(VideoGenDataset):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger)
self.data_file = cfg.DATA_FILE
self.delimiter = cfg.get('DELIMITER', '#;#')
self.fields = cfg.get('FIELDS', ['video_path', 'prompt'])
self.use_num = cfg.get('USE_NUM', -1)
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)
if self.use_num and self.use_num > 0:
self.items = self.items[:self.use_num]
self.data = self.encode(self.items)
self.real_number = len(self.data)
if model_cfg is not None:
self.model.to('cpu')
del self.model
torch.cuda.empty_cache()
def parse_data(self, data_file, delimiter, fields):
items = list()
with FS.get_object(data_file) as local_data:
rows = [
i.split(delimiter,
len(fields) - 1)
for i in local_data.decode('utf-8').strip().split('\n')
]
for i, row in enumerate(rows):
item = {}
for key, value in zip(self.fields, row):
if key in ['oss_key', 'path', 'video_path']:
item['video_path'] = value
elif key in ['prompt', 'caption', 'text']:
item['prompt'] = value
elif key in ['width', 'height']:
item[key] = int(value)
else:
item[key] = value
items.append(item)
return items
def encode(self, 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 = self._preprocess_video_data(video_path)
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':
item['image'] = item['video'][:, :1, :, :]
return items
def _get(self, index):
return self.data[index % self.real_number]