mirror of
https://github.com/hpcaitech/Open-Sora.git
synced 2026-05-21 11:59:01 +02:00
[feat] generate feature
This commit is contained in:
parent
c2d499129f
commit
80f9ecbde3
|
|
@ -59,5 +59,4 @@ text_encoder = dict(
|
|||
save_text_features = True
|
||||
save_compressed_text_features = True
|
||||
bin_size = 10
|
||||
save_dir = f"/mnt/nfs-207/sora_data/feat/test_b{bin_size}"
|
||||
log_time = False
|
||||
|
|
|
|||
|
|
@ -27,8 +27,10 @@ bucket_config = { # 12s/it
|
|||
|
||||
grad_checkpoint = True
|
||||
|
||||
load_text_features = True
|
||||
|
||||
# Acceleration settings
|
||||
num_workers = 8
|
||||
num_workers = 0
|
||||
num_bucket_build_workers = 16
|
||||
dtype = "bf16"
|
||||
plugin = "zero2"
|
||||
|
|
@ -1,3 +1,4 @@
|
|||
import collections
|
||||
import random
|
||||
from typing import Optional
|
||||
|
||||
|
|
@ -8,7 +9,7 @@ from torch.distributed.distributed_c10d import _get_default_group
|
|||
from torch.utils.data import DataLoader
|
||||
|
||||
from .datasets import VariableVideoTextDataset, VideoTextDataset
|
||||
from .sampler import StatefulDistributedSampler, VariableVideoBatchSampler, BatchDistributedSampler
|
||||
from .sampler import BatchDistributedSampler, StatefulDistributedSampler, VariableVideoBatchSampler
|
||||
|
||||
|
||||
# Deterministic dataloader
|
||||
|
|
@ -83,6 +84,7 @@ def prepare_dataloader(
|
|||
else:
|
||||
raise ValueError(f"Unsupported dataset type: {type(dataset)}")
|
||||
|
||||
|
||||
def prepare_variable_dataloader(
|
||||
dataset,
|
||||
batch_size,
|
||||
|
|
@ -173,4 +175,43 @@ def build_batch_dataloader(
|
|||
pin_memory=pin_memory,
|
||||
num_workers=num_workers,
|
||||
**_kwargs,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def collate_fn_default(batch):
|
||||
# HACK: for loading text features
|
||||
use_mask = False
|
||||
if "mask" in batch[0] and isinstance(batch[0]["mask"], int):
|
||||
masks = [x.pop("mask") for x in batch]
|
||||
|
||||
texts = [x.pop("text") for x in batch]
|
||||
texts = torch.cat(texts, dim=1)
|
||||
use_mask = True
|
||||
|
||||
ret = torch.utils.data.default_collate(batch)
|
||||
|
||||
if use_mask:
|
||||
ret["mask"] = masks
|
||||
ret["text"] = texts
|
||||
return ret
|
||||
|
||||
|
||||
def collate_fn_batch(batch):
|
||||
"""
|
||||
Used only with BatchDistributedSampler
|
||||
"""
|
||||
res = torch.utils.data.default_collate(batch)
|
||||
|
||||
# squeeze the first dimension, which is due to torch.stack() in default_collate()
|
||||
if isinstance(res, collections.abc.Mapping):
|
||||
for k, v in res.items():
|
||||
if isinstance(v, torch.Tensor):
|
||||
res[k] = v.squeeze(0)
|
||||
elif isinstance(res, collections.abc.Sequence):
|
||||
res = [x.squeeze(0) if isinstance(x, torch.Tensor) else x for x in res]
|
||||
elif isinstance(res, torch.Tensor):
|
||||
res = res.squeeze(0)
|
||||
else:
|
||||
raise TypeError
|
||||
|
||||
return res
|
||||
|
|
|
|||
|
|
@ -179,8 +179,9 @@ class VariableVideoTextDataset(VideoTextDataset):
|
|||
if self.get_text:
|
||||
ret["text"] = sample["text"]
|
||||
if self.dummy_text_feature:
|
||||
ret["text"] = "dummy text"
|
||||
ret["mask"] = None
|
||||
text_len = 50
|
||||
ret["text"] = torch.zeros((1, text_len, 1152))
|
||||
ret["mask"] = text_len
|
||||
return ret
|
||||
|
||||
def __getitem__(self, index):
|
||||
|
|
@ -205,14 +206,14 @@ class BatchDataset(torch.utils.data.Dataset):
|
|||
def __init__(self):
|
||||
# self.meta = read_file(data_path)
|
||||
# self.path_list = self.meta['path'].tolist()
|
||||
self.path_list = [f'/mnt/nfs-207/sora_data/webvid-10M/feat_text/data/{idx}.bin' for idx in range(5)]
|
||||
self.path_list = [f"/mnt/nfs-207/sora_data/webvid-10M/feat_text/data/{idx}.bin" for idx in range(5)]
|
||||
|
||||
self._len_buffer = len(torch.load(self.path_list[0]))
|
||||
self._num_buffers = len(self.path_list)
|
||||
self.num_samples = self.len_buffer * len(self.path_list)
|
||||
|
||||
self.cur_file_idx = -1
|
||||
|
||||
|
||||
@property
|
||||
def num_buffers(self):
|
||||
return self._num_buffers
|
||||
|
|
@ -220,7 +221,7 @@ class BatchDataset(torch.utils.data.Dataset):
|
|||
@property
|
||||
def len_buffer(self):
|
||||
return self._len_buffer
|
||||
|
||||
|
||||
def _load_buffer(self, idx):
|
||||
file_idx = idx // self.len_buffer
|
||||
if file_idx == self.cur_file_idx:
|
||||
|
|
@ -236,4 +237,3 @@ class BatchDataset(torch.utils.data.Dataset):
|
|||
|
||||
batch = self.cur_buffer[idx % self.len_buffer] # dict; keys are {'x', 'fps'} and text related
|
||||
return batch
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
import os
|
||||
import re
|
||||
import collections
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
|
@ -215,32 +214,3 @@ def resize_crop_to_fill(pil_image, image_size):
|
|||
arr = np.array(image)
|
||||
assert i + th <= arr.shape[0] and j + tw <= arr.shape[1]
|
||||
return Image.fromarray(arr[i : i + th, j : j + tw])
|
||||
|
||||
|
||||
def collate_fn_ignore_none(batch):
|
||||
# we filter out the None values
|
||||
# None value is returned when the get_item fails for an index
|
||||
batch = [val for val in batch if val is not None]
|
||||
return torch.utils.data.default_collate(batch)
|
||||
|
||||
|
||||
def collate_fn_batch(batch):
|
||||
"""
|
||||
Used only with BatchDistributedSampler
|
||||
"""
|
||||
res = torch.utils.data.default_collate(batch)
|
||||
|
||||
# squeeze the first dimension, which is due to torch.stack() in default_collate()
|
||||
if isinstance(res, collections.abc.Mapping):
|
||||
for k, v in res.items():
|
||||
if isinstance(v, torch.Tensor):
|
||||
res[k] = v.squeeze(0)
|
||||
elif isinstance(res, collections.abc.Sequence):
|
||||
res = [x.squeeze(0) if isinstance(x, torch.Tensor) else x for x in res]
|
||||
elif isinstance(res, torch.Tensor):
|
||||
res = res.squeeze(0)
|
||||
else:
|
||||
raise TypeError
|
||||
|
||||
return res
|
||||
|
||||
|
|
|
|||
|
|
@ -370,10 +370,9 @@ class STDiT3(PreTrainedModel):
|
|||
|
||||
# === get y embed ===
|
||||
if self.config.skip_y_embedder:
|
||||
y_lens = mask.tolist()
|
||||
y_lens = mask
|
||||
else:
|
||||
y, y_lens = self.encode_text(y, mask)
|
||||
breakpoint()
|
||||
|
||||
# === get x embed ===
|
||||
x = self.x_embedder(x) # [B, N, C]
|
||||
|
|
|
|||
|
|
@ -3,11 +3,11 @@ from pprint import pformat
|
|||
|
||||
import colossalai
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from tqdm import tqdm
|
||||
|
||||
from opensora.acceleration.parallel_states import get_data_parallel_group
|
||||
from opensora.datasets import prepare_variable_dataloader
|
||||
from opensora.datasets.utils import collate_fn_ignore_none
|
||||
from opensora.acceleration.parallel_states import get_data_parallel_group, set_data_parallel_group
|
||||
from opensora.datasets.dataloader import collate_fn_default, prepare_dataloader
|
||||
from opensora.registry import DATASETS, MODELS, build_module
|
||||
from opensora.utils.config_utils import parse_configs, save_training_config
|
||||
from opensora.utils.misc import FeatureSaver, Timer, create_logger, format_numel_str, get_model_numel, to_torch_dtype
|
||||
|
|
@ -36,6 +36,7 @@ def main():
|
|||
torch.backends.cudnn.allow_tf32 = True
|
||||
|
||||
colossalai.launch_from_torch({})
|
||||
set_data_parallel_group(dist.group.WORLD)
|
||||
|
||||
# == init logger, tensorboard & wandb ==
|
||||
logger = create_logger()
|
||||
|
|
@ -59,14 +60,14 @@ def main():
|
|||
drop_last=True,
|
||||
pin_memory=True,
|
||||
process_group=get_data_parallel_group(),
|
||||
collate_fn=collate_fn_ignore_none,
|
||||
collate_fn=collate_fn_default,
|
||||
)
|
||||
dataloader = prepare_variable_dataloader(
|
||||
dataloader, sampler = prepare_dataloader(
|
||||
bucket_config=cfg.get("bucket_config", None),
|
||||
num_bucket_build_workers=cfg.get("num_bucket_build_workers", 1),
|
||||
**dataloader_args,
|
||||
)
|
||||
num_batch = dataloader.batch_sampler.get_num_batch()
|
||||
num_steps_per_epoch = len(dataloader)
|
||||
|
||||
# ======================================================
|
||||
# 3. build model
|
||||
|
|
@ -107,8 +108,8 @@ def main():
|
|||
save_compressed_text_features = cfg.get("save_compressed_text_features", False)
|
||||
|
||||
# == number of bins ==
|
||||
num_bin = num_batch // bin_size
|
||||
logger.info("Number of batches: %s", num_batch)
|
||||
num_bin = num_steps_per_epoch // bin_size
|
||||
logger.info("Number of batches: %s", num_steps_per_epoch)
|
||||
logger.info("Bin size: %s", bin_size)
|
||||
logger.info("Number of bins: %s", num_bin)
|
||||
|
||||
|
|
|
|||
|
|
@ -4,12 +4,14 @@ set -x
|
|||
set -e
|
||||
|
||||
START_SPLIT=0
|
||||
NUM_SPLIT=100
|
||||
NUM_SPLIT=10
|
||||
|
||||
DATA_PATH=$1
|
||||
SAVE_PATH=$2
|
||||
DATA_ARG="--data-path $DATA_PATH"
|
||||
SAVE_ARG="--save-dir $SAVE_PATH"
|
||||
|
||||
CMD="torchrun --standalone --nproc_per_node 1 scripts/misc/extract_feat.py configs/opensora-v1-2/misc/extract.py $DATA_ARG"
|
||||
CMD="torchrun --standalone --nproc_per_node 1 scripts/misc/extract_feat.py configs/opensora-v1-2/misc/extract.py $DATA_ARG $SAVE_ARG"
|
||||
declare -a GPUS=(0 1 2 3 4 5 6 7)
|
||||
|
||||
mkdir -p logs/extract_feat
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@ from tqdm import tqdm
|
|||
from opensora.acceleration.checkpoint import set_grad_checkpoint
|
||||
from opensora.acceleration.parallel_states import get_data_parallel_group
|
||||
from opensora.datasets import prepare_dataloader
|
||||
from opensora.datasets.utils import collate_fn_ignore_none
|
||||
from opensora.datasets.dataloader import collate_fn_default
|
||||
from opensora.registry import DATASETS, MODELS, SCHEDULERS, build_module
|
||||
from opensora.utils.ckpt_utils import load, model_gathering, model_sharding, record_model_param_shape, save
|
||||
from opensora.utils.config_utils import define_experiment_workspace, parse_configs, save_training_config
|
||||
|
|
@ -99,7 +99,7 @@ def main():
|
|||
drop_last=True,
|
||||
pin_memory=True,
|
||||
process_group=get_data_parallel_group(),
|
||||
collate_fn=collate_fn_ignore_none,
|
||||
collate_fn=collate_fn_default,
|
||||
)
|
||||
dataloader, sampler = prepare_dataloader(
|
||||
bucket_config=cfg.get("bucket_config", None),
|
||||
|
|
|
|||
|
|
@ -14,8 +14,7 @@ from tqdm import tqdm
|
|||
|
||||
from opensora.acceleration.checkpoint import set_grad_checkpoint
|
||||
from opensora.acceleration.parallel_states import get_data_parallel_group
|
||||
from opensora.datasets.dataloader import prepare_dataloader
|
||||
from opensora.datasets.utils import collate_fn_ignore_none
|
||||
from opensora.datasets.dataloader import collate_fn_default, prepare_dataloader
|
||||
from opensora.registry import DATASETS, MODELS, SCHEDULERS, build_module
|
||||
from opensora.utils.ckpt_utils import load, model_gathering, model_sharding, record_model_param_shape, save
|
||||
from opensora.utils.config_utils import define_experiment_workspace, parse_configs, save_training_config
|
||||
|
|
@ -96,7 +95,7 @@ def main():
|
|||
drop_last=True,
|
||||
pin_memory=True,
|
||||
process_group=get_data_parallel_group(),
|
||||
collate_fn=collate_fn_ignore_none,
|
||||
collate_fn=collate_fn_default,
|
||||
)
|
||||
dataloader, sampler = prepare_dataloader(
|
||||
bucket_config=cfg.get("bucket_config", None),
|
||||
|
|
@ -232,7 +231,11 @@ def main():
|
|||
x = vae.encode(x) # [B, C, T, H/P, W/P]
|
||||
# Prepare text inputs
|
||||
if cfg.get("load_text_features", False):
|
||||
model_args = {"y": y.to(device, dtype), "mask": batch.pop("mask").to(device, dtype)}
|
||||
model_args = {"y": y.to(device, dtype)}
|
||||
mask = batch.pop("mask")
|
||||
if isinstance(mask, torch.Tensor):
|
||||
mask = mask.to(device, dtype)
|
||||
model_args["mask"] = mask
|
||||
else:
|
||||
model_args = text_encoder.encode(y)
|
||||
|
||||
|
|
@ -244,7 +247,8 @@ def main():
|
|||
|
||||
# == video meta info ==
|
||||
for k, v in batch.items():
|
||||
model_args[k] = v.to(device, dtype)
|
||||
if isinstance(v, torch.Tensor):
|
||||
model_args[k] = v.to(device, dtype)
|
||||
|
||||
# == diffusion loss computation ==
|
||||
loss_dict = scheduler.training_losses(model, x, model_args, mask=mask)
|
||||
|
|
|
|||
|
|
@ -14,8 +14,7 @@ from tqdm import tqdm
|
|||
|
||||
from opensora.acceleration.checkpoint import set_grad_checkpoint
|
||||
from opensora.acceleration.parallel_states import get_data_parallel_group
|
||||
from opensora.datasets import build_batch_dataloader
|
||||
from opensora.datasets.utils import collate_fn_batch
|
||||
from opensora.datasets.dataloader import build_batch_dataloader, collate_fn_batch
|
||||
from opensora.registry import DATASETS, MODELS, SCHEDULERS, build_module
|
||||
from opensora.utils.ckpt_utils import load, model_gathering, model_sharding, record_model_param_shape, save
|
||||
from opensora.utils.config_utils import define_experiment_workspace, parse_configs, save_training_config
|
||||
|
|
@ -105,15 +104,15 @@ def main():
|
|||
num_steps_per_epoch = len(dataset) // dist.get_world_size()
|
||||
sampler_to_io = None
|
||||
|
||||
'''
|
||||
TODO:
|
||||
"""
|
||||
TODO:
|
||||
- prefetch
|
||||
- collate fn
|
||||
- resume
|
||||
- sampler_to_io ?
|
||||
- remove text_encoder & caption_embedder
|
||||
- currently only support 1 epoch; every epoch is the same
|
||||
'''
|
||||
"""
|
||||
|
||||
# if cfg.dataset.type == DEFAULT_DATASET_NAME:
|
||||
# dataloader = prepare_dataloader(**dataloader_args)
|
||||
|
|
@ -253,8 +252,8 @@ def main():
|
|||
) as pbar:
|
||||
for step, batch in pbar:
|
||||
# modify here
|
||||
x = batch['x'].to(device, dtype) # feat of vae encoder
|
||||
print(step, dist.get_rank(), batch['x'].shape)
|
||||
x = batch["x"].to(device, dtype) # feat of vae encoder
|
||||
print(step, dist.get_rank(), batch["x"].shape)
|
||||
continue
|
||||
|
||||
# x = batch.pop("video").to(device, dtype) # [B, C, T, H, W]
|
||||
|
|
|
|||
|
|
@ -8,10 +8,10 @@ from functools import partial
|
|||
from glob import glob
|
||||
|
||||
import cv2
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import torchvision
|
||||
from PIL import Image
|
||||
from tqdm import tqdm
|
||||
|
||||
from .utils import IMG_EXTENSIONS
|
||||
|
|
@ -84,8 +84,8 @@ def get_info(path):
|
|||
return 0, 0, 0, np.nan, np.nan, np.nan
|
||||
|
||||
|
||||
def get_image_info(path, backend='pillow'):
|
||||
if backend == 'pillow':
|
||||
def get_image_info(path, backend="pillow"):
|
||||
if backend == "pillow":
|
||||
try:
|
||||
with open(path, "rb") as f:
|
||||
img = Image.open(f)
|
||||
|
|
@ -97,7 +97,7 @@ def get_image_info(path, backend='pillow'):
|
|||
return num_frames, height, width, aspect_ratio, fps, hw
|
||||
except:
|
||||
return 0, 0, 0, np.nan, np.nan, np.nan
|
||||
elif backend == 'cv2':
|
||||
elif backend == "cv2":
|
||||
try:
|
||||
im = cv2.imread(path)
|
||||
if im is None:
|
||||
|
|
@ -113,8 +113,8 @@ def get_image_info(path, backend='pillow'):
|
|||
raise ValueError
|
||||
|
||||
|
||||
def get_video_info(path, backend='torchvision'):
|
||||
if backend == 'torchvision':
|
||||
def get_video_info(path, backend="torchvision"):
|
||||
if backend == "torchvision":
|
||||
try:
|
||||
vframes, _, infos = torchvision.io.read_video(filename=path, pts_unit="sec", output_format="TCHW")
|
||||
num_frames, height, width = vframes.shape[0], vframes.shape[2], vframes.shape[3]
|
||||
|
|
@ -127,7 +127,7 @@ def get_video_info(path, backend='torchvision'):
|
|||
return num_frames, height, width, aspect_ratio, fps, hw
|
||||
except:
|
||||
return 0, 0, 0, np.nan, np.nan, np.nan
|
||||
elif backend == 'cv2':
|
||||
elif backend == "cv2":
|
||||
try:
|
||||
cap = cv2.VideoCapture(path)
|
||||
num_frames, height, width, fps = (
|
||||
|
|
@ -603,13 +603,14 @@ def main(args):
|
|||
data = data[data["flow"] >= args.flowmin]
|
||||
if args.remove_text_duplication:
|
||||
data = data.drop_duplicates(subset=["text"], keep="first")
|
||||
print(f"Filtered number of samples: {len(data)}.")
|
||||
|
||||
# process data
|
||||
if args.shuffle:
|
||||
data = data.sample(frac=1).reset_index(drop=True) # shuffle
|
||||
if args.get_first_n_data is not None:
|
||||
data = data.head(args.get_first_n_data)
|
||||
if args.head is not None:
|
||||
data = data.head(args.head)
|
||||
|
||||
print(f"Filtered number of samples: {len(data)}.")
|
||||
|
||||
# shard data
|
||||
if args.shard is not None:
|
||||
|
|
@ -689,7 +690,7 @@ def parse_args():
|
|||
|
||||
# data processing
|
||||
parser.add_argument("--shuffle", default=False, action="store_true", help="shuffle the dataset")
|
||||
parser.add_argument("--get_first_n_data", type=int, default=None, help="return the first n rows of data")
|
||||
parser.add_argument("--head", type=int, default=None, help="return the first n rows of data")
|
||||
|
||||
return parser.parse_args()
|
||||
|
||||
|
|
@ -768,8 +769,8 @@ def get_output_path(args, input_name):
|
|||
# processing
|
||||
if args.shuffle:
|
||||
name += f"_shuffled_seed{args.seed}"
|
||||
if args.get_first_n_data is not None:
|
||||
name += f"_first_{args.get_first_n_data}_data"
|
||||
if args.head is not None:
|
||||
name += f"_first_{args.head}_data"
|
||||
|
||||
output_path = os.path.join(dir_path, f"{name}.{args.format}")
|
||||
return output_path
|
||||
|
|
|
|||
Loading…
Reference in a new issue