[feat] generate feature

This commit is contained in:
zhengzangw 2024-05-21 04:05:02 +00:00
parent c2d499129f
commit 80f9ecbde3
12 changed files with 97 additions and 79 deletions

View file

@ -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

View file

@ -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"

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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]

View file

@ -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)

View file

@ -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

View file

@ -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),

View file

@ -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)

View file

@ -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]

View file

@ -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