improve dataset preparation (#43)

* update

* update

* update

* update
This commit is contained in:
Aryan
2024-10-18 02:43:22 +05:30
committed by GitHub
parent 8c12f34d4e
commit f2a1626de2
8 changed files with 129 additions and 75 deletions
+2 -2
View File
@@ -18,7 +18,7 @@ The framework supports resolutions and frame counts that meet the following cond
- Any resolution as long as it is divisible by 32. For example, `720 * 480`, `1920 * 1020`, etc.
- **Supported Frame Counts (Frames)**:
- Must satisfy (4K + 1), i.e., multiples of 4 such as 16, 24, 32, 48, 64, 80.
- Must be `4 * k` or `4 * k + 1` (example: 16, 32, 49, 81)
It is recommended to place all videos in a single folder.
@@ -58,4 +58,4 @@ huggingface-cli download --repo-type dataset Wild-Heart/Disney-VideoGeneration-D
This dataset has been prepared in the expected format and can be used directly. However, directly using the video dataset may cause Out of Memory (OOM) issues on GPUs with smaller VRAM because it requires loading the [VAE](https://huggingface.co/THUDM/CogVideoX-5b/tree/main/vae) (which encodes videos into latent space) and the large [T5-XXL](https://huggingface.co/google/t5-v1_1-xxl/) text encoder. To reduce memory usage, you can use the `training/prepare_dataset.py` script to precompute latents and embeddings.
Fill or modify the parameters in `prepare_dataset.sh` and execute it to get precomputed latents and embeddings (make sure to specify `--save_tensors` to save the precomputed artifacts). When using these artifacts during training, ensure that you specify the `--load_tensors` flag, or else the videos will be used directly, requiring the text encoder and VAE to be loaded. The script also supports PyTorch DDP so that large datasets can be encoded in parallel across multiple GPUs (modify the `NUM_GPUS` parameter).
Fill or modify the parameters in `prepare_dataset.sh` and execute it to get precomputed latents and embeddings (make sure to specify `--save_latents_and_embeddings` to save the precomputed artifacts). If preparing for image-to-video training, make sure to pass `--save_image_latents`, which encodes and saves image latents along with videos. When using these artifacts during training, ensure that you specify the `--load_tensors` flag, or else the videos will be used directly, requiring the text encoder and VAE to be loaded. The script also supports PyTorch DDP so that large datasets can be encoded in parallel across multiple GPUs (modify the `NUM_GPUS` parameter).
+4 -3
View File
@@ -18,7 +18,7 @@ A black and white animated sequence on a ship’s deck features a bulldog charac
- 任意分辨率且必须能被32整除。例如,`720 * 480`, `1920 * 1020` 等分辨率。
- **支持的帧数(Frames)**:
- 满足 (4K +1),即4的倍数,例如,16, 24, 32, 48, 64, 80。
- 必须是 `4 * k` 或 `4 * k + 1`(例如:16, 32, 49, 81)
所有的视频建议放在一个文件夹中。
@@ -66,6 +66,7 @@ OOM(内存不足),因为它需要加载 [VAE](https://huggingface.co/THUDM
文本编码器。为了降低内存需求,您可以使用 `training/prepare_dataset.py` 脚本预先计算潜在变量和嵌入。
填写或修改 `prepare_dataset.sh` 中的参数并执行它以获得预先计算的潜在变量和嵌入(请确保指定 `--save_tensors`
以保存预计算的工件)。在训练期间使用这些工件时,确保指定 `--load_tensors` 标志,否则将直接使用视频并需要加载文本编码器和
填写或修改 `prepare_dataset.sh` 中的参数并执行它以获得预先计算的潜在变量和嵌入(请确保指定 `--save_latents_and_embeddings`
以保存预计算的工件)。如果准备图像到视频的训练,请确保传递 `--save_image_latents`,它对沙子进行编码,将图像潜在值与视频一起保存。
在训练期间使用这些工件时,确保指定 `--load_tensors` 标志,否则将直接使用视频并需要加载文本编码器和
VAE。该脚本还支持 PyTorch DDP,以便可以使用多个 GPU 并行编码大型数据集(修改 `NUM_GPUS` 参数)。
+1 -1
View File
@@ -38,7 +38,7 @@ CMD_WITHOUT_PRE_ENCODING="\
--dtype $DTYPE
"
CMD_WITH_PRE_ENCODING="$CMD_WITHOUT_PRE_ENCODING --save_tensors"
CMD_WITH_PRE_ENCODING="$CMD_WITHOUT_PRE_ENCODING --save_latents_and_embeddings"
# Select which you'd like to run
CMD=$CMD_WITH_PRE_ENCODING
+5 -3
View File
@@ -26,7 +26,6 @@ from typing import Any, Dict
import diffusers
import torch
import transformers
import wandb
from accelerate import Accelerator, DistributedType
from accelerate.logging import get_logger
from accelerate.utils import (
@@ -53,6 +52,8 @@ from torch.utils.data import DataLoader
from tqdm.auto import tqdm
from transformers import AutoTokenizer, T5EncoderModel
import wandb
from args import get_args # isort:skip
from dataset import BucketSampler, VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop # isort:skip
@@ -523,7 +524,7 @@ def main(args):
# Scheduler and math around the number of training steps.
overrode_max_train_steps = False
num_update_steps_per_epoch = math.ceil(len(train_dataset) / args.gradient_accumulation_steps)
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
if args.max_train_steps is None:
args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
overrode_max_train_steps = True
@@ -560,7 +561,7 @@ def main(args):
)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataset) / args.gradient_accumulation_steps)
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
if overrode_max_train_steps:
args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
# Afterwards we recalculate our number of training epochs
@@ -582,6 +583,7 @@ def main(args):
accelerator.print("***** Running training *****")
accelerator.print(f" Num trainable parameters = {num_trainable_parameters}")
accelerator.print(f" Num examples = {len(train_dataset)}")
accelerator.print(f" Num batches each epoch = {len(train_dataloader)}")
accelerator.print(f" Num epochs = {args.num_train_epochs}")
accelerator.print(f" Instantaneous batch size per device = {args.train_batch_size}")
accelerator.print(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}")
+5 -3
View File
@@ -25,7 +25,6 @@ from typing import Any, Dict
import diffusers
import torch
import transformers
import wandb
from accelerate import Accelerator, DistributedType
from accelerate.logging import get_logger
from accelerate.utils import (
@@ -52,6 +51,8 @@ from torch.utils.data import DataLoader
from tqdm.auto import tqdm
from transformers import AutoTokenizer, T5EncoderModel
import wandb
from args import get_args # isort:skip
from dataset import BucketSampler, VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop # isort:skip
@@ -507,7 +508,7 @@ def main(args):
# Scheduler and math around the number of training steps.
overrode_max_train_steps = False
num_update_steps_per_epoch = math.ceil(len(train_dataset) / args.gradient_accumulation_steps)
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
if args.max_train_steps is None:
args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
overrode_max_train_steps = True
@@ -544,7 +545,7 @@ def main(args):
)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataset) / args.gradient_accumulation_steps)
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
if overrode_max_train_steps:
args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
# Afterwards we recalculate our number of training epochs
@@ -566,6 +567,7 @@ def main(args):
accelerator.print("***** Running training *****")
accelerator.print(f" Num trainable parameters = {num_trainable_parameters}")
accelerator.print(f" Num examples = {len(train_dataset)}")
accelerator.print(f" Num batches each epoch = {len(train_dataloader)}")
accelerator.print(f" Num epochs = {args.num_train_epochs}")
accelerator.print(f" Instantaneous batch size per device = {args.train_batch_size}")
accelerator.print(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}")
+5 -3
View File
@@ -25,7 +25,6 @@ from typing import Any, Dict
import diffusers
import torch
import transformers
import wandb
from accelerate import Accelerator, DistributedType
from accelerate.logging import get_logger
from accelerate.utils import (
@@ -51,6 +50,8 @@ from torch.utils.data import DataLoader
from tqdm.auto import tqdm
from transformers import AutoTokenizer, T5EncoderModel
import wandb
from args import get_args # isort:skip
from dataset import BucketSampler, VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop # isort:skip
@@ -471,7 +472,7 @@ def main(args):
# Scheduler and math around the number of training steps.
overrode_max_train_steps = False
num_update_steps_per_epoch = math.ceil(len(train_dataset) / args.gradient_accumulation_steps)
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
if args.max_train_steps is None:
args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
overrode_max_train_steps = True
@@ -508,7 +509,7 @@ def main(args):
)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataset) / args.gradient_accumulation_steps)
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
if overrode_max_train_steps:
args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
# Afterwards we recalculate our number of training epochs
@@ -530,6 +531,7 @@ def main(args):
accelerator.print("***** Running training *****")
accelerator.print(f" Num trainable parameters = {num_trainable_parameters}")
accelerator.print(f" Num examples = {len(train_dataset)}")
accelerator.print(f" Num batches each epoch = {len(train_dataloader)}")
accelerator.print(f" Num epochs = {args.num_train_epochs}")
accelerator.print(f" Instantaneous batch size per device = {args.train_batch_size}")
accelerator.print(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}")
+11 -1
View File
@@ -375,7 +375,7 @@ class BucketSampler(Sampler):
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:
@@ -386,6 +386,16 @@ class BucketSampler(Sampler):
self.buckets = {resolution: [] for resolution in data_source.resolutions}
self._raised_warning_for_drop_last = False
def __len__(self):
if self.drop_last and not self._raised_warning_for_drop_last:
self._raised_warning_for_drop_last = True
logger.warning(
"Calculating the length for bucket sampler is not possible when `drop_last` is set to True. This may cause problems when setting the number of epochs used for training."
)
return (len(self.data_source) + self.batch_size - 1) // self.batch_size
def __iter__(self):
for index, data in enumerate(self.data_source):
video_metadata = data["video_metadata"]
+96 -59
View File
@@ -41,21 +41,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.")
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.")
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.")
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
@@ -145,23 +151,21 @@ def get_args() -> Dict[str, Any]:
parser.add_argument(
"--max_sequence_length", type=int, default=226, help="Max sequence length of prompt embeddings."
)
parser.add_argument("--target_fps", type=int, default=8, help="Frame rate of output videos.")
parser.add_argument(
"--target_fps", type=int, default=8, help="Frame rate of output videos if `--save_tensors` is unspecified."
)
parser.add_argument(
"--save_tensors",
"--save_latents_and_embeddings",
action="store_true",
help="Whether to encode videos/captions to latents/embeddings and save them in pytorch serializable format.",
)
parser.add_argument(
"--use_slicing",
action="store_true",
help="Whether to enable sliced encoding/decoding in the VAE. Only used if `--save_tensors` is also used.",
help="Whether to enable sliced encoding/decoding in the VAE. Only used if `--save_latents_and_embeddings` is also used.",
)
parser.add_argument(
"--use_tiling",
action="store_true",
help="Whether to enable tiled encoding/decoding in the VAE. Only used if `--save_tensors` is also used.",
help="Whether to enable tiled encoding/decoding in the VAE. Only used if `--save_latents_and_embeddings` is also used.",
)
parser.add_argument("--batch_size", type=int, default=1, help="Number of videos to process at once in the VAE.")
parser.add_argument(
@@ -293,10 +297,15 @@ def save_video(video: torch.Tensor, path: pathlib.Path, fps: int = 8) -> None:
def save_prompt(prompt: str, path: pathlib.Path) -> None:
with open(path, "w", encoding="utf=8") as file:
with open(path, "w", encoding="utf-8") as file:
file.write(prompt)
def save_metadata(metadata: Dict[str, Any], path: pathlib.Path) -> None:
with open(path, "w", encoding="utf-8") as file:
file.write(json.dumps(metadata))
@torch.no_grad()
def serialize_artifacts(
batch_size: int,
@@ -314,10 +323,8 @@ def serialize_artifacts(
prompts: Optional[List[str]] = None,
prompt_embeds: Optional[torch.Tensor] = None,
) -> None:
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
num_frames, height, width = videos.size(1), videos.size(3), videos.size(4)
metadata = [{"num_frames": num_frames, "height": height, "width": width}]
data_folder_mapper_list = [
(images, images_dir, lambda img, path: save_image(img[0], path), "png"),
@@ -326,6 +333,7 @@ def serialize_artifacts(
(video_latents, video_latents_dir, torch.save, "pt"),
(prompts, prompts_dir, save_prompt, "txt"),
(prompt_embeds, prompt_embeds_dir, torch.save, "pt"),
(metadata, videos_dir, save_metadata, "txt"),
]
filenames = [uuid.uuid4() for _ in range(batch_size)]
@@ -334,7 +342,7 @@ def serialize_artifacts(
continue
for slice, filename in zip(data, filenames):
if isinstance(slice, torch.Tensor):
slice = slice.clone()
slice = slice.clone().to("cpu")
path = folder.joinpath(f"{filename}.{extension}")
save_fn(slice, path)
@@ -443,8 +451,10 @@ def main():
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)
images = None
if args.save_image_latents:
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)
@@ -466,20 +476,23 @@ def main():
# 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)
if args.save_latents_and_embeddings:
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)
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()
if args.use_slicing:
vae.enable_slicing()
if args.use_tiling:
vae.enable_tiling()
# 4. Compute latents and embeddings and save
if rank == 0:
@@ -492,45 +505,65 @@ def main():
for step, batch in enumerate(iterator):
try:
images = None
image_latents = None
video_latents = None
prompt_embeds = None
if args.save_image_latents:
images = batch["images"].to(device, non_blocking=True)
images = images.permute(0, 2, 1, 3, 4) # [B, C, F, H, W]
videos = batch["videos"].to(device, non_blocking=True)
videos = videos.permute(0, 2, 1, 3, 4) # [B, C, F, H, W]
prompts = batch["prompts"]
# 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
if args.save_latents_and_embeddings:
if args.save_image_latents:
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)
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,
prompts,
args.max_sequence_length,
device,
weight_dtype,
requires_grad=False,
)
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)
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)
if images is not None:
images = (images.permute(0, 2, 1, 3, 4) + 1) / 2
# Encode prompts
prompt_embeds = compute_prompt_embeddings(
tokenizer,
text_encoder,
prompts,
args.max_sequence_length,
device,
weight_dtype,
requires_grad=False,
)
videos = (videos.permute(0, 2, 1, 3, 4) + 1) / 2
output_queue.put(
{
"batch_size": prompt_embeds.shape[0],
"batch_size": len(prompts),
"fps": target_fps,
"images_dir": images_dir,
"image_latents_dir": image_latents_dir,
@@ -576,6 +609,7 @@ def main():
("video_latents", "pt"),
("prompts", "txt"),
("prompt_embeds", "pt"),
("videos", "txt"),
]:
tmp_subfolder = tmp_dir.joinpath(subfolder)
combined_subfolder = output_dir.joinpath(subfolder)
@@ -616,10 +650,14 @@ def main():
with open(videos_txt, "w") as file:
for stem in stems:
file.write(f"videos/{stem}.pt\n")
file.write(f"videos/{stem}.mp4\n")
with open(data_jsonl, "w") as file:
for prompt, stem in zip(prompts, stems):
video_metadata_txt = output_dir.joinpath(f"videos/{stem}.txt")
with open(video_metadata_txt, "r", encoding="utf-8") as metadata_file:
metadata = json.loads(metadata_file.read())
data = {
"prompt": prompt,
"prompt_embed": f"prompt_embeds/{stem}.pt",
@@ -627,12 +665,11 @@ def main():
"image_latent": f"image_latents/{stem}.pt",
"video": f"videos/{stem}.mp4",
"video_latent": f"video_latents/{stem}.pt",
"metadata": metadata,
}
file.write(json.dumps(data) + "\n")
print(
f"Completed preprocessing. All files saved to `{output_dir.as_posix()}`"
)
print(f"Completed preprocessing. All files saved to `{output_dir.as_posix()}`")
if __name__ == "__main__":