From b8352abf70157b44018d7fc4473760967034db46 Mon Sep 17 00:00:00 2001 From: Sayak Paul Date: Fri, 3 Jan 2025 10:53:16 +0530 Subject: [PATCH] Support CogVideoX T2V (#165) * support cog t2v. * generator. * updates * style * fixes * fix padding frames. Co-authored-by: zRzRzRzRzRzRzR * revert changes related to generator. * refactor a lot of things. * accept revision and cache_dir. * remove unused var * refactoring fixes. * refactor * update * update --------- Co-authored-by: zRzRzRzRzRzRzR Co-authored-by: Aryan --- .gitignore | 1 + finetrainers/args.py | 7 +- finetrainers/cogvideox/__init__.py | 1 + finetrainers/cogvideox/cogvideox_lora.py | 327 ++++++++++++++++++ finetrainers/cogvideox/utils.py | 51 +++ .../hunyuan_video/hunyuan_video_lora.py | 6 +- finetrainers/ltx_video/ltx_video_lora.py | 2 + finetrainers/models.py | 4 + finetrainers/trainer.py | 158 ++++++--- finetrainers/utils/__init__.py | 10 +- finetrainers/utils/diffusion_utils.py | 115 ++++++ finetrainers/utils/model_utils.py | 25 ++ finetrainers/utils/torch_utils.py | 2 +- 13 files changed, 654 insertions(+), 55 deletions(-) create mode 100644 finetrainers/cogvideox/__init__.py create mode 100644 finetrainers/cogvideox/cogvideox_lora.py create mode 100644 finetrainers/cogvideox/utils.py create mode 100644 finetrainers/utils/model_utils.py diff --git a/.gitignore b/.gitignore index ecadaee..60f4682 100644 --- a/.gitignore +++ b/.gitignore @@ -170,5 +170,6 @@ wandb/ dump* outputs* *.slurm +.vscode/ !requirements.txt diff --git a/finetrainers/args.py b/finetrainers/args.py index 247d11a..2608ee3 100644 --- a/finetrainers/args.py +++ b/finetrainers/args.py @@ -4,6 +4,7 @@ from typing import Any, Dict, List, Optional, Tuple import torch from .constants import DEFAULT_IMAGE_RESOLUTION_BUCKETS, DEFAULT_VIDEO_RESOLUTION_BUCKETS +from .models import SUPPORTED_MODEL_CONFIGS class Args: @@ -226,7 +227,11 @@ def validate_args(args: Args): def _add_model_arguments(parser: argparse.ArgumentParser) -> None: parser.add_argument( - "--model_name", type=str, required=True, choices=["hunyuan_video", "ltx_video"], help="Name of model to train." + "--model_name", + type=str, + required=True, + choices=list(SUPPORTED_MODEL_CONFIGS.keys()), + help="Name of model to train.", ) parser.add_argument( "--pretrained_model_name_or_path", diff --git a/finetrainers/cogvideox/__init__.py b/finetrainers/cogvideox/__init__.py new file mode 100644 index 0000000..6a3f826 --- /dev/null +++ b/finetrainers/cogvideox/__init__.py @@ -0,0 +1 @@ +from .cogvideox_lora import COGVIDEOX_T2V_LORA_CONFIG diff --git a/finetrainers/cogvideox/cogvideox_lora.py b/finetrainers/cogvideox/cogvideox_lora.py new file mode 100644 index 0000000..c3b754a --- /dev/null +++ b/finetrainers/cogvideox/cogvideox_lora.py @@ -0,0 +1,327 @@ +from typing import Any, Dict, List, Optional, Union + +import torch +from diffusers import AutoencoderKLCogVideoX, CogVideoXDDIMScheduler, CogVideoXPipeline, CogVideoXTransformer3DModel +from PIL import Image +from transformers import T5EncoderModel, T5Tokenizer + +from .utils import prepare_rotary_positional_embeddings + + +def load_condition_models( + model_id: str = "THUDM/CogVideoX-5b", + text_encoder_dtype: torch.dtype = torch.bfloat16, + revision: Optional[str] = None, + cache_dir: Optional[str] = None, + **kwargs, +): + tokenizer = T5Tokenizer.from_pretrained(model_id, subfolder="tokenizer", revision=revision, cache_dir=cache_dir) + text_encoder = T5EncoderModel.from_pretrained( + model_id, subfolder="text_encoder", torch_dtype=text_encoder_dtype, revision=revision, cache_dir=cache_dir + ) + return {"tokenizer": tokenizer, "text_encoder": text_encoder} + + +def load_latent_models( + model_id: str = "THUDM/CogVideoX-5b", + vae_dtype: torch.dtype = torch.bfloat16, + revision: Optional[str] = None, + cache_dir: Optional[str] = None, + **kwargs, +): + vae = AutoencoderKLCogVideoX.from_pretrained( + model_id, subfolder="vae", torch_dtype=vae_dtype, revision=revision, cache_dir=cache_dir + ) + return {"vae": vae} + + +def load_diffusion_models( + model_id: str = "THUDM/CogVideoX-5b", + transformer_dtype: torch.dtype = torch.bfloat16, + revision: Optional[str] = None, + cache_dir: Optional[str] = None, + **kwargs, +): + transformer = CogVideoXTransformer3DModel.from_pretrained( + model_id, subfolder="transformer", torch_dtype=transformer_dtype, revision=revision, cache_dir=cache_dir + ) + scheduler = CogVideoXDDIMScheduler.from_pretrained(model_id, subfolder="scheduler") + return {"transformer": transformer, "scheduler": scheduler} + + +def initialize_pipeline( + model_id: str = "THUDM/CogVideoX-5b", + text_encoder_dtype: torch.dtype = torch.bfloat16, + transformer_dtype: torch.dtype = torch.bfloat16, + vae_dtype: torch.dtype = torch.bfloat16, + tokenizer: Optional[T5Tokenizer] = None, + text_encoder: Optional[T5EncoderModel] = None, + transformer: Optional[CogVideoXTransformer3DModel] = None, + vae: Optional[AutoencoderKLCogVideoX] = None, + scheduler: Optional[CogVideoXDDIMScheduler] = None, + device: Optional[torch.device] = None, + revision: Optional[str] = None, + cache_dir: Optional[str] = None, + enable_slicing: bool = False, + enable_tiling: bool = False, + enable_model_cpu_offload: bool = False, + **kwargs, +) -> CogVideoXPipeline: + component_name_pairs = [ + ("tokenizer", tokenizer), + ("text_encoder", text_encoder), + ("transformer", transformer), + ("vae", vae), + ("scheduler", scheduler), + ] + components = {} + for name, component in component_name_pairs: + if component is not None: + components[name] = component + + pipe = CogVideoXPipeline.from_pretrained(model_id, **components, revision=revision, cache_dir=cache_dir) + pipe.text_encoder = pipe.text_encoder.to(dtype=text_encoder_dtype) + pipe.transformer = pipe.transformer.to(dtype=transformer_dtype) + pipe.vae = pipe.vae.to(dtype=vae_dtype) + + if enable_slicing: + pipe.vae.enable_slicing() + if enable_tiling: + pipe.vae.enable_tiling() + + if enable_model_cpu_offload: + pipe.enable_model_cpu_offload(device=device) + else: + pipe.to(device=device) + + return pipe + + +def prepare_conditions( + tokenizer, + text_encoder, + prompt: Union[str, List[str]], + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, + max_sequence_length: int = 226, # TODO: this should be configurable + **kwargs, +): + device = device or text_encoder.device + dtype = dtype or text_encoder.dtype + return _get_t5_prompt_embeds( + tokenizer=tokenizer, + text_encoder=text_encoder, + prompt=prompt, + max_sequence_length=max_sequence_length, + device=device, + dtype=dtype, + ) + + +def prepare_latents( + vae: AutoencoderKLCogVideoX, + image_or_video: torch.Tensor, + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, + generator: Optional[torch.Generator] = None, + precompute: bool = False, + **kwargs, +) -> torch.Tensor: + device = device or vae.device + dtype = dtype or vae.dtype + + if image_or_video.ndim == 4: + image_or_video = image_or_video.unsqueeze(2) + assert image_or_video.ndim == 5, f"Expected 5D tensor, got {image_or_video.ndim}D tensor" + + image_or_video = image_or_video.to(device=device, dtype=vae.dtype) + image_or_video = image_or_video.permute(0, 2, 1, 3, 4) # [B, C, F, H, W] + if not precompute: + latents = vae.encode(image_or_video).latent_dist.sample(generator=generator) + if not vae.config.invert_scale_latents: + latents = latents * vae.config.scaling_factor + # For training Cog 1.5, we don't need to handle the scaling factor here. + # The CogVideoX team forgot to multiply here, so we should not do it too. Invert scale latents + # is probably only needed for image-to-video training. + # TODO(aryan): investigate this + # else: + # latents = 1 / vae.config.scaling_factor * latents + latents = latents.to(dtype=dtype) + return {"latents": latents} + else: + # handle vae scaling in the `train()` method directly. + if vae.use_slicing and image_or_video.shape[0] > 1: + encoded_slices = [vae._encode(x_slice) for x_slice in image_or_video.split(1)] + h = torch.cat(encoded_slices) + else: + h = vae._encode(image_or_video) + return {"latents": h} + + +def post_latent_preparation( + vae_config: Dict[str, Any], latents: torch.Tensor, patch_size_t: Optional[int] = None, **kwargs +) -> torch.Tensor: + if not vae_config.invert_scale_latents: + latents = latents * vae_config.scaling_factor + # For training Cog 1.5, we don't need to handle the scaling factor here. + # The CogVideoX team forgot to multiply here, so we should not do it too. Invert scale latents + # is probably only needed for image-to-video training. + # TODO(aryan): investigate this + # else: + # latents = 1 / vae_config.scaling_factor * latents + latents = _pad_frames(latents, patch_size_t) + latents = latents.permute(0, 2, 1, 3, 4) # [B, F, C, H, W] + return {"latents": latents} + + +def collate_fn_t2v(batch: List[List[Dict[str, torch.Tensor]]]) -> Dict[str, torch.Tensor]: + return { + "prompts": [x["prompt"] for x in batch[0]], + "videos": torch.stack([x["video"] for x in batch[0]]), + } + + +def calculate_noisy_latents( + scheduler: CogVideoXDDIMScheduler, + noise: torch.Tensor, + latents: torch.Tensor, + timesteps: torch.LongTensor, +) -> torch.Tensor: + noisy_latents = scheduler.add_noise(latents, noise, timesteps) + return noisy_latents + + +def forward_pass( + transformer: CogVideoXTransformer3DModel, + scheduler: CogVideoXDDIMScheduler, + prompt_embeds: torch.Tensor, + latents: torch.Tensor, + noisy_latents: torch.Tensor, + timesteps: torch.LongTensor, + ofs_emb: Optional[torch.Tensor] = None, + **kwargs, +) -> torch.Tensor: + # Just hardcode for now. In Diffusers, we will refactor such that RoPE would be handled within the model itself. + VAE_SPATIAL_SCALE_FACTOR = 8 + transformer_config = transformer.module.config if hasattr(transformer, "module") else transformer.config + batch_size, num_frames, num_channels, height, width = noisy_latents.shape + rope_base_height = transformer_config.sample_height * VAE_SPATIAL_SCALE_FACTOR + rope_base_width = transformer_config.sample_width * VAE_SPATIAL_SCALE_FACTOR + + image_rotary_emb = ( + prepare_rotary_positional_embeddings( + height=height * VAE_SPATIAL_SCALE_FACTOR, + width=width * VAE_SPATIAL_SCALE_FACTOR, + num_frames=num_frames, + vae_scale_factor_spatial=VAE_SPATIAL_SCALE_FACTOR, + patch_size=transformer_config.patch_size, + patch_size_t=transformer_config.patch_size_t if hasattr(transformer_config, "patch_size_t") else None, + attention_head_dim=transformer_config.attention_head_dim, + device=transformer.device, + base_height=rope_base_height, + base_width=rope_base_width, + ) + if transformer_config.use_rotary_positional_embeddings + else None + ) + ofs_emb = None if transformer_config.ofs_embed_dim is None else latents.new_full((batch_size,), fill_value=2.0) + + velocity = transformer( + hidden_states=noisy_latents, + timestep=timesteps, + encoder_hidden_states=prompt_embeds, + ofs=ofs_emb, + image_rotary_emb=image_rotary_emb, + return_dict=False, + )[0] + # For CogVideoX, the transformer predicts the velocity. The denoised output is calculated by applying the same + # code paths as scheduler.get_velocity(), which can be confusing to understand. + denoised_latents = scheduler.get_velocity(velocity, noisy_latents, timesteps) + + return {"latents": denoised_latents} + + +def validation( + pipeline: CogVideoXPipeline, + prompt: str, + image: Optional[Image.Image] = None, + video: Optional[List[Image.Image]] = None, + height: Optional[int] = None, + width: Optional[int] = None, + num_frames: Optional[int] = None, + num_videos_per_prompt: int = 1, + generator: Optional[torch.Generator] = None, + **kwargs, +): + generation_kwargs = { + "prompt": prompt, + "height": height, + "width": width, + "num_frames": num_frames, + "num_videos_per_prompt": num_videos_per_prompt, + "generator": generator, + "return_dict": True, + "output_type": "pil", + } + generation_kwargs = {k: v for k, v in generation_kwargs.items() if v is not None} + output = pipeline(**generation_kwargs).frames[0] + return [("video", output)] + + +def _get_t5_prompt_embeds( + tokenizer: T5Tokenizer, + text_encoder: T5EncoderModel, + prompt: Union[str, List[str]] = None, + max_sequence_length: int = 226, + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, +): + prompt = [prompt] if isinstance(prompt, str) else prompt + + text_inputs = tokenizer( + prompt, + padding="max_length", + max_length=max_sequence_length, + truncation=True, + add_special_tokens=True, + return_tensors="pt", + ) + text_input_ids = text_inputs.input_ids + + prompt_embeds = text_encoder(text_input_ids.to(device))[0] + prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) + + return {"prompt_embeds": prompt_embeds} + + +def _pad_frames(latents: torch.Tensor, patch_size_t: int): + if patch_size_t is None or patch_size_t == 1: + return latents + + # `latents` should be of the following format: [B, C, F, H, W]. + # For CogVideoX 1.5, the latent frames should be padded to make it divisible by patch_size_t + latent_num_frames = latents.shape[2] + additional_frames = patch_size_t - latent_num_frames % patch_size_t + + if additional_frames > 0: + last_frame = latents[:, :, -1:, :, :] + padding_frames = last_frame.repeat(1, 1, additional_frames, 1, 1) + latents = torch.cat([latents, padding_frames], dim=2) + + return latents + + +COGVIDEOX_T2V_LORA_CONFIG = { + "pipeline_cls": CogVideoXPipeline, + "load_condition_models": load_condition_models, + "load_latent_models": load_latent_models, + "load_diffusion_models": load_diffusion_models, + "initialize_pipeline": initialize_pipeline, + "prepare_conditions": prepare_conditions, + "prepare_latents": prepare_latents, + "post_latent_preparation": post_latent_preparation, + "collate_fn": collate_fn_t2v, + "calculate_noisy_latents": calculate_noisy_latents, + "forward_pass": forward_pass, + "validation": validation, +} diff --git a/finetrainers/cogvideox/utils.py b/finetrainers/cogvideox/utils.py new file mode 100644 index 0000000..bd98c1f --- /dev/null +++ b/finetrainers/cogvideox/utils.py @@ -0,0 +1,51 @@ +from typing import Optional, Tuple + +import torch +from diffusers.models.embeddings import get_3d_rotary_pos_embed +from diffusers.pipelines.cogvideo.pipeline_cogvideox import get_resize_crop_region_for_grid + + +def prepare_rotary_positional_embeddings( + height: int, + width: int, + num_frames: int, + vae_scale_factor_spatial: int = 8, + patch_size: int = 2, + patch_size_t: int = None, + attention_head_dim: int = 64, + device: Optional[torch.device] = None, + base_height: int = 480, + base_width: int = 720, +) -> Tuple[torch.Tensor, torch.Tensor]: + grid_height = height // (vae_scale_factor_spatial * patch_size) + grid_width = width // (vae_scale_factor_spatial * patch_size) + base_size_width = base_width // (vae_scale_factor_spatial * patch_size) + base_size_height = base_height // (vae_scale_factor_spatial * patch_size) + + if patch_size_t is None: + # CogVideoX 1.0 + grid_crops_coords = get_resize_crop_region_for_grid( + (grid_height, grid_width), base_size_width, base_size_height + ) + freqs_cos, freqs_sin = get_3d_rotary_pos_embed( + embed_dim=attention_head_dim, + crops_coords=grid_crops_coords, + grid_size=(grid_height, grid_width), + temporal_size=num_frames, + ) + else: + # CogVideoX 1.5 + base_num_frames = (num_frames + patch_size_t - 1) // patch_size_t + + freqs_cos, freqs_sin = get_3d_rotary_pos_embed( + embed_dim=attention_head_dim, + crops_coords=None, + grid_size=(grid_height, grid_width), + temporal_size=base_num_frames, + grid_type="slice", + max_size=(base_size_height, base_size_width), + ) + + freqs_cos = freqs_cos.to(device=device) + freqs_sin = freqs_sin.to(device=device) + return freqs_cos, freqs_sin diff --git a/finetrainers/hunyuan_video/hunyuan_video_lora.py b/finetrainers/hunyuan_video/hunyuan_video_lora.py index 719425f..dc0ecf7 100644 --- a/finetrainers/hunyuan_video/hunyuan_video_lora.py +++ b/finetrainers/hunyuan_video/hunyuan_video_lora.py @@ -198,10 +198,7 @@ def prepare_latents( return {"latents": h} -def post_latent_preparation( - latents: torch.Tensor, - **kwargs, -) -> torch.Tensor: +def post_latent_preparation(latents: torch.Tensor, **kwargs) -> torch.Tensor: return {"latents": latents} @@ -221,6 +218,7 @@ def forward_pass( latents: torch.Tensor, noisy_latents: torch.Tensor, timesteps: torch.LongTensor, + **kwargs, ) -> torch.Tensor: denoised_latents = transformer( hidden_states=noisy_latents, diff --git a/finetrainers/ltx_video/ltx_video_lora.py b/finetrainers/ltx_video/ltx_video_lora.py index e104a8c..0e1af9b 100644 --- a/finetrainers/ltx_video/ltx_video_lora.py +++ b/finetrainers/ltx_video/ltx_video_lora.py @@ -173,6 +173,7 @@ def post_latent_preparation( width: int, patch_size: int = 1, patch_size_t: int = 1, + **kwargs, ) -> torch.Tensor: latents = _normalize_latents(latents, latents_mean, latents_std) latents = _pack_latents(latents, patch_size, patch_size_t) @@ -196,6 +197,7 @@ def forward_pass( num_frames: int, height: int, width: int, + **kwargs, ) -> torch.Tensor: # TODO(aryan): make configurable rope_interpolation_scale = [1 / 25, 32, 32] diff --git a/finetrainers/models.py b/finetrainers/models.py index f2310b8..c7d95ae 100644 --- a/finetrainers/models.py +++ b/finetrainers/models.py @@ -1,5 +1,6 @@ from typing import Any, Dict +from .cogvideox import COGVIDEOX_T2V_LORA_CONFIG from .hunyuan_video import HUNYUAN_VIDEO_T2V_LORA_CONFIG from .ltx_video import LTX_VIDEO_T2V_LORA_CONFIG @@ -11,6 +12,9 @@ SUPPORTED_MODEL_CONFIGS = { "ltx_video": { "lora": LTX_VIDEO_T2V_LORA_CONFIG, }, + "cogvideox": { + "lora": COGVIDEOX_T2V_LORA_CONFIG, + }, } diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index 22ae323..eb1a6de 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -3,7 +3,7 @@ import logging import math import os import random -from datetime import timedelta +from datetime import datetime, timedelta from pathlib import Path from typing import Any, Dict @@ -11,7 +11,6 @@ import diffusers import torch import torch.backends import transformers -import wandb from accelerate import Accelerator, DistributedType from accelerate.logging import get_logger from accelerate.utils import ( @@ -21,18 +20,17 @@ from accelerate.utils import ( gather_object, set_seed, ) +from diffusers.configuration_utils import FrozenDict from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution from diffusers.optimization import get_scheduler -from diffusers.training_utils import ( - cast_training_params, - compute_density_for_timestep_sampling, - compute_loss_weighting_for_sd3, -) +from diffusers.training_utils import cast_training_params from diffusers.utils import export_to_video, load_image, load_video from huggingface_hub import create_repo, upload_folder from peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict from tqdm import tqdm +import wandb + from .args import _INVERSE_DTYPE_MAP, Args, validate_args from .constants import ( FINETRAINERS_LOG_LEVEL, @@ -45,10 +43,18 @@ from .models import get_config_from_model_name from .state import State from .utils.checkpointing import get_intermediate_ckpt_path, get_latest_ckpt_path_to_resume_from from .utils.data_utils import should_perform_precomputation +from .utils.diffusion_utils import ( + get_scheduler_alphas, + get_scheduler_sigmas, + prepare_loss_weights, + prepare_sigmas, + prepare_target, +) from .utils.file_utils import string_to_filename from .utils.memory_utils import free_memory, get_memory_statistics, make_contiguous +from .utils.model_utils import resolve_vae_cls_from_ckpt_path from .utils.optimizer_utils import get_optimizer -from .utils.torch_utils import align_device_and_dtype, expand_tensor_to_dims, unwrap_model +from .utils.torch_utils import align_device_and_dtype, expand_tensor_dims, unwrap_model logger = get_logger("finetrainers") @@ -60,6 +66,7 @@ class Trainer: validate_args(args) self.args = args + self.args.seed = self.args.seed or datetime.now().year self.state = State() # Tokenizers @@ -82,6 +89,9 @@ class Trainer: # Scheduler self.scheduler = None + self.transformer_config = None + self.vae_config = None + self._init_distributed() self._init_logging() self._init_directories_and_repositories() @@ -125,6 +135,7 @@ class Trainer: return load_component_kwargs def _set_components(self, components: Dict[str, Any]) -> None: + # Set models self.tokenizer = components.get("tokenizer", self.tokenizer) self.tokenizer_2 = components.get("tokenizer_2", self.tokenizer_2) self.tokenizer_3 = components.get("tokenizer_3", self.tokenizer_3) @@ -136,6 +147,10 @@ class Trainer: self.vae = components.get("vae", self.vae) self.scheduler = components.get("scheduler", self.scheduler) + # Set configs + self.transformer_config = self.transformer.config if self.transformer is not None else self.transformer_config + self.vae_config = self.vae.config if self.vae is not None else self.vae_config + def _delete_components(self) -> None: self.tokenizer = None self.tokenizer_2 = None @@ -172,8 +187,6 @@ class Trainer: if self.args.enable_tiling: self.vae.enable_tiling() - self.transformer_config = self.transformer.config if self.transformer is not None else None - def prepare_precomputations(self) -> None: if not self.args.precompute_conditions: return @@ -242,19 +255,21 @@ class Trainer: conditions_dir.mkdir(parents=True, exist_ok=True) latents_dir.mkdir(parents=True, exist_ok=True) + accelerator = self.state.accelerator + # Precompute conditions progress_bar = tqdm( - range(0, len(self.dataset)), + range(0, (len(self.dataset) + accelerator.num_processes - 1) // accelerator.num_processes), desc="Precomputing conditions", - disable=not self.state.accelerator.is_local_main_process, + disable=not accelerator.is_local_main_process, ) index = 0 for i, data in enumerate(self.dataset): - if i % self.state.accelerator.num_processes != self.state.accelerator.process_index: + if i % accelerator.num_processes != accelerator.process_index: continue logger.debug( - f"Precomputing conditions and latents for batch {i + 1}/{len(self.dataset)} on process {self.state.accelerator.process_index}" + f"Precomputing conditions for batch {i + 1}/{len(self.dataset)} on process {accelerator.process_index}" ) text_conditions = self.model_config["prepare_conditions"]( @@ -265,10 +280,10 @@ class Trainer: text_encoder_2=self.text_encoder_2, text_encoder_3=self.text_encoder_3, prompt=data["prompt"], - device=self.state.accelerator.device, + device=accelerator.device, dtype=self.state.weight_dtype, ) - filename = conditions_dir / f"conditions-{i}-{index}.pt" + filename = conditions_dir / f"conditions-{accelerator.process_index}-{index}.pt" torch.save(text_conditions, filename.as_posix()) index += 1 progress_bar.update(1) @@ -276,7 +291,7 @@ class Trainer: memory_statistics = get_memory_statistics() logger.info(f"Memory after precomputing conditions: {json.dumps(memory_statistics, indent=4)}") - torch.cuda.reset_peak_memory_stats(self.state.accelerator.device) + torch.cuda.reset_peak_memory_stats(accelerator.device) # Precompute latents latent_components = self.model_config["load_latent_models"](**self._get_load_components_kwargs()) @@ -296,39 +311,39 @@ class Trainer: self.vae.enable_tiling() progress_bar = tqdm( - range(0, len(self.dataset)), + range(0, (len(self.dataset) + accelerator.num_processes - 1) // accelerator.num_processes), desc="Precomputing latents", - disable=not self.state.accelerator.is_local_main_process, + disable=not accelerator.is_local_main_process, ) index = 0 for i, data in enumerate(self.dataset): - if i % self.state.accelerator.num_processes != self.state.accelerator.process_index: + if i % accelerator.num_processes != accelerator.process_index: continue logger.debug( - f"Precomputing latents for batch {i + 1}/{len(self.dataset)} on process {self.state.accelerator.process_index}" + f"Precomputing latents for batch {i + 1}/{len(self.dataset)} on process {accelerator.process_index}" ) latent_conditions = self.model_config["prepare_latents"]( vae=self.vae, image_or_video=data["video"].unsqueeze(0), - device=self.state.accelerator.device, + device=accelerator.device, dtype=self.state.weight_dtype, generator=self.state.generator, precompute=True, ) - filename = latents_dir / f"latents-{self.state.accelerator.process_index}-{index}.pt" + filename = latents_dir / f"latents-{accelerator.process_index}-{index}.pt" torch.save(latent_conditions, filename.as_posix()) index += 1 progress_bar.update(1) self._delete_components() - self.state.accelerator.wait_for_everyone() + accelerator.wait_for_everyone() logger.info("Precomputation complete") memory_statistics = get_memory_statistics() logger.info(f"Memory after precomputing latents: {json.dumps(memory_statistics, indent=4)}") - torch.cuda.reset_peak_memory_stats(self.state.accelerator.device) + torch.cuda.reset_peak_memory_stats(accelerator.device) # Update dataloader to use precomputed conditions and latents self.dataloader = torch.utils.data.DataLoader( @@ -569,6 +584,20 @@ class Trainer: memory_statistics = get_memory_statistics() logger.info(f"Memory before training start: {json.dumps(memory_statistics, indent=4)}") + if self.vae_config is None: + # If we've precomputed conditions and latents already, and are now re-using it, we will never load + # the VAE so self.vae_config will not be set. So, we need to load it here. + vae_cls_name = resolve_vae_cls_from_ckpt_path( + self.args.pretrained_model_name_or_path, revision=self.args.revision, cache_dir=self.args.cache_dir + ) + vae_config = vae_cls_name.load_config( + self.args.pretrained_model_name_or_path, + subfolder="vae", + revision=self.args.revision, + cache_dir=self.args.cache_dir, + ) + self.vae_config = FrozenDict(**vae_config) + self.state.train_batch_size = ( self.args.batch_size * self.state.accelerator.num_processes * self.args.gradient_accumulation_steps ) @@ -611,12 +640,14 @@ class Trainer: accelerator = self.state.accelerator weight_dtype = self.state.weight_dtype - scheduler_sigmas = self.scheduler.sigmas.clone().to(device=accelerator.device, dtype=weight_dtype) generator = torch.Generator(device=accelerator.device) if self.args.seed is not None: generator = generator.manual_seed(self.args.seed) self.state.generator = generator + scheduler_sigmas = get_scheduler_sigmas(self.scheduler).to(device=accelerator.device, dtype=torch.float32) + scheduler_alphas = get_scheduler_alphas(self.scheduler).to(device=accelerator.device, dtype=torch.float32) + for epoch in range(first_epoch, self.state.train_epochs): logger.debug(f"Starting epoch ({epoch + 1}/{self.state.train_epochs})") @@ -644,7 +675,7 @@ class Trainer: patch_size_t=self.transformer_config.patch_size_t, device=accelerator.device, dtype=weight_dtype, - generator=generator, + generator=self.state.generator, ) text_conditions = self.model_config["prepare_conditions"]( tokenizer=self.tokenizer, @@ -660,9 +691,16 @@ class Trainer: text_conditions = batch["text_conditions"] latent_conditions["latents"] = DiagonalGaussianDistribution( latent_conditions["latents"] - ).sample(generator) - if "post_latent_preparation" in self.model_config.keys(): - latent_conditions = self.model_config["post_latent_preparation"](**latent_conditions) + ).sample(self.state.generator) + + # This method should only be called for precomputed latents. + # TODO(aryan): rename this in separate PR + latent_conditions = self.model_config["post_latent_preparation"]( + vae_config=self.vae_config, + patch_size=self.transformer_config.patch_size, + patch_size_t=self.transformer_config.patch_size_t, + **latent_conditions, + ) align_device_and_dtype(latent_conditions, accelerator.device, weight_dtype) align_device_and_dtype(text_conditions, accelerator.device, weight_dtype) batch_size = latent_conditions["latents"].shape[0] @@ -679,40 +717,64 @@ class Trainer: if "pooled_prompt_embeds" in text_conditions: text_conditions["pooled_prompt_embeds"].fill_(0) - # These weighting schemes use a uniform timestep sampling and instead post-weight the loss - weights = compute_density_for_timestep_sampling( - weighting_scheme=self.args.flow_weighting_scheme, + sigmas = prepare_sigmas( + scheduler=self.scheduler, + sigmas=scheduler_sigmas, batch_size=batch_size, - logit_mean=self.args.flow_logit_mean, - logit_std=self.args.flow_logit_std, - mode_scale=self.args.flow_mode_scale, + num_train_timesteps=self.scheduler.config.num_train_timesteps, + flow_weighting_scheme=self.args.flow_weighting_scheme, + flow_logit_mean=self.args.flow_logit_mean, + flow_logit_std=self.args.flow_logit_std, + flow_mode_scale=self.args.flow_mode_scale, + device=accelerator.device, + generator=self.state.generator, ) - indices = (weights * self.scheduler.config.num_train_timesteps).long() - sigmas = scheduler_sigmas[indices] timesteps = (sigmas * 1000.0).long() noise = torch.randn( latent_conditions["latents"].shape, - generator=generator, + generator=self.state.generator, device=accelerator.device, dtype=weight_dtype, ) - sigmas = expand_tensor_to_dims(sigmas, ndim=latent_conditions["latents"].ndim) - noisy_latents = (1.0 - sigmas) * latent_conditions["latents"] + sigmas * noise + sigmas = expand_tensor_dims(sigmas, ndim=noise.ndim) + + # TODO(aryan): We probably don't need calculate_noisy_latents because we can determine the type of + # scheduler and calculate the noisy latents accordingly. Look into this later. + if "calculate_noisy_latents" in self.model_config.keys(): + noisy_latents = self.model_config["calculate_noisy_latents"]( + scheduler=self.scheduler, + noise=noise, + latents=latent_conditions["latents"], + timesteps=timesteps, + ) + else: + # Default to flow-matching noise addition + noisy_latents = (1.0 - sigmas) * latent_conditions["latents"] + sigmas * noise latent_conditions.update({"noisy_latents": noisy_latents}) - # These weighting schemes use a uniform timestep sampling and instead post-weight the loss - weights = compute_loss_weighting_for_sd3( - weighting_scheme=self.args.flow_weighting_scheme, sigmas=sigmas + weights = prepare_loss_weights( + scheduler=self.scheduler, + alphas=scheduler_alphas[timesteps] if scheduler_alphas is not None else None, + sigmas=sigmas, + flow_weighting_scheme=self.args.flow_weighting_scheme, ) + weights = expand_tensor_dims(weights, noise.ndim) + pred = self.model_config["forward_pass"]( - transformer=self.transformer, timesteps=timesteps, **latent_conditions, **text_conditions + transformer=self.transformer, + scheduler=self.scheduler, + timesteps=timesteps, + **latent_conditions, + **text_conditions, + ) + target = prepare_target( + scheduler=self.scheduler, noise=noise, latents=latent_conditions["latents"] ) - target = noise - latent_conditions["latents"] loss = weights.float() * (pred["latents"].float() - target.float()).pow(2) - # Average loss across channel dimension + # Average loss across all but batch dimension loss = loss.mean(list(range(1, loss.ndim))) # Average loss across batch dimension loss = loss.mean() @@ -949,7 +1011,7 @@ class Trainer: self.transformer.train() def evaluate(self) -> None: - raise NotImplementedError + raise NotImplementedError("Evaluation has not been implemented yet.") def _init_distributed(self) -> None: logging_dir = Path(self.args.output_dir, self.args.logging_dir) diff --git a/finetrainers/utils/__init__.py b/finetrainers/utils/__init__.py index 85ffd2f..9f0f45c 100644 --- a/finetrainers/utils/__init__.py +++ b/finetrainers/utils/__init__.py @@ -1,4 +1,12 @@ -from .diffusion_utils import default_flow_shift, resolution_dependant_timestep_flow_shift +from .diffusion_utils import ( + default_flow_shift, + get_scheduler_alphas, + get_scheduler_sigmas, + prepare_loss_weights, + prepare_sigmas, + prepare_target, + resolution_dependant_timestep_flow_shift, +) from .file_utils import delete_files, find_files from .memory_utils import bytes_to_gigabytes, free_memory, get_memory_statistics, make_contiguous from .optimizer_utils import get_optimizer, gradient_norm, max_gradient diff --git a/finetrainers/utils/diffusion_utils.py b/finetrainers/utils/diffusion_utils.py index be37bc9..1d46fbc 100644 --- a/finetrainers/utils/diffusion_utils.py +++ b/finetrainers/utils/diffusion_utils.py @@ -1,4 +1,9 @@ +import math +from typing import Optional, Union + import torch +from diffusers import CogVideoXDDIMScheduler, FlowMatchEulerDiscreteScheduler +from diffusers.training_utils import compute_loss_weighting_for_sd3 # Default values copied from https://github.com/huggingface/diffusers/blob/8957324363d8b239d82db4909fbf8c0875683e3d/src/diffusers/schedulers/scheduling_flow_match_euler_discrete.py#L47 @@ -28,3 +33,113 @@ def resolution_dependant_timestep_flow_shift( def default_flow_shift(sigmas: torch.Tensor, shift: float = 1.0) -> torch.Tensor: sigmas = (sigmas * shift) / (1 + (shift - 1) * sigmas) return sigmas + + +def compute_density_for_timestep_sampling( + weighting_scheme: str, + batch_size: int, + logit_mean: float = None, + logit_std: float = None, + mode_scale: float = None, + device: torch.device = torch.device("cpu"), + generator: Optional[torch.Generator] = None, +) -> torch.Tensor: + r""" + Compute the density for sampling the timesteps when doing SD3 training. + + Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528. + + SD3 paper reference: https://arxiv.org/abs/2403.03206v1. + """ + if weighting_scheme == "logit_normal": + # See 3.1 in the SD3 paper ($rf/lognorm(0.00,1.00)$). + u = torch.normal(mean=logit_mean, std=logit_std, size=(batch_size,), device=device, generator=generator) + u = torch.nn.functional.sigmoid(u) + elif weighting_scheme == "mode": + u = torch.rand(size=(batch_size,), device=device, generator=generator) + u = 1 - u - mode_scale * (torch.cos(math.pi * u / 2) ** 2 - 1 + u) + else: + u = torch.rand(size=(batch_size,), device=device, generator=generator) + return u + + +def get_scheduler_alphas(scheduler: Union[CogVideoXDDIMScheduler, FlowMatchEulerDiscreteScheduler]) -> torch.Tensor: + if isinstance(scheduler, FlowMatchEulerDiscreteScheduler): + return None + elif isinstance(scheduler, CogVideoXDDIMScheduler): + return scheduler.alphas_cumprod.clone() + else: + raise ValueError(f"Unsupported scheduler type {type(scheduler)}") + + +def get_scheduler_sigmas(scheduler: Union[CogVideoXDDIMScheduler, FlowMatchEulerDiscreteScheduler]) -> torch.Tensor: + if isinstance(scheduler, FlowMatchEulerDiscreteScheduler): + return scheduler.sigmas.clone() + elif isinstance(scheduler, CogVideoXDDIMScheduler): + return scheduler.timesteps.clone().float() / float(scheduler.config.num_train_timesteps) + else: + raise ValueError(f"Unsupported scheduler type {type(scheduler)}") + + +def prepare_sigmas( + scheduler: Union[CogVideoXDDIMScheduler, FlowMatchEulerDiscreteScheduler], + sigmas: torch.Tensor, + batch_size: int, + num_train_timesteps: int, + flow_weighting_scheme: str = "none", + flow_logit_mean: float = 0.0, + flow_logit_std: float = 1.0, + flow_mode_scale: float = 1.29, + device: torch.device = torch.device("cpu"), + generator: Optional[torch.Generator] = None, +) -> torch.Tensor: + if isinstance(scheduler, FlowMatchEulerDiscreteScheduler): + weights = compute_density_for_timestep_sampling( + weighting_scheme=flow_weighting_scheme, + batch_size=batch_size, + logit_mean=flow_logit_mean, + logit_std=flow_logit_std, + mode_scale=flow_mode_scale, + device=device, + generator=generator, + ) + indices = (weights * num_train_timesteps).long() + elif isinstance(scheduler, CogVideoXDDIMScheduler): + # TODO(aryan): Currently, only uniform sampling is supported. Add more sampling schemes. + weights = torch.rand(size=(batch_size,), device=device, generator=generator) + indices = (weights * num_train_timesteps).long() + else: + raise ValueError(f"Unsupported scheduler type {type(scheduler)}") + + return sigmas[indices] + + +def prepare_loss_weights( + scheduler: Union[CogVideoXDDIMScheduler, FlowMatchEulerDiscreteScheduler], + alphas: Optional[torch.Tensor] = None, + sigmas: Optional[torch.Tensor] = None, + flow_weighting_scheme: str = "none", +) -> torch.Tensor: + if isinstance(scheduler, FlowMatchEulerDiscreteScheduler): + return compute_loss_weighting_for_sd3(sigmas, weighting_scheme=flow_weighting_scheme) + elif isinstance(scheduler, CogVideoXDDIMScheduler): + # SNR is computed as (alphas / (1 - alphas)), but for some reason CogVideoX uses 1 / (1 - alphas). + # TODO(aryan): Experiment if using alphas / (1 - alphas) gives better results. + return 1 / (1 - alphas) + else: + raise ValueError(f"Unsupported scheduler type {type(scheduler)}") + + +def prepare_target( + scheduler: Union[CogVideoXDDIMScheduler, FlowMatchEulerDiscreteScheduler], + noise: torch.Tensor, + latents: torch.Tensor, +) -> torch.Tensor: + if isinstance(scheduler, FlowMatchEulerDiscreteScheduler): + target = noise - latents + elif isinstance(scheduler, CogVideoXDDIMScheduler): + target = latents + else: + raise ValueError(f"Unsupported scheduler type {type(scheduler)}") + + return target diff --git a/finetrainers/utils/model_utils.py b/finetrainers/utils/model_utils.py new file mode 100644 index 0000000..1451ebf --- /dev/null +++ b/finetrainers/utils/model_utils.py @@ -0,0 +1,25 @@ +import importlib +import json +import os + +from huggingface_hub import hf_hub_download + + +def resolve_vae_cls_from_ckpt_path(ckpt_path, **kwargs): + ckpt_path = str(ckpt_path) + if os.path.exists(str(ckpt_path)) and os.path.isdir(ckpt_path): + index_path = os.path.join(ckpt_path, "model_index.json") + else: + revision = kwargs.get("revision", None) + cache_dir = kwargs.get("cache_dir", None) + index_path = hf_hub_download( + repo_id=ckpt_path, filename="model_index.json", revision=revision, cache_dir=cache_dir + ) + + with open(index_path, "r") as f: + model_index_dict = json.load(f) + assert "vae" in model_index_dict, "No VAE found in the modelx index dict." + + vae_cls_config = model_index_dict["vae"] + library = importlib.import_module(vae_cls_config[0]) + return getattr(library, vae_cls_config[1]) diff --git a/finetrainers/utils/torch_utils.py b/finetrainers/utils/torch_utils.py index 13989ad..1c6ef5d 100644 --- a/finetrainers/utils/torch_utils.py +++ b/finetrainers/utils/torch_utils.py @@ -29,7 +29,7 @@ def align_device_and_dtype( return x -def expand_tensor_to_dims(tensor, ndim): +def expand_tensor_dims(tensor, ndim): while len(tensor.shape) < ndim: tensor = tensor.unsqueeze(-1) return tensor