mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
+2
-2
@@ -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).
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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}")
|
||||
|
||||
@@ -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}")
|
||||
|
||||
@@ -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
@@ -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
@@ -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__":
|
||||
|
||||
Reference in New Issue
Block a user