Support CogVideoX T2V (#165)

* support cog t2v.

* generator.

* updates

* style

* fixes

* fix padding frames.

Co-authored-by: zRzRzRzRzRzRzR <Yuxuan.Zhang2104@student.xjtlu.edu.cn>

* 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 <Yuxuan.Zhang2104@student.xjtlu.edu.cn>
Co-authored-by: Aryan <aryan@huggingface.co>
This commit is contained in:
Sayak Paul
2025-01-03 10:53:16 +05:30
committed by GitHub
parent 2a1aa05191
commit b8352abf70
13 changed files with 654 additions and 55 deletions
+1
View File
@@ -170,5 +170,6 @@ wandb/
dump*
outputs*
*.slurm
.vscode/
!requirements.txt
+6 -1
View File
@@ -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",
+1
View File
@@ -0,0 +1 @@
from .cogvideox_lora import COGVIDEOX_T2V_LORA_CONFIG
+327
View File
@@ -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,
}
+51
View File
@@ -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
@@ -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,
+2
View File
@@ -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]
+4
View File
@@ -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,
},
}
+110 -48
View File
@@ -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)
+9 -1
View File
@@ -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
+115
View File
@@ -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
+25
View File
@@ -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])
+1 -1
View File
@@ -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