mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
Multi-GPU parallel encoding support for training videos.
This commit is contained in:
@@ -3,6 +3,9 @@ __pycache__/
|
||||
*.py[cod]
|
||||
*$py.class
|
||||
|
||||
# JetBrains
|
||||
.idea
|
||||
|
||||
# C extensions
|
||||
*.so
|
||||
|
||||
|
||||
@@ -1,4 +1,9 @@
|
||||
# Finetuning CogVideoX
|
||||
# CogVideoX Factory
|
||||
|
||||
## Introduction
|
||||
|
||||
This is a repos for CogVideoX Fine-tuning.
|
||||
|
||||
|
||||
|
||||
## Dataset Preparation
|
||||
|
||||
@@ -0,0 +1,95 @@
|
||||
# CogVideoX Factory
|
||||
|
||||
## 简介
|
||||
|
||||
这是用于 CogVideoX 微调的仓库。
|
||||
|
||||
## 数据集准备
|
||||
|
||||
创建两个文件,一个文件包含以换行符分隔的提示词,另一个文件包含以换行符分隔的视频数据路径(视频文件的路径必须相对于您在指定 `--data_root` 时传递的路径)。让我们通过一个例子来更好地理解这一点!
|
||||
|
||||
假设您将 `--data_root` 指定为 `/dataset`,并且该目录包含文件:`prompts.txt` 和 `videos.txt`。
|
||||
|
||||
`prompts.txt` 文件应包含以换行符分隔的提示词:
|
||||
|
||||
```
|
||||
一段黑白动画序列,主角是一只名为 Rabbity Ribfried 的兔子和一只拟人化的山羊,在一个充满音乐和趣味的环境中,展示他们不断发展的互动。
|
||||
一段黑白动画序列,场景在船甲板上,主角是一只名为 Bully Bulldoger 的斗牛犬角色,展示了夸张的面部表情和肢体语言。角色从自信到专注,再到紧张和痛苦,展示了一系列情绪,随着它克服挑战。船的内部在背景中保持静止,只有简单的细节,如钟声和开着的门。角色的动态动作和变化的表情推动了故事的发展,没有镜头移动,确保观众专注于其不断变化的反应和肢体动作。
|
||||
...
|
||||
```
|
||||
|
||||
`videos.txt` 文件应包含以换行符分隔的视频文件路径。请注意,路径应相对于 `--data_root` 目录。
|
||||
|
||||
```bash
|
||||
videos/00000.mp4
|
||||
videos/00001.mp4
|
||||
...
|
||||
```
|
||||
|
||||
总体而言,如果您在数据集根目录运行 `tree` 命令,您的数据集应如下所示:
|
||||
|
||||
```bash
|
||||
/dataset
|
||||
├── prompts.txt
|
||||
├── videos.txt
|
||||
├── videos
|
||||
├── videos/00000.mp4
|
||||
├── videos/00001.mp4
|
||||
├── ...
|
||||
```
|
||||
|
||||
使用此格式时,`--caption_column` 必须是 `prompts.txt`,`--video_column` 必须是 `videos.txt`。如果您的数据存储在 CSV 文件中,您也可以指定 `--dataset_file` 为 CSV 的路径,`--caption_column` 和 `--video_column` 为 CSV 文件中的实际列名。
|
||||
|
||||
例如,让我们使用这个 [Disney 数据集](https://huggingface.co/datasets/Wild-Heart/Disney-VideoGeneration-Dataset) 进行微调。要下载,可以使用 🤗 Hugging Face CLI。
|
||||
|
||||
```bash
|
||||
huggingface-cli download --repo-type dataset Wild-Heart/Disney-VideoGeneration-Dataset --local-dir video-dataset-disney
|
||||
```
|
||||
|
||||
## 训练
|
||||
|
||||
TODO
|
||||
|
||||
请查看 `training/*.sh`
|
||||
|
||||
注意:未在 MPS 上测试
|
||||
|
||||
## 内存需求
|
||||
|
||||
训练支持并验证的内存优化包括:
|
||||
|
||||
- 来自 [TorchAO](https://github.com/pytorch/ao) 的 `CPUOffloadOptimizer`。
|
||||
- 来自 [bitsandbytes](https://huggingface.co/docs/bitsandbytes/optimizers) 的低位优化器。
|
||||
|
||||
### LoRA 微调
|
||||
|
||||
<details>
|
||||
<summary> AdamW </summary>
|
||||
|
||||
With `train_batch_size = 1`:
|
||||
|
||||
| model | lora rank | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|
||||
|:------------------:|:---------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
|
||||
| THUDM/CogVideoX-2b | 16 | False | 12.945 | 43.764 | 46.918 | 24.234 |
|
||||
| THUDM/CogVideoX-2b | 16 | True | 12.945 | 12.945 | 21.121 | 24.234 |
|
||||
| THUDM/CogVideoX-2b | 64 | False | 13.035 | 44.314 | 47.469 | 24.469 |
|
||||
| THUDM/CogVideoX-2b | 64 | True | 13.036 | 13.035 | 21.564 | 24.500 |
|
||||
| THUDM/CogVideoX-2b | 256 | False | 13.095 | 45.826 | 48.990 | 25.543 |
|
||||
| THUDM/CogVideoX-2b | 256 | True | 13.094 | 13.095 | 22.344 | 25.537 |
|
||||
| THUDM/CogVideoX-5b | 16 | True | 19.742 | 19.742 | 28.746 | 38.123 |
|
||||
| THUDM/CogVideoX-5b | 64 | True | 20.006 | 20.818 | 30.338 | 38.738 |
|
||||
| THUDM/CogVideoX-5b | 256 | True | 20.771 | 22.119 | 31.939 | 41.537 |
|
||||
|
||||
With `train_batch_size = 4`:
|
||||
|
||||
| model | lora rank | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|
||||
|:------------------:|:---------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
|
||||
| THUDM/CogVideoX-2b | 16 | True | 12.945 | 21.803 | 21.814 | 24.322 |
|
||||
| THUDM/CogVideoX-2b | 64 | True | 13.035 | 22.254 | 22.254 | 24.572 |
|
||||
| THUDM/CogVideoX-2b | 256 | True | 13.094 | 22.020 | 22.033 | 25.574 |
|
||||
| THUDM/CogVideoX-5b | 16 | True | 19.742 | 46.492 | 46.492 | 38.197 |
|
||||
| THUDM/CogVideoX-5b | 64 | True | 20.006 | 47.805 | 47.805 | 39.365 |
|
||||
| THUDM/CogVideoX-5b | 256 | True | 20.771 | 47.268 | 47.332 | 41.008 |
|
||||
|
||||
> [!NOTE]
|
||||
>
|
||||
+20
-26
@@ -1,12 +1,11 @@
|
||||
#!/bin/bash
|
||||
|
||||
MODEL_ID="THUDM/CogVideoX-2b"
|
||||
MODEL_ID="/share/official_pretrains/hf_home/CogVideoX-5b"
|
||||
|
||||
# For more details on the expected data format, please refer to the README.
|
||||
DATA_ROOT="/raid/aryan/video-dataset-tom-and-jerry" # This needs to be the path to the base directory where your videos are located.
|
||||
DATA_ROOT="/share/home/zyx/disney_cogvideox"
|
||||
CAPTION_COLUMN="prompts.txt"
|
||||
VIDEO_COLUMN="videos.txt"
|
||||
OUTPUT_DIR="/raid/aryan/video-dataset-tom-and-jerry-encoded"
|
||||
OUTPUT_DIR="/share/home/zyx/disney_cogvideox-encoded-multi"
|
||||
HEIGHT=480
|
||||
WIDTH=720
|
||||
MAX_NUM_FRAMES=49
|
||||
@@ -14,29 +13,24 @@ MAX_SEQUENCE_LENGTH=226
|
||||
TARGET_FPS=8
|
||||
BATCH_SIZE=1
|
||||
DTYPE=fp32
|
||||
NUM_GPUS=8
|
||||
|
||||
# To create a folder-style dataset structure without pre-encoding videos and captions'
|
||||
CMD_WITHOUT_PRE_ENCODING="\
|
||||
python3 training/prepare_dataset.py \
|
||||
--model_id $MODEL_ID \
|
||||
--data_root $DATA_ROOT \
|
||||
--caption_column $CAPTION_COLUMN \
|
||||
--video_column $VIDEO_COLUMN \
|
||||
--output_dir $OUTPUT_DIR \
|
||||
--height $HEIGHT \
|
||||
--width $WIDTH \
|
||||
--max_num_frames $MAX_NUM_FRAMES \
|
||||
--max_sequence_length $MAX_SEQUENCE_LENGTH \
|
||||
--target_fps $TARGET_FPS \
|
||||
--batch_size $BATCH_SIZE \
|
||||
--dtype $DTYPE
|
||||
"
|
||||
|
||||
CMD_WITH_PRE_ENCODING="$CMD_WITHOUT_PRE_ENCODING --save_tensors"
|
||||
|
||||
# Select which you'd like to run
|
||||
CMD=$CMD_WITH_PRE_ENCODING
|
||||
CMD="torchrun --nproc_per_node=$NUM_GPUS \
|
||||
training/prepare_dataset.py \
|
||||
--model_id $MODEL_ID \
|
||||
--data_root $DATA_ROOT \
|
||||
--caption_column $CAPTION_COLUMN \
|
||||
--video_column $VIDEO_COLUMN \
|
||||
--output_dir $OUTPUT_DIR \
|
||||
--height $HEIGHT \
|
||||
--width $WIDTH \
|
||||
--max_num_frames $MAX_NUM_FRAMES \
|
||||
--max_sequence_length $MAX_SEQUENCE_LENGTH \
|
||||
--target_fps $TARGET_FPS \
|
||||
--batch_size $BATCH_SIZE \
|
||||
--dtype $DTYPE \
|
||||
--save_tensors"
|
||||
|
||||
echo "===== Running \`$CMD\` ====="
|
||||
eval $CMD
|
||||
echo -ne "===== Finished running script =====\n"
|
||||
echo -ne "===== Finished running script =====\n"
|
||||
@@ -4,7 +4,7 @@ export WANDB_MODE="offline"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
|
||||
GPU_IDS="0"
|
||||
GPU_IDS="0,1,2,3,4,5,6,7"
|
||||
|
||||
# Training Configurations
|
||||
# Experiment with as many hyperparameters as you want!
|
||||
@@ -19,22 +19,26 @@ ACCELERATE_CONFIG_FILE="accelerate_configs/uncompiled_1.yaml"
|
||||
# Absolute path to where the data is located. Make sure to have read the README for how to prepare data.
|
||||
# This example assumes you downloaded an already prepared dataset from HF CLI as follows:
|
||||
# huggingface-cli download --repo-type dataset Wild-Heart/Disney-VideoGeneration-Dataset --local-dir /path/to/my/datasets/disney-dataset
|
||||
DATA_ROOT="/path/to/my/datasets/disney-dataset"
|
||||
CAPTION_COLUMN="prompt.txt"
|
||||
|
||||
DATA_ROOT="/share/home/zyx/disney_cogvideox-encoded-multi"
|
||||
CAPTION_COLUMN="prompts.txt"
|
||||
VIDEO_COLUMN="videos.txt"
|
||||
MODEL_PATH="/share/official_pretrains/hf_home/CogVideoX-5b"
|
||||
|
||||
|
||||
# Launch experiments with different hyperparameters
|
||||
for learning_rate in "${LEARNING_RATES[@]}"; do
|
||||
for lr_schedule in "${LR_SCHEDULES[@]}"; do
|
||||
for optimizer in "${OPTIMIZERS[@]}"; do
|
||||
for steps in "${MAX_TRAIN_STEPS[@]}"; do
|
||||
output_dir="/path/to/my/models/cogvideox-lora__optimizer_${optimizer}__steps_${steps}__lr-schedule_${lr_schedule}__learning-rate_${learning_rate}/"
|
||||
output_dir="cogvideox-lora__optimizer_${optimizer}__steps_${steps}__lr-schedule_${lr_schedule}__learning-rate_${learning_rate}/"
|
||||
|
||||
cmd="accelerate launch --config_file $ACCELERATE_CONFIG_FILE --gpu_ids $GPU_IDS training/cogvideox_text_to_video_lora.py \
|
||||
--pretrained_model_name_or_path THUDM/CogVideoX-5b \
|
||||
--pretrained_model_name_or_path $MODEL_PATH \
|
||||
--data_root $DATA_ROOT \
|
||||
--caption_column $CAPTION_COLUMN \
|
||||
--video_column $VIDEO_COLUMN \
|
||||
--load_tensors \
|
||||
--id_token BW_STYLE \
|
||||
--height_buckets 480 \
|
||||
--width_buckets 720 \
|
||||
|
||||
+105
-46
@@ -1,23 +1,21 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
# For folder structure dataset: python3 prepare_dataset.py --model_id THUDM/CogVideoX-2b --data_root /raid/aryan/video-dataset-disney/ --caption_column prompts.txt --video_column videos.txt --output_dir dump --height 480 --width 720 --max_num_frames 49 --max_sequence_length 226 --target_fps 8 --batch_size 1 --dtype fp32
|
||||
# For latent/embed structure dataset: python3 prepare_dataset.py --model_id THUDM/CogVideoX-2b --data_root /raid/aryan/video-dataset-disney/ --caption_column prompts.txt --video_column videos.txt --output_dir dump --height 480 --width 720 --max_num_frames 49 --max_sequence_length 226 --target_fps 8 --batch_size 1 --dtype fp32 --save_tensors
|
||||
|
||||
import argparse
|
||||
import gc
|
||||
import os
|
||||
import pathlib
|
||||
import traceback
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import pandas as pd
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from diffusers import AutoencoderKLCogVideoX
|
||||
from diffusers.utils import export_to_video, get_logger
|
||||
from torchvision import transforms
|
||||
from transformers import T5EncoderModel, T5Tokenizer
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
# Must import after importing torch, otherwise there's a nasty segfault when loading text_encoder/vae
|
||||
import decord # isort:skip
|
||||
|
||||
decord.bridge.set_bridge("torch")
|
||||
@@ -104,7 +102,7 @@ def get_args() -> Dict[str, Any]:
|
||||
|
||||
def load_dataset_from_local_path(
|
||||
data_root: pathlib.Path, caption_column: str, video_column: str
|
||||
) -> Tuple[List[str], List[str]]:
|
||||
) -> Tuple[List[str], List[pathlib.Path]]:
|
||||
if not data_root.exists():
|
||||
raise ValueError("Root folder for videos does not exist")
|
||||
|
||||
@@ -127,7 +125,7 @@ def load_dataset_from_local_path(
|
||||
|
||||
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 atleast one path that is not a valid file."
|
||||
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
|
||||
@@ -135,7 +133,7 @@ def load_dataset_from_local_path(
|
||||
|
||||
def load_dataset_from_csv(
|
||||
data_root: pathlib.Path, dataset_file: pathlib.Path, caption_column: str, video_column: str
|
||||
) -> Tuple[List[str], List[str]]:
|
||||
) -> Tuple[List[str], List[pathlib.Path]]:
|
||||
df = pd.read_csv(dataset_file)
|
||||
prompts = df[caption_column].tolist()
|
||||
video_paths = df[video_column].tolist()
|
||||
@@ -143,7 +141,7 @@ def load_dataset_from_csv(
|
||||
|
||||
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 atleast one path that is not a valid file."
|
||||
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
|
||||
@@ -151,7 +149,7 @@ def load_dataset_from_csv(
|
||||
|
||||
def load_and_preprocess_video(
|
||||
path: pathlib.Path, height: int, width: int, max_num_frames: int, video_transforms, num_threads: int = 0
|
||||
) -> torch.Tensor:
|
||||
) -> Optional[torch.Tensor]:
|
||||
frames = None
|
||||
|
||||
try:
|
||||
@@ -160,11 +158,11 @@ def load_and_preprocess_video(
|
||||
|
||||
if video_num_frames < max_num_frames:
|
||||
logger.warning(
|
||||
f"Video at `{path.as_posix()}` should have atleast `{max_num_frames=}`, but got only `{video_num_frames=}`. Skipping it."
|
||||
f"Video at `{path.as_posix()}` should have at least `{max_num_frames=}`, but got only `{video_num_frames=}`. Skipping it."
|
||||
)
|
||||
return
|
||||
return None
|
||||
|
||||
indices = list(range(0, video_num_frames, video_num_frames // max_num_frames))
|
||||
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()
|
||||
@@ -241,7 +239,7 @@ def encode_prompt(
|
||||
def compute_prompt_embeddings(
|
||||
tokenizer: T5Tokenizer,
|
||||
text_encoder: T5EncoderModel,
|
||||
prompt: str,
|
||||
prompts: List[str],
|
||||
max_sequence_length: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
@@ -251,7 +249,7 @@ def compute_prompt_embeddings(
|
||||
prompt_embeds = encode_prompt(
|
||||
tokenizer,
|
||||
text_encoder,
|
||||
prompt,
|
||||
prompts,
|
||||
num_videos_per_prompt=1,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
@@ -262,7 +260,7 @@ def compute_prompt_embeddings(
|
||||
prompt_embeds = encode_prompt(
|
||||
tokenizer,
|
||||
text_encoder,
|
||||
prompt,
|
||||
prompts,
|
||||
num_videos_per_prompt=1,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
@@ -272,7 +270,7 @@ def compute_prompt_embeddings(
|
||||
|
||||
|
||||
def save_videos(
|
||||
videos: torch.Tensor, video_paths: List[str], prompts: List[str], output_dir: pathlib.Path, target_fps: int = 8
|
||||
videos: torch.Tensor, video_paths: List[pathlib.Path], prompts: List[str], output_dir: pathlib.Path, target_fps: int = 8
|
||||
) -> None:
|
||||
assert videos.size(0) == len(video_paths)
|
||||
|
||||
@@ -289,13 +287,13 @@ def save_videos(
|
||||
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(pathlib.Path(video_path).name)
|
||||
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(), "w", encoding="utf-8") as file:
|
||||
for video_path in video_paths:
|
||||
file.write(f"videos/{pathlib.Path(video_path).name}\n")
|
||||
file.write(f"videos/{video_path.name}\n")
|
||||
|
||||
with open(output_dir.joinpath("prompts.txt").as_posix(), "w", encoding="utf-8") as file:
|
||||
for prompt in prompts:
|
||||
@@ -305,7 +303,7 @@ def save_videos(
|
||||
def save_latents_and_embeddings(
|
||||
latents: torch.Tensor,
|
||||
prompt_embeds: torch.Tensor,
|
||||
video_paths: List[str],
|
||||
video_paths: List[pathlib.Path],
|
||||
prompts: List[str],
|
||||
output_dir: pathlib.Path,
|
||||
) -> None:
|
||||
@@ -321,27 +319,20 @@ def save_latents_and_embeddings(
|
||||
embeds_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
for latent, embed, video_path in zip(latents, prompt_embeds, video_paths):
|
||||
# Need to perform the clone, otherwise the entire `latents` or `prompt_embeds` tensor is
|
||||
# saved for every single video/prompt embedding. This is due to us viewing a slice of a
|
||||
# large tensor when iteratively saving stuff here.
|
||||
latent = latent.clone()
|
||||
embed = embed.clone()
|
||||
|
||||
video_path = pathlib.Path(video_path)
|
||||
filename_without_ext = video_path.name.split(".")[0]
|
||||
filename_without_ext = video_path.stem
|
||||
|
||||
latent_filename = latents_dir.joinpath(filename_without_ext)
|
||||
embed_filename = embeds_dir.joinpath(filename_without_ext)
|
||||
|
||||
latent_filename = f"{latent_filename}.pt"
|
||||
embed_filename = f"{embed_filename}.pt"
|
||||
latent_filename = latents_dir.joinpath(f"{filename_without_ext}.pt")
|
||||
embed_filename = embeds_dir.joinpath(f"{filename_without_ext}.pt")
|
||||
|
||||
torch.save(latent, latent_filename)
|
||||
torch.save(embed, embed_filename)
|
||||
|
||||
with open(output_dir.joinpath("videos.txt").as_posix(), "w", encoding="utf-8") as file:
|
||||
for video_path in video_paths:
|
||||
file.write(f"videos/{pathlib.Path(video_path).name}\n")
|
||||
file.write(f"videos/{video_path.name}\n")
|
||||
|
||||
with open(output_dir.joinpath("prompts.txt").as_posix(), "w", encoding="utf-8") as file:
|
||||
for prompt in prompts:
|
||||
@@ -349,7 +340,23 @@ def save_latents_and_embeddings(
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def main(args: Dict[str, Any]) -> None:
|
||||
def main():
|
||||
args = get_args()
|
||||
|
||||
# Initialize distributed processing
|
||||
if 'LOCAL_RANK' in os.environ:
|
||||
local_rank = int(os.environ['LOCAL_RANK'])
|
||||
torch.cuda.set_device(local_rank)
|
||||
dist.init_process_group(backend="nccl")
|
||||
world_size = dist.get_world_size()
|
||||
rank = dist.get_rank()
|
||||
else:
|
||||
# Single GPU
|
||||
local_rank = 0
|
||||
world_size = 1
|
||||
rank = 0
|
||||
torch.cuda.set_device(local_rank)
|
||||
|
||||
data_root = pathlib.Path(args.data_root)
|
||||
dataset_file = None
|
||||
if args.dataset_file:
|
||||
@@ -367,10 +374,18 @@ def main(args: Dict[str, Any]) -> None:
|
||||
]
|
||||
)
|
||||
|
||||
# Preprocess videos with progress bar
|
||||
prompts_usable = []
|
||||
video_paths_usable = []
|
||||
videos = []
|
||||
for prompt, path in zip(prompts, video_paths):
|
||||
|
||||
# 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
|
||||
)
|
||||
@@ -378,18 +393,47 @@ def main(args: Dict[str, Any]) -> 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)
|
||||
|
||||
# Split data among GPUs
|
||||
if world_size > 1:
|
||||
total_samples = len(prompts_usable)
|
||||
samples_per_gpu = total_samples // 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
|
||||
|
||||
# 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]
|
||||
else:
|
||||
pass
|
||||
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
|
||||
if not args.save_tensors:
|
||||
save_videos(videos, video_paths_usable, prompts_usable, pathlib.Path(args.output_dir), args.target_fps)
|
||||
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("cuda")
|
||||
text_encoder = text_encoder.to(device)
|
||||
|
||||
prompt_embeds_list = []
|
||||
for start_index in range(0, len(prompts_usable), args.batch_size):
|
||||
|
||||
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)
|
||||
|
||||
for start_index in iterator:
|
||||
end_index = min(len(prompts_usable), start_index + args.batch_size)
|
||||
batch_prompts = prompts_usable[start_index:end_index]
|
||||
|
||||
@@ -398,20 +442,20 @@ def main(args: Dict[str, Any]) -> None:
|
||||
text_encoder,
|
||||
batch_prompts,
|
||||
max_sequence_length=args.max_sequence_length,
|
||||
device="cuda",
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
prompt_embeds_list.append(prompt_embeds)
|
||||
prompt_embeds_list.append(prompt_embeds.to("cpu"))
|
||||
|
||||
prompt_embeds = torch.cat(prompt_embeds_list).to("cpu")
|
||||
prompt_embeds = torch.cat(prompt_embeds_list)
|
||||
|
||||
del tokenizer, text_encoder
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.synchronize("cuda")
|
||||
torch.cuda.synchronize(device)
|
||||
|
||||
vae = AutoencoderKLCogVideoX.from_pretrained(args.model_id, subfolder="vae", torch_dtype=dtype)
|
||||
vae = vae.to("cuda")
|
||||
vae = vae.to(device)
|
||||
|
||||
if args.use_slicing:
|
||||
vae.enable_slicing()
|
||||
@@ -419,11 +463,17 @@ def main(args: Dict[str, Any]) -> None:
|
||||
vae.enable_tiling()
|
||||
|
||||
encoded_videos = []
|
||||
for start_index in range(0, len(video_paths_usable), args.batch_size):
|
||||
|
||||
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("cuda")
|
||||
batch_videos = batch_videos.to(device)
|
||||
batch_videos = batch_videos.permute(0, 2, 1, 3, 4) # [B, C, F, H, W]
|
||||
|
||||
if args.use_slicing:
|
||||
@@ -432,19 +482,28 @@ def main(args: Dict[str, Any]) -> None:
|
||||
else:
|
||||
encoded_video = vae._encode(batch_videos)
|
||||
|
||||
encoded_videos.append(encoded_video)
|
||||
encoded_videos.append(encoded_video.to("cpu"))
|
||||
|
||||
encoded_videos = torch.cat(encoded_videos).to("cpu")
|
||||
encoded_videos = torch.cat(encoded_videos)
|
||||
|
||||
del vae
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.synchronize("cuda")
|
||||
torch.cuda.synchronize(device)
|
||||
|
||||
# Ensure that only one process creates the output directories
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
save_latents_and_embeddings(
|
||||
encoded_videos, prompt_embeds, video_paths_usable, prompts_usable, pathlib.Path(args.output_dir)
|
||||
)
|
||||
|
||||
# Finalize distributed processing
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = get_args()
|
||||
@@ -455,4 +514,4 @@ if __name__ == "__main__":
|
||||
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(args)
|
||||
main()
|
||||
Reference in New Issue
Block a user