Improve dataset preparation support + multiresolution prep (#39)

* update

* make style

* renormalize correctly

* apply suggestions from review

* apply suggestions from review

* update
This commit is contained in:
Aryan
2024-10-17 16:48:52 +05:30
committed by GitHub
parent a6c246c29d
commit feb2e26547
5 changed files with 427 additions and 363 deletions
+2 -2
View File
@@ -46,7 +46,7 @@ from diffusers import export_to_video
pipe = CogVideoXPipeline.from_pretrained(
"THUDM/CogVideoX-5b", torch_dtype=torch.bfloat16
).to("cuda")
+ pipe.load_lora_weights("my-awesome-name/my-awesome-lora", adapter_name=["cogvideox-lora"])
+ pipe.load_lora_weights("my-awesome-name/my-awesome-lora", adapter_name="cogvideox-lora")
+ pipe.set_adapters(["cogvideox-lora"], [1.0])
video = pipe("<my-awesome-prompt>").frames[0]
@@ -429,7 +429,7 @@ With `train_batch_size = 4`:
- [ ] Make scripts compatible with FSDP
- [x] Make scripts compatible with DeepSpeed
- [ ] vLLM-powered captioning script
- [ ] Multi-resolution/frame support in `prepare_dataset.py`
- [x] Multi-resolution/frame support in `prepare_dataset.py`
- [ ] Analyzing traces for potential speedups and removing as many syncs as possible
- [ ] Support for QLoRA (priority), and other types of high usage LoRAs methods
- [x] Test scripts with memory-efficient optimizer from bitsandbytes
+1 -1
View File
@@ -440,7 +440,7 @@ diffusers(该分支为 CogVideoX 的图像到视频添加了 LoRA 加载支持
- [ ] 使脚本兼容 FSDP
- [x] 使脚本兼容 DeepSpeed
- [ ] 基于 vLLM 的字幕脚本
- [ ] 在 `prepare_dataset.py` 中支持多分辨率/帧数
- [x] 在 `prepare_dataset.py` 中支持多分辨率/帧数
- [ ] 分析性能瓶颈并尽可能减少同步操作
- [ ] 支持 QLoRA(优先),以及其他高使用率的 LoRA 方法
- [x] 使用 bitsandbytes 的节省内存优化器测试脚本
+3 -1
View File
@@ -200,6 +200,7 @@ def log_validation(
return videos
class CollateFunction:
def __init__(self, weight_dtype, load_tensors):
self.weight_dtype = weight_dtype
@@ -223,6 +224,7 @@ class CollateFunction:
"prompts": prompts,
}
def main(args):
if args.report_to == "wandb" and args.hub_token is not None:
raise ValueError(
@@ -647,7 +649,7 @@ def main(args):
# Encode videos
if not args.load_tensors:
images = images.permute(0, 2, 1, 3, 4) # [B, C, F, H, W]
images = images.permute(0, 2, 1, 3, 4) # [B, C, F, H, W]
image_noise_sigma = torch.normal(
mean=-3.0, std=0.5, size=(images.size(0),), device=accelerator.device, dtype=weight_dtype
)
+38 -5
View File
@@ -78,15 +78,16 @@ class VideoDataset(Dataset):
self.video_paths,
) = self._load_dataset_from_csv()
self.num_videos = len(self.video_paths)
if self.num_videos != len(self.prompts):
if len(self.video_paths) != len(self.prompts):
raise ValueError(
f"Expected length of prompts and videos to be the same but found {len(self.prompts)=} and {len(self.video_paths)=}. Please ensure that the number of caption prompts and videos match in your dataset."
)
self.video_transforms = transforms.Compose(
[
transforms.RandomHorizontalFlip(random_flip) if random_flip else transforms.Lambda(self.identity_transform),
transforms.RandomHorizontalFlip(random_flip)
if random_flip
else transforms.Lambda(self.identity_transform),
transforms.Lambda(self.scale_transform),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
]
@@ -101,7 +102,7 @@ class VideoDataset(Dataset):
return x / 255.0
def __len__(self) -> int:
return self.num_videos
return len(self.video_paths)
def __getitem__(self, index: int) -> Dict[str, Any]:
if isinstance(index, list):
@@ -358,10 +359,30 @@ class VideoDatasetWithResizeAndRectangleCrop(VideoDataset):
class BucketSampler(Sampler):
def __init__(self, data_source: VideoDataset, batch_size: int = 8, shuffle: bool = True) -> None:
r"""
PyTorch Sampler that groups 3D data by height, width and frames.
Args:
data_source (`VideoDataset`):
A PyTorch dataset object that is an instance of `VideoDataset`.
batch_size (`int`, defaults to `8`):
The batch size to use for training.
shuffle (`bool`, defaults to `True`):
Whether or not to shuffle the data in each batch before dispatching to dataloader.
drop_last (`bool`, defaults to `False`):
Whether or not to drop incomplete buckets of data after completely iterating over all data
in the dataset. If set to True, only batches that have `batch_size` number of entries will
be yielded. If set to False, it is guaranteed that all data in the dataset will be processed
and batches that do not have `batch_size` number of entries will also be yielded.
"""
def __init__(
self, data_source: VideoDataset, batch_size: int = 8, shuffle: bool = True, drop_last: bool = False
) -> None:
self.data_source = data_source
self.batch_size = batch_size
self.shuffle = shuffle
self.drop_last = drop_last
self.buckets = {resolution: [] for resolution in data_source.resolutions}
@@ -377,3 +398,15 @@ class BucketSampler(Sampler):
yield self.buckets[(f, h, w)]
del self.buckets[(f, h, w)]
self.buckets[(f, h, w)] = []
if self.drop_last:
return
for fhw, bucket in list(self.buckets.items()):
if len(bucket) == 0:
continue
if self.shuffle:
random.shuffle(bucket)
yield bucket
del self.buckets[fhw]
self.buckets[fhw] = []
+383 -354
View File
@@ -1,28 +1,32 @@
#!/usr/bin/env python3
import argparse
import gc
import functools
import json
import os
import pathlib
import queue
import traceback
from typing import Any, Dict, List, Optional, Tuple, Union
import uuid
from concurrent.futures import ThreadPoolExecutor
from typing import Any, Dict, List, Optional, Union
import numpy as np
import pandas as pd
import torch
import torch.distributed as dist
import torchvision.transforms as TT
from diffusers import AutoencoderKLCogVideoX
from diffusers.training_utils import set_seed
from diffusers.utils import export_to_video, get_logger
from torch.utils.data import DataLoader
from torchvision import transforms
from torchvision.transforms import InterpolationMode
from torchvision.transforms.functional import resize
from tqdm import tqdm
from transformers import T5EncoderModel, T5Tokenizer
import decord # isort:skip
from dataset import BucketSampler, VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop # isort:skip
decord.bridge.set_bridge("torch")
logger = get_logger(__name__)
@@ -34,6 +38,27 @@ DTYPE_MAPPING = {
}
def check_height(x: Any) -> int:
x = int(x)
if x % 16 != 0:
raise argparse.ArgumentTypeError(f"`--height_buckets` must be divisible by 16, but got {x} which does not fit criteria.")
return x
def check_width(x: Any) -> int:
x = int(x)
if x % 16 != 0:
raise argparse.ArgumentTypeError(f"`--width_buckets` must be divisible by 16, but got {x} which does not fit criteria.")
return x
def check_frames(x: Any) -> int:
x = int(x)
if x % 4 != 0 and x % 4 != 1:
raise argparse.ArgumentTypeError(f"`--frames_buckets` must be of form `4 * k` or `4 * k + 1`, but got {x} which does not fit criteria.")
return x
def get_args() -> Dict[str, Any]:
parser = argparse.ArgumentParser()
parser.add_argument(
@@ -58,6 +83,53 @@ def get_args() -> Dict[str, Any]:
default="video",
help="If using a CSV file via the `--dataset_file` argument, this should be the name of the column containing the video paths. If using the folder structure format for data loading, this should be the name of the file containing line-separated video paths (the file should be located in `--data_root`).",
)
parser.add_argument(
"--id_token",
type=str,
default=None,
help="Identifier token appended to the start of each prompt if provided.",
)
parser.add_argument(
"--height_buckets",
nargs="+",
type=check_height,
default=[256, 320, 384, 480, 512, 576, 720, 768, 960, 1024, 1280, 1536],
)
parser.add_argument(
"--width_buckets",
nargs="+",
type=check_width,
default=[256, 320, 384, 480, 512, 576, 720, 768, 960, 1024, 1280, 1536],
)
parser.add_argument(
"--frame_buckets",
nargs="+",
type=check_frames,
default=[49],
)
parser.add_argument(
"--random_flip",
type=float,
default=None,
help="If random horizontal flip augmentation is to be used, this should be the flip probability.",
)
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=0,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--pin_memory",
action="store_true",
help="Whether or not to use the pinned memory setting in pytorch dataloader.",
)
parser.add_argument(
"--video_reshape_mode",
type=str,
default=None,
help="All input videos are reshaped to this mode. Choose between ['center', 'random', 'none']",
)
parser.add_argument(
"--save_image_latents",
action="store_true",
@@ -69,8 +141,6 @@ def get_args() -> Dict[str, Any]:
required=True,
help="Path to output directory where preprocessed videos/latents/embeddings will be saved.",
)
parser.add_argument("--height", type=int, default=480, help="Height of the resized output video.")
parser.add_argument("--width", type=int, default=720, help="Width of the resized output video.")
parser.add_argument("--max_num_frames", type=int, default=49, help="Maximum number of frames in output video.")
parser.add_argument(
"--max_sequence_length", type=int, default=226, help="Max sequence length of prompt embeddings."
@@ -107,122 +177,13 @@ def get_args() -> Dict[str, Any]:
default="fp32",
help="Data type to use when generating latents and prompt embeddings.",
)
parser.add_argument("--seed", type=int, default=42, help="Seed for reproducibility.")
parser.add_argument(
"--num_artifact_workers", type=int, default=4, help="Number of worker threads for serializing artifacts."
)
return parser.parse_args()
def load_dataset_from_local_path(
data_root: pathlib.Path, caption_column: str, video_column: str
) -> Tuple[List[str], List[pathlib.Path]]:
if not data_root.exists():
raise ValueError("Root folder for videos does not exist")
prompt_path = data_root.joinpath(caption_column)
video_path = data_root.joinpath(video_column)
if not prompt_path.exists() or not prompt_path.is_file():
raise ValueError(
"Expected `--caption_column` to be path to a file in `--data_root` containing line-separated text prompts."
)
if not video_path.exists() or not video_path.is_file():
raise ValueError(
"Expected `--video_column` to be path to a file in `--data_root` containing line-separated paths to video data in the same directory."
)
with open(prompt_path, "r", encoding="utf-8") as file:
prompts = [line.strip() for line in file.readlines() if len(line.strip()) > 0]
with open(video_path, "r", encoding="utf-8") as file:
video_paths = [data_root.joinpath(line.strip()) for line in file.readlines() if len(line.strip()) > 0]
if any(not path.is_file() for path in video_paths):
raise ValueError(
f"Expected `{video_column}` to be a path to a file in `{data_root}` containing line-separated paths to video data but found at least one path that is not a valid file."
)
return prompts, video_paths
def load_dataset_from_csv(
data_root: pathlib.Path, dataset_file: pathlib.Path, caption_column: str, video_column: str
) -> Tuple[List[str], List[pathlib.Path]]:
df = pd.read_csv(dataset_file)
prompts = df[caption_column].tolist()
video_paths = df[video_column].tolist()
video_paths = [data_root.joinpath(line.strip()) for line in video_paths]
if any(not path.is_file() for path in video_paths):
raise ValueError(
f"Expected `{video_column}` to be a path to a file in `{data_root}` containing line-separated paths to video data but found at least one path that is not a valid file."
)
return prompts, video_paths
def resize_for_rectangle_crop(arr, height, width, reshape_mode):
image_size = height, width
if arr.shape[3] / arr.shape[2] > image_size[1] / image_size[0]:
arr = resize(
arr,
size=[image_size[0], int(arr.shape[3] * image_size[0] / arr.shape[2])],
interpolation=InterpolationMode.BICUBIC,
)
else:
arr = resize(
arr,
size=[int(arr.shape[2] * image_size[1] / arr.shape[3]), image_size[1]],
interpolation=InterpolationMode.BICUBIC,
)
h, w = arr.shape[2], arr.shape[3]
arr = arr.squeeze(0)
delta_h = h - image_size[0]
delta_w = w - image_size[1]
if reshape_mode == "random" or reshape_mode == "none":
top = np.random.randint(0, delta_h + 1)
left = np.random.randint(0, delta_w + 1)
elif reshape_mode == "center":
top, left = delta_h // 2, delta_w // 2
else:
raise NotImplementedError
arr = TT.functional.crop(arr, top=top, left=left, height=image_size[0], width=image_size[1])
return arr
def load_and_preprocess_video(
path: pathlib.Path,
height: int,
width: int,
max_num_frames: int,
video_transforms,
num_threads: int = 0,
video_reshape_mode: str = "center",
) -> Optional[torch.Tensor]:
frames = None
try:
video_reader = decord.VideoReader(uri=path.as_posix(), num_threads=num_threads)
video_num_frames = len(video_reader)
if video_num_frames < max_num_frames:
logger.warning(
f"Video at `{path.as_posix()}` should have at least `{max_num_frames=}`, but got only `{video_num_frames=}`. Skipping it."
)
return None
indices = list(range(0, video_num_frames, max(video_num_frames // max_num_frames, 1)))
frames: torch.Tensor = video_reader.get_batch(indices)
frames = frames[:max_num_frames].float()
frames = frames.permute(0, 3, 1, 2).contiguous()
frames = resize_for_rectangle_crop(frames, height, width, video_reshape_mode)
frames = torch.stack([video_transforms(frame) for frame in frames], dim=0)
except Exception as e:
logger.error(f"Error: {e}. Skipping video located at `{path.as_posix()}`")
traceback.print_exc()
return frames
def _get_t5_prompt_embeds(
tokenizer: T5Tokenizer,
text_encoder: T5EncoderModel,
@@ -318,94 +279,93 @@ def compute_prompt_embeddings(
return prompt_embeds
def save_videos(
videos: torch.Tensor,
video_paths: List[pathlib.Path],
prompts: List[str],
output_dir: pathlib.Path,
target_fps: int = 8,
to_pil_image = transforms.ToPILImage(mode="RGB")
def save_image(image: torch.Tensor, path: pathlib.Path) -> None:
image = to_pil_image(image)
image.save(path)
def save_video(video: torch.Tensor, path: pathlib.Path, fps: int = 8) -> None:
video = [to_pil_image(frame) for frame in video]
export_to_video(video, path, fps=fps)
def save_prompt(prompt: str, path: pathlib.Path) -> None:
with open(path, "w", encoding="utf=8") as file:
file.write(prompt)
@torch.no_grad()
def serialize_artifacts(
batch_size: int,
fps: int,
images_dir: Optional[pathlib.Path] = None,
image_latents_dir: Optional[pathlib.Path] = None,
videos_dir: Optional[pathlib.Path] = None,
video_latents_dir: Optional[pathlib.Path] = None,
prompts_dir: Optional[pathlib.Path] = None,
prompt_embeds_dir: Optional[pathlib.Path] = None,
images: Optional[torch.Tensor] = None,
image_latents: Optional[torch.Tensor] = None,
videos: Optional[torch.Tensor] = None,
video_latents: Optional[torch.Tensor] = None,
prompts: Optional[List[str]] = None,
prompt_embeds: Optional[torch.Tensor] = None,
) -> None:
assert videos.size(0) == len(video_paths)
if images is not None:
images = (images.permute(0, 2, 1, 3, 4) + 1) / 2
if videos is not None:
videos = (videos.permute(0, 2, 1, 3, 4) + 1) / 2
videos = (videos + 1) / 2
videos = (videos * 255.0).clip(0, 255)
videos = videos.to(dtype=torch.uint8)
data_folder_mapper_list = [
(images, images_dir, lambda img, path: save_image(img[0], path), "png"),
(image_latents, image_latents_dir, torch.save, "pt"),
(videos, videos_dir, functools.partial(save_video, fps=fps), "mp4"),
(video_latents, video_latents_dir, torch.save, "pt"),
(prompts, prompts_dir, save_prompt, "txt"),
(prompt_embeds, prompt_embeds_dir, torch.save, "pt"),
]
filenames = [uuid.uuid4() for _ in range(batch_size)]
video_dir = output_dir.joinpath("videos")
output_dir.mkdir(parents=True, exist_ok=True)
video_dir.mkdir(parents=True, exist_ok=True)
to_pil_image = transforms.ToPILImage()
videos_pil = [[to_pil_image(frame) for frame in video] for video in videos]
for video, video_path in zip(videos_pil, video_paths):
filename = video_dir.joinpath(video_path.name)
logger.debug(f"Saving video to `{filename}`")
export_to_video(video, filename.as_posix(), fps=target_fps)
with open(output_dir.joinpath("videos.txt").as_posix(), "a", encoding="utf-8") as file:
for video_path in video_paths:
file.write(f"videos/{video_path.name}\n")
with open(output_dir.joinpath("prompt.txt").as_posix(), "a", encoding="utf-8") as file:
for prompt in prompts:
file.write(f"{prompt}\n")
for data, folder, save_fn, extension in data_folder_mapper_list:
if data is None:
continue
for slice, filename in zip(data, filenames):
if isinstance(slice, torch.Tensor):
slice = slice.clone()
path = folder.joinpath(f"{filename}.{extension}")
save_fn(slice, path)
def save_latents_and_embeddings(
image_latents: torch.Tensor,
latents: torch.Tensor,
prompt_embeds: torch.Tensor,
video_paths: List[pathlib.Path],
prompts: List[str],
output_dir: pathlib.Path,
save_image_latents: bool = False,
) -> None:
assert latents.size(0) == prompt_embeds.size(0)
assert latents.size(0) == len(video_paths)
assert prompt_embeds.size(0) == len(prompts)
if save_image_latents:
assert image_latents.size(0) == latents.size(0)
else:
image_latents = [None] * latents.size(0)
def save_intermediates(output_queue: queue.Queue) -> None:
while True:
try:
item = output_queue.get(timeout=30)
if item is None:
break
serialize_artifacts(**item)
image_latents_dir = output_dir.joinpath("image_latents")
latents_dir = output_dir.joinpath("latents")
embeds_dir = output_dir.joinpath("embeddings")
output_dir.mkdir(parents=True, exist_ok=True)
image_latents_dir.mkdir(parents=True, exist_ok=True)
latents_dir.mkdir(parents=True, exist_ok=True)
embeds_dir.mkdir(parents=True, exist_ok=True)
for image_latent, latent, embed, video_path in zip(image_latents, latents, prompt_embeds, video_paths):
image_latent = image_latent.clone()
latent = latent.clone()
embed = embed.clone()
filename_without_ext = video_path.stem
image_latent_filename = image_latents_dir.joinpath(f"{filename_without_ext}.pt")
latent_filename = latents_dir.joinpath(f"{filename_without_ext}.pt")
embed_filename = embeds_dir.joinpath(f"{filename_without_ext}.pt")
torch.save(image_latent, image_latent_filename)
torch.save(latent, latent_filename)
torch.save(embed, embed_filename)
with open(output_dir.joinpath("videos.txt").as_posix(), "a", encoding="utf-8") as file:
for video_path in video_paths:
file.write(f"videos/{video_path.name}\n")
with open(output_dir.joinpath("prompt.txt").as_posix(), "a", encoding="utf-8") as file:
for prompt in prompts:
file.write(f"{prompt}\n")
except queue.Empty:
continue
@torch.no_grad()
def main():
args = get_args()
set_seed(args.seed)
output_dir = pathlib.Path(args.output_dir)
tmp_dir = output_dir.joinpath("tmp")
output_dir.mkdir(parents=True, exist_ok=True)
tmp_dir.mkdir(parents=True, exist_ok=True)
# Create task queue for non-blocking serializing of artifacts
output_queue = queue.Queue()
save_thread = ThreadPoolExecutor(max_workers=args.num_artifact_workers)
save_future = save_thread.submit(save_intermediates, output_queue)
# Initialize distributed processing
if "LOCAL_RANK" in os.environ:
@@ -419,192 +379,261 @@ def main():
local_rank = 0
world_size = 1
rank = 0
torch.cuda.set_device(local_rank)
torch.cuda.set_device(rank)
data_root = pathlib.Path(args.data_root)
dataset_file = None
if args.dataset_file:
dataset_file = pathlib.Path(args.dataset_file)
# Create folders where intermediate tensors from each rank will be saved
images_dir = tmp_dir.joinpath(f"images/{rank}")
image_latents_dir = tmp_dir.joinpath(f"image_latents/{rank}")
videos_dir = tmp_dir.joinpath(f"videos/{rank}")
video_latents_dir = tmp_dir.joinpath(f"video_latents/{rank}")
prompts_dir = tmp_dir.joinpath(f"prompts/{rank}")
prompt_embeds_dir = tmp_dir.joinpath(f"prompt_embeds/{rank}")
if dataset_file is None:
prompts, video_paths = load_dataset_from_local_path(data_root, args.caption_column, args.video_column)
images_dir.mkdir(parents=True, exist_ok=True)
image_latents_dir.mkdir(parents=True, exist_ok=True)
videos_dir.mkdir(parents=True, exist_ok=True)
video_latents_dir.mkdir(parents=True, exist_ok=True)
prompts_dir.mkdir(parents=True, exist_ok=True)
prompt_embeds_dir.mkdir(parents=True, exist_ok=True)
weight_dtype = DTYPE_MAPPING[args.dtype]
target_fps = args.target_fps
# 1. Dataset
dataset_init_kwargs = {
"data_root": args.data_root,
"dataset_file": args.dataset_file,
"caption_column": args.caption_column,
"video_column": args.video_column,
"max_num_frames": args.max_num_frames,
"id_token": args.id_token,
"height_buckets": args.height_buckets,
"width_buckets": args.width_buckets,
"frame_buckets": args.frame_buckets,
"load_tensors": False,
"random_flip": args.random_flip,
"image_to_video": args.save_image_latents,
}
if args.video_reshape_mode is None:
dataset = VideoDatasetWithResizing(**dataset_init_kwargs)
else:
prompts, video_paths = load_dataset_from_csv(data_root, dataset_file, args.caption_column, args.video_column)
video_transforms = transforms.Compose(
[
transforms.Lambda(lambda x: x / 255.0),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
]
)
# Preprocess videos with progress bar
prompts_usable = []
video_paths_usable = []
videos = []
# Only show progress bar on the main process
if rank == 0:
iterator = tqdm(zip(prompts, video_paths), total=len(prompts), desc="Load and Preprocess videos")
else:
iterator = zip(prompts, video_paths)
for prompt, path in iterator:
video = load_and_preprocess_video(
path, args.height, args.width, args.max_num_frames, video_transforms, args.num_decode_threads
dataset = VideoDatasetWithResizeAndRectangleCrop(
video_reshape_mode=args.video_reshape_mode, **dataset_init_kwargs
)
if video is not None:
prompts_usable.append(prompt)
video_paths_usable.append(path)
videos.append(video)
if len(videos) == 0:
logger.error("No usable videos found after preprocessing.")
return
videos = torch.stack(videos)
original_dataset_size = len(dataset)
# Split data among GPUs
if world_size > 1:
total_samples = len(prompts_usable)
samples_per_gpu = total_samples // world_size
samples_per_gpu = original_dataset_size // world_size
start_index = rank * samples_per_gpu
end_index = start_index + samples_per_gpu
if rank == world_size - 1:
end_index = total_samples # Make sure the last GPU gets the remaining data
end_index = original_dataset_size # Make sure the last GPU gets the remaining data
# Slice the data
prompts_usable = prompts_usable[start_index:end_index]
video_paths_usable = video_paths_usable[start_index:end_index]
videos = videos[start_index:end_index]
dataset.prompts = dataset.prompts[start_index:end_index]
dataset.video_paths = dataset.video_paths[start_index:end_index]
else:
pass
device = torch.device(f"cuda:{local_rank}")
rank_dataset_size = len(dataset)
if not args.save_tensors:
save_videos(videos, video_paths_usable, prompts_usable, pathlib.Path(args.output_dir), args.target_fps)
# 2. Dataloader
def collate_fn(data):
prompts = [x["prompt"] for x in data[0]]
images = [x["image"] for x in data[0]]
images = torch.stack(images).to(dtype=weight_dtype, non_blocking=True)
videos = [x["video"] for x in data[0]]
videos = torch.stack(videos).to(dtype=weight_dtype, non_blocking=True)
return {
"images": images,
"videos": videos,
"prompts": prompts,
}
dataloader = DataLoader(
dataset,
batch_size=1,
sampler=BucketSampler(dataset, batch_size=args.batch_size, shuffle=True, drop_last=False),
collate_fn=collate_fn,
num_workers=args.dataloader_num_workers,
pin_memory=args.pin_memory,
)
# 3. Prepare models
device = f"cuda:{rank}"
generator = torch.Generator(device).manual_seed(args.seed)
tokenizer = T5Tokenizer.from_pretrained(args.model_id, subfolder="tokenizer")
text_encoder = T5EncoderModel.from_pretrained(args.model_id, subfolder="text_encoder", torch_dtype=weight_dtype)
text_encoder = text_encoder.to(device)
vae = AutoencoderKLCogVideoX.from_pretrained(args.model_id, subfolder="vae", torch_dtype=weight_dtype)
vae = vae.to(device)
if args.use_slicing:
vae.enable_slicing()
if args.use_tiling:
vae.enable_tiling()
# 4. Compute latents and embeddings and save
if rank == 0:
iterator = tqdm(
dataloader, desc="Encoding", total=(rank_dataset_size + args.batch_size - 1) // args.batch_size
)
else:
dtype = DTYPE_MAPPING[args.dtype]
tokenizer = T5Tokenizer.from_pretrained(args.model_id, subfolder="tokenizer")
text_encoder = T5EncoderModel.from_pretrained(args.model_id, subfolder="text_encoder", torch_dtype=dtype)
text_encoder = text_encoder.to(device)
iterator = dataloader
prompt_embeds_list = []
for step, batch in enumerate(iterator):
try:
images = None
if args.save_image_latents:
images = batch["images"].to(device, non_blocking=True)
videos = batch["videos"].to(device, non_blocking=True)
prompts = batch["prompts"]
if rank == 0:
iterator = tqdm(range(0, len(prompts_usable), args.batch_size), desc="Encoding prompts")
else:
iterator = range(0, len(prompts_usable), args.batch_size)
# Encode videos & images
image_latents = None
if args.save_image_latents:
images = images.permute(0, 2, 1, 3, 4) # [B, C, F, H, W]
image_noise_sigma = torch.normal(
mean=-3.0, std=0.5, size=(images.size(0),), generator=generator, device=device, dtype=weight_dtype
)
image_noise_sigma = torch.exp(image_noise_sigma)
noisy_images = images + torch.empty_like(images).normal_(generator=generator) * image_noise_sigma[:, None, None, None, None]
image_latent_dist = vae.encode(noisy_images).latent_dist
image_latents = image_latent_dist.sample() * vae.config.scaling_factor
image_latents = image_latents.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
image_latents = image_latents.to(memory_format=torch.contiguous_format, dtype=weight_dtype)
for start_index in iterator:
end_index = min(len(prompts_usable), start_index + args.batch_size)
batch_prompts = prompts_usable[start_index:end_index]
videos = videos.permute(0, 2, 1, 3, 4) # [B, C, F, H, W]
latent_dist = vae.encode(videos).latent_dist
video_latents = latent_dist.sample(generator=generator) * vae.config.scaling_factor
video_latents = video_latents.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
video_latents = video_latents.to(memory_format=torch.contiguous_format, dtype=weight_dtype)
# Encode prompts
prompt_embeds = compute_prompt_embeddings(
tokenizer,
text_encoder,
batch_prompts,
max_sequence_length=args.max_sequence_length,
device=device,
dtype=dtype,
)
prompt_embeds_list.append(prompt_embeds.to("cpu"))
prompt_embeds = None
if len(prompt_embeds_list) > 0:
prompt_embeds = torch.cat(prompt_embeds_list)
del tokenizer, text_encoder
gc.collect()
torch.cuda.empty_cache()
torch.cuda.synchronize(device)
vae = AutoencoderKLCogVideoX.from_pretrained(args.model_id, subfolder="vae", torch_dtype=dtype)
vae = vae.to(device)
if args.use_slicing:
vae.enable_slicing()
if args.use_tiling:
vae.enable_tiling()
encoded_videos_list = []
encoded_images_list = []
if rank == 0:
iterator = tqdm(range(0, len(video_paths_usable), args.batch_size), desc="Encoding videos")
else:
iterator = range(0, len(video_paths_usable), args.batch_size)
for start_index in iterator:
end_index = min(len(video_paths_usable), start_index + args.batch_size)
batch_videos = videos[start_index:end_index]
batch_videos = batch_videos.to(device)
batch_videos = batch_videos.permute(0, 2, 1, 3, 4) # [B, C, F, H, W]
if args.save_image_latents:
batch_images = batch_videos[:, :, :1].clone()
if args.use_slicing:
encoded_slices = [vae._encode(video_slice) for video_slice in batch_videos.split(1)]
encoded_video = torch.cat(encoded_slices)
encoded_videos_list.append(encoded_video.to("cpu"))
if args.save_image_latents:
encoded_slices = [vae._encode(image_slice) for image_slice in batch_images.split(1)]
encoded_image = torch.cat(encoded_slices)
encoded_images_list.append(encoded_image.to("cpu"))
else:
encoded_video = vae._encode(batch_videos)
encoded_videos_list.append(encoded_video.to("cpu"))
if args.save_image_latents:
encoded_image = vae._encode(batch_images)
encoded_images_list.append(encoded_image.to("cpu"))
encoded_videos = None
if len(encoded_videos_list) > 0:
encoded_videos = torch.cat(encoded_videos_list)
encoded_images = None
if len(encoded_images_list) > 0:
encoded_images = torch.cat(encoded_images_list)
del vae
gc.collect()
torch.cuda.empty_cache()
torch.cuda.synchronize(device)
# Ensure that only one process creates the output directories
if world_size > 1:
dist.barrier()
if prompt_embeds is not None:
assert encoded_videos is not None
save_latents_and_embeddings(
encoded_images,
encoded_videos,
prompt_embeds,
video_paths_usable,
prompts_usable,
pathlib.Path(args.output_dir),
args.save_image_latents,
prompts,
args.max_sequence_length,
device,
weight_dtype,
requires_grad=False,
)
# Finalize distributed processing
output_queue.put(
{
"batch_size": prompt_embeds.shape[0],
"fps": target_fps,
"images_dir": images_dir,
"image_latents_dir": image_latents_dir,
"videos_dir": videos_dir,
"video_latents_dir": video_latents_dir,
"prompts_dir": prompts_dir,
"prompt_embeds_dir": prompt_embeds_dir,
"images": images,
"image_latents": image_latents,
"videos": videos,
"video_latents": video_latents,
"prompts": prompts,
"prompt_embeds": prompt_embeds,
}
)
except Exception:
print("-------------------------")
print(f"An exception occurred while processing data: {rank=}, {world_size=}, {step=}")
traceback.print_exc()
print("-------------------------")
# 5. Complete distributed processing
if world_size > 1:
dist.barrier()
dist.destroy_process_group()
output_queue.put(None)
save_thread.shutdown(wait=True)
save_future.result()
# 6. Combine results from each rank
if rank == 0:
print(
f"Completed preprocessing latents and embeddings. Temporary files from all ranks saved to `{tmp_dir.as_posix()}`"
)
# Move files from each rank to common directory
for subfolder, extension in [
("images", "png"),
("image_latents", "pt"),
("videos", "mp4"),
("video_latents", "pt"),
("prompts", "txt"),
("prompt_embeds", "pt"),
]:
tmp_subfolder = tmp_dir.joinpath(subfolder)
combined_subfolder = output_dir.joinpath(subfolder)
combined_subfolder.mkdir(parents=True, exist_ok=True)
pattern = f"*.{extension}"
for file in tmp_subfolder.rglob(pattern):
file.replace(combined_subfolder / file.name)
# Remove temporary directories
def rmdir_recursive(dir: pathlib.Path) -> None:
for child in dir.iterdir():
if child.is_file():
child.unlink()
else:
rmdir_recursive(child)
dir.rmdir()
rmdir_recursive(tmp_dir)
# Combine prompts and videos into individual text files and single jsonl
prompts_folder = output_dir.joinpath("prompts")
prompts = []
stems = []
for filename in prompts_folder.rglob("*.txt"):
with open(filename, "r") as file:
prompts.append(file.read().strip())
stems.append(filename.stem)
prompts_txt = output_dir.joinpath("prompts.txt")
videos_txt = output_dir.joinpath("videos.txt")
data_jsonl = output_dir.joinpath("data.jsonl")
with open(prompts_txt, "w") as file:
for prompt in prompts:
file.write(f"{prompt}\n")
with open(videos_txt, "w") as file:
for stem in stems:
file.write(f"videos/{stem}.pt\n")
with open(data_jsonl, "w") as file:
for prompt, stem in zip(prompts, stems):
data = {
"prompt": prompt,
"prompt_embed": f"prompt_embeds/{stem}.pt",
"image": f"images/{stem}.png",
"image_latent": f"image_latents/{stem}.pt",
"video": f"videos/{stem}.mp4",
"video_latent": f"video_latents/{stem}.pt",
}
file.write(json.dumps(data) + "\n")
print(
f"Completed preprocessing. All files saved to `{output_dir.as_posix()}`"
)
if __name__ == "__main__":
args = get_args()
assert args.height % 16 == 0, "CogVideoX requires input video height to be divisible by 16."
assert args.width % 16 == 0, "CogVideoX requires input video width to be divisible by 16."
assert (
args.max_num_frames % 4 == 0 or args.max_num_frames % 4 == 1
), "`--max_num_frames` must be of form 4 * k or 4 * k + 1 to be compatible with VAE."
main()