Hunyuan Video LoRA (#126)

* add hunyuan-video lora support

* minor fixes; make style

* update readme

* update

* update

* Update README.md

* Update README.md

Co-authored-by: Sayak Paul <spsayakpaul@gmail.com>

* update

* update

* change move train scripts to internal directory

* update

---------

Co-authored-by: Sayak Paul <spsayakpaul@gmail.com>
This commit is contained in:
Aryan
2024-12-20 08:02:08 +05:30
committed by GitHub
parent 9ef58e2f3a
commit 223add1a59
13 changed files with 521 additions and 31 deletions
+142 -9
View File
@@ -10,6 +10,11 @@ FineTrainers is a work-in-progress library to support training of video models.
</tr>
</table>
## News
- 🔥 **2024-12-20**: Support for LoRA finetuning of [Hunyuan Video](https://huggingface.co/tencent/HunyuanVideo) added! We would like to thank @SHYuanBest for his work on a training script [here](https://github.com/huggingface/diffusers/pull/10254).
- 🔥 **2024-12-18**: Support for LoRA finetuning of [LTX Video](https://huggingface.co/Lightricks/LTX-Video) added!
## Quickstart
Clone the repository and make sure the requirements are installed: `pip install -r requirements.txt` and install diffusers from source by `pip install git+https://github.com/huggingface/diffusers`.
@@ -40,13 +45,12 @@ export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export FINETRAINERS_LOG_LEVEL=DEBUG
# Modify this based on the number of GPUs available
GPU_IDS="0,1"
DATA_ROOT="/path/to/dataset/cakify"
DATA_ROOT="/raid/aryan/video-dataset-disney"
CAPTION_COLUMN="prompts.txt"
VIDEO_COLUMN="videos.txt"
OUTPUT_DIR="/path/to/output/directory/ltx-video/ltxv_cakify"
OUTPUT_DIR="/path/to/output/directory/ltx-video/ltxv_disney"
# Model arguments
model_cmd="--model_name ltx_video \
@@ -57,7 +61,7 @@ dataset_cmd="--data_root $DATA_ROOT \
--video_column $VIDEO_COLUMN \
--caption_column $CAPTION_COLUMN \
--id_token BW_STYLE \
--video_resolution_buckets 17x512x768 49x512x768 61x512x768 129x512x768 \
--video_resolution_buckets 49x512x768 \
--caption_dropout_p 0.05"
# Dataloader arguments
@@ -71,7 +75,7 @@ training_cmd="--training_type lora \
--seed 42 \
--mixed_precision bf16 \
--batch_size 1 \
--train_steps 2000 \
--train_steps 1200 \
--rank 128 \
--lora_alpha 128 \
--target_modules to_q to_k to_v to_out.0 \
@@ -84,8 +88,8 @@ training_cmd="--training_type lora \
# Optimizer arguments
optimizer_cmd="--optimizer adamw \
--lr 1e-5 \
--lr_scheduler constant \
--lr 3e-5 \
--lr_scheduler constant_with_warmup \
--lr_warmup_steps 100 \
--lr_num_cycles 1 \
--beta1 0.9 \
@@ -95,7 +99,7 @@ optimizer_cmd="--optimizer adamw \
--max_grad_norm 1.0"
# Validation arguments
validation_cmd="--validation_prompts \"BW_STYLE A black and white animated scene unfolds with an anthropomorphic goat surrounded by musical notes and symbols, suggesting a playful environment. Mickey Mouse appears, leaning forward in curiosity as the goat remains still. The goat then engages with Mickey, who bends down to converse or react. The dynamics shift as Mickey grabs the goat, potentially in surprise or playfulness, amidst a minimalistic background. The scene captures the evolving relationship between the two characters in a whimsical, animated setting, emphasizing their interactions and emotions@@@49x512x768:::BW_STYLE A black and white animated scene unfolds with an anthropomorphic goat surrounded by musical notes and symbols, suggesting a playful environment. Mickey Mouse appears, leaning forward in curiosity as the goat remains still. The goat then engages with Mickey, who bends down to converse or react. The dynamics shift as Mickey grabs the goat, potentially in surprise or playfulness, amidst a minimalistic background. The scene captures the evolving relationship between the two characters in a whimsical, animated setting, emphasizing their interactions and emotions@@@129x512x768:::BW_STYLE A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance@@@49x512x768\" \
validation_cmd="--validation_prompts \"afkx A black and white animated scene unfolds with an anthropomorphic goat surrounded by musical notes and symbols, suggesting a playful environment. Mickey Mouse appears, leaning forward in curiosity as the goat remains still. The goat then engages with Mickey, who bends down to converse or react. The dynamics shift as Mickey grabs the goat, potentially in surprise or playfulness, amidst a minimalistic background. The scene captures the evolving relationship between the two characters in a whimsical, animated setting, emphasizing their interactions and emotions.@@@49x512x768:::A woman with long brown hair and light skin smiles at another woman with long blonde hair. The woman with brown hair wears a black jacket and has a small, barely noticeable mole on her right cheek. The camera angle is a close-up, focused on the woman with brown hair's face. The lighting is warm and natural, likely from the setting sun, casting a soft glow on the scene. The scene appears to be real-life footage@@@49x512x768\" \
--num_validation_videos 1 \
--validation_steps 100"
@@ -133,7 +137,7 @@ pipe = LTXPipeline.from_pretrained(
"Lightricks/LTX-Video", torch_dtype=torch.bfloat16
).to("cuda")
+ pipe.load_lora_weights("my-awesome-name/my-awesome-lora", adapter_name="ltxv-lora")
+ pipe.set_adapters(["ltxv-lora"], [1.0])
+ pipe.set_adapters(["ltxv-lora"], [0.75])
video = pipe("<my-awesome-prompt>").frames[0]
export_to_video(video, "output.mp4", fps=8)
@@ -141,6 +145,135 @@ export_to_video(video, "output.mp4", fps=8)
</details>
<details>
<summary> Hunyuan Video </summary>
### Training:
```bash
#!/bin/bash
# export TORCH_LOGS="+dynamo,recompiles,graph_breaks"
# export TORCHDYNAMO_VERBOSE=1
export WANDB_MODE="offline"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export FINETRAINERS_LOG_LEVEL=DEBUG
GPU_IDS="0,1,2,3,4,5,6,7"
DATA_ROOT="/path/to/dataset"
CAPTION_COLUMN="prompts.txt"
VIDEO_COLUMN="videos.txt"
OUTPUT_DIR="/path/to/models/hunyuan-video/hunyuan-video-loras/hunyuan-video_cakify_500_3e-5_constant_with_warmup"
# Model arguments
model_cmd="--model_name hunyuan_video \
--pretrained_model_name_or_path tencent/HunyuanVideo
--revision refs/pr/18"
# Dataset arguments
dataset_cmd="--data_root $DATA_ROOT \
--video_column $VIDEO_COLUMN \
--caption_column $CAPTION_COLUMN \
--id_token afkx \
--video_resolution_buckets 17x512x768 49x512x768 61x512x768 129x512x768 \
--caption_dropout_p 0.05"
# Dataloader arguments
dataloader_cmd="--dataloader_num_workers 0"
# Diffusion arguments
diffusion_cmd=""
# Training arguments
training_cmd="--training_type lora \
--seed 42 \
--mixed_precision bf16 \
--batch_size 1 \
--train_steps 500 \
--rank 128 \
--lora_alpha 128 \
--target_modules to_q to_k to_v to_out.0 \
--gradient_accumulation_steps 1 \
--gradient_checkpointing \
--checkpointing_steps 500 \
--checkpointing_limit 2 \
--enable_slicing \
--enable_tiling"
# Optimizer arguments
optimizer_cmd="--optimizer adamw \
--lr 2e-5 \
--lr_scheduler constant_with_warmup \
--lr_warmup_steps 100 \
--lr_num_cycles 1 \
--beta1 0.9 \
--beta2 0.95 \
--weight_decay 1e-4 \
--epsilon 1e-8 \
--max_grad_norm 1.0"
# Validation arguments
validation_cmd="--validation_prompts \"afkx A baker carefully cuts a green bell pepper cake on a white plate against a bright yellow background, followed by a strawberry cake with a similar slice of cake being cut before the interior of the bell pepper cake is revealed with the surrounding cake-to-object sequence.@@@49x512x768:::afkx A cake shaped like a Nutella container is carefully sliced, revealing a light interior, amidst a Nutella-themed setup, showcasing deliberate cutting and preserved details for an appetizing dessert presentation on a white base with accompanying jello and cutlery, highlighting culinary skills and creative cake designs.@@@49x512x768:::afkx A cake shaped like a Nutella container is carefully sliced, revealing a light interior, amidst a Nutella-themed setup, showcasing deliberate cutting and preserved details for an appetizing dessert presentation on a white base with accompanying jello and cutlery, highlighting culinary skills and creative cake designs.@@@61x512x768:::afkx A vibrant orange cake disguised as a Nike packaging box sits on a dark surface, meticulous in its detail and design, complete with a white swoosh and 'NIKE' logo. A person's hands, holding a knife, hover over the cake, ready to make a precise cut, amidst a simple and clean background.@@@61x512x768:::afkx A vibrant orange cake disguised as a Nike packaging box sits on a dark surface, meticulous in its detail and design, complete with a white swoosh and 'NIKE' logo. A person's hands, holding a knife, hover over the cake, ready to make a precise cut, amidst a simple and clean background.@@@97x512x768:::afkx A vibrant orange cake disguised as a Nike packaging box sits on a dark surface, meticulous in its detail and design, complete with a white swoosh and 'NIKE' logo. A person's hands, holding a knife, hover over the cake, ready to make a precise cut, amidst a simple and clean background.@@@129x512x768:::A person with gloved hands carefully cuts a cake shaped like a Skittles bottle, beginning with a precise incision at the lid, followed by careful sequential cuts around the neck, eventually detaching the lid from the body, revealing the chocolate interior of the cake while showcasing the layered design's detail.@@@61x512x768:::afkx A woman with long brown hair and light skin smiles at another woman with long blonde hair. The woman with brown hair wears a black jacket and has a small, barely noticeable mole on her right cheek. The camera angle is a close-up, focused on the woman with brown hair's face. The lighting is warm and natural, likely from the setting sun, casting a soft glow on the scene. The scene appears to be real-life footage@@@61x512x768\" \
--num_validation_videos 1 \
--validation_steps 100"
# Miscellaneous arguments
miscellaneous_cmd="--tracker_name finetrainers-hunyuan-video \
--output_dir $OUTPUT_DIR \
--nccl_timeout 1800 \
--report_to wandb"
cmd="accelerate launch --config_file accelerate_configs/uncompiled_8.yaml --gpu_ids $GPU_IDS train.py \
$model_cmd \
$dataset_cmd \
$dataloader_cmd \
$diffusion_cmd \
$training_cmd \
$optimizer_cmd \
$validation_cmd \
$miscellaneous_cmd"
echo "Running command: $cmd"
eval $cmd
echo -ne "-------------------- Finished executing script --------------------\n\n"
```
### Inference:
Assuming your LoRA is saved and pushed to the HF Hub, and named `my-awesome-name/my-awesome-lora`, we can now use the finetuned model for inference:
```py
import torch
from diffusers import HunyuanVideoPipeline
import torch
from diffusers import HunyuanVideoPipeline, HunyuanVideoTransformer3DModel
from diffusers.utils import export_to_video
model_id = "tencent/HunyuanVideo"
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
model_id, subfolder="transformer", torch_dtype=torch.bfloat16
)
pipe = HunyuanVideoPipeline.from_pretrained(model_id, transformer=transformer, torch_dtype=torch.float16)
pipe.load_lora_weights("my-awesome-name/my-awesome-lora", adapter_name="hunyuanvideo-lora")
pipe.set_adapters(["hunyuanvideo-lora"], [0.6])
pipe.vae.enable_tiling()
pipe.to("cuda")
output = pipe(
prompt="A cat walks on the grass, realistic",
height=320,
width=512,
num_frames=61,
num_inference_steps=30,
).frames[0]
export_to_video(output, "output.mp4", fps=15)
```
</details>
If you would like to use a custom dataset, refer to the dataset preparation guide [here](./assets/dataset.md).
## Memory requirements
+3 -1
View File
@@ -206,7 +206,9 @@ def validate_args(args: Args):
def _add_model_arguments(parser: argparse.ArgumentParser) -> None:
parser.add_argument("--model_name", type=str, required=True, choices=["ltx_video"], help="Name of model to train.")
parser.add_argument(
"--model_name", type=str, required=True, choices=["hunyuan_video", "ltx_video"], help="Name of model to train."
)
parser.add_argument(
"--pretrained_model_name_or_path",
type=str,
+1
View File
@@ -0,0 +1 @@
from .hunyuan_video_lora import HUNYUAN_VIDEO_T2V_LORA_CONFIG
@@ -0,0 +1,319 @@
from typing import Any, Dict, List, Optional, Tuple, Union
import torch
import torch.nn as nn
from accelerate.logging import get_logger
from diffusers import (
AutoencoderKLHunyuanVideo,
FlowMatchEulerDiscreteScheduler,
HunyuanVideoPipeline,
HunyuanVideoTransformer3DModel,
)
from transformers import AutoTokenizer, CLIPTextModel, CLIPTokenizer, LlamaModel, LlamaTokenizer
from PIL import Image
logger = get_logger("finetrainers") # pylint: disable=invalid-name
def load_components(
model_id: str = "tencent/HunyuanVideo",
text_encoder_dtype: torch.dtype = torch.float16,
text_encoder_2_dtype: torch.dtype = torch.float16,
transformer_dtype: torch.dtype = torch.bfloat16,
vae_dtype: torch.dtype = torch.float16,
revision: Optional[str] = None,
cache_dir: Optional[str] = None,
) -> Dict[str, nn.Module]:
tokenizer = AutoTokenizer.from_pretrained(model_id, subfolder="tokenizer", revision=revision, cache_dir=cache_dir)
text_encoder = LlamaModel.from_pretrained(
model_id, subfolder="text_encoder", torch_dtype=text_encoder_dtype, revision=revision, cache_dir=cache_dir
)
tokenizer_2 = CLIPTokenizer.from_pretrained(
model_id, subfolder="tokenizer_2", revision=revision, cache_dir=cache_dir
)
text_encoder_2 = CLIPTextModel.from_pretrained(
model_id, subfolder="text_encoder_2", torch_dtype=text_encoder_2_dtype, revision=revision, cache_dir=cache_dir
)
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
model_id, subfolder="transformer", torch_dtype=transformer_dtype, revision=revision, cache_dir=cache_dir
)
vae = AutoencoderKLHunyuanVideo.from_pretrained(
model_id, subfolder="vae", torch_dtype=vae_dtype, revision=revision, cache_dir=cache_dir
)
scheduler = FlowMatchEulerDiscreteScheduler()
return {
"tokenizer": tokenizer,
"text_encoder": text_encoder,
"tokenizer_2": tokenizer_2,
"text_encoder_2": text_encoder_2,
"transformer": transformer,
"vae": vae,
"scheduler": scheduler,
}
def initialize_pipeline(
model_id: str = "tencent/HunyuanVideo",
text_encoder_dtype: torch.dtype = torch.float16,
text_encoder_2_dtype: torch.dtype = torch.float16,
transformer_dtype: torch.dtype = torch.bfloat16,
vae_dtype: torch.dtype = torch.float16,
tokenizer: Optional[LlamaTokenizer] = None,
text_encoder: Optional[LlamaModel] = None,
tokenizer_2: Optional[CLIPTokenizer] = None,
text_encoder_2: Optional[CLIPTextModel] = None,
transformer: Optional[HunyuanVideoTransformer3DModel] = None,
vae: Optional[AutoencoderKLHunyuanVideo] = None,
scheduler: Optional[FlowMatchEulerDiscreteScheduler] = 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,
) -> HunyuanVideoPipeline:
component_name_pairs = [
("tokenizer", tokenizer),
("text_encoder", text_encoder),
("tokenizer_2", tokenizer_2),
("text_encoder_2", text_encoder_2),
("transformer", transformer),
("vae", vae),
("scheduler", scheduler),
]
components = {}
for name, component in component_name_pairs:
if component is not None:
components[name] = component
pipe = HunyuanVideoPipeline.from_pretrained(model_id, **components, revision=revision, cache_dir=cache_dir)
pipe.text_encoder = pipe.text_encoder.to(dtype=text_encoder_dtype)
pipe.text_encoder_2 = pipe.text_encoder_2.to(dtype=text_encoder_2_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: LlamaTokenizer,
text_encoder: LlamaModel,
tokenizer_2: CLIPTokenizer,
text_encoder_2: CLIPTextModel,
prompt: Union[str, List[str]],
guidance: float = 1.0,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
max_sequence_length: int = 128,
# TODO(aryan): make configurable
prompt_template: Dict[str, Any] = {
"template": (
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
"1. The main content and theme of the video."
"2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects."
"3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
"4. background environment, light, style and atmosphere."
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
),
"crop_start": 95,
},
) -> torch.Tensor:
device = device or text_encoder.device
dtype = dtype or text_encoder.dtype
if isinstance(prompt, str):
prompt = [prompt]
conditions = {}
conditions.update(
_get_llama_prompt_embeds(tokenizer, text_encoder, prompt, prompt_template, device, dtype, max_sequence_length)
)
conditions.update(_get_clip_prompt_embeds(tokenizer_2, text_encoder_2, prompt, device, dtype))
guidance = torch.tensor([guidance], device=device, dtype=dtype) * 1000.0
conditions["guidance"] = guidance
return conditions
def prepare_latents(
vae: AutoencoderKLHunyuanVideo,
image_or_video: torch.Tensor,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
generator: Optional[torch.Generator] = None,
**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).contiguous() # [B, C, F, H, W] -> [B, F, C, H, W]
latents = vae.encode(image_or_video).latent_dist.sample(generator=generator)
latents = latents * vae.config.scaling_factor
latents = latents.to(dtype=dtype)
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 forward_pass(
transformer: HunyuanVideoTransformer3DModel,
prompt_embeds: torch.Tensor,
pooled_prompt_embeds: torch.Tensor,
prompt_attention_mask: torch.Tensor,
guidance: torch.Tensor,
latents: torch.Tensor,
noisy_latents: torch.Tensor,
timesteps: torch.LongTensor,
) -> torch.Tensor:
denoised_latents = transformer(
hidden_states=noisy_latents,
timestep=timesteps,
encoder_hidden_states=prompt_embeds,
pooled_projections=pooled_prompt_embeds,
encoder_attention_mask=prompt_attention_mask,
guidance=guidance,
return_dict=False,
)[0]
return {"latents": denoised_latents}
def validation(
pipeline: HunyuanVideoPipeline,
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_llama_prompt_embeds(
tokenizer: LlamaTokenizer,
text_encoder: LlamaModel,
prompt: List[str],
prompt_template: Dict[str, Any],
device: torch.device,
dtype: torch.dtype,
max_sequence_length: int = 256,
num_hidden_layers_to_skip: int = 2,
) -> Tuple[torch.Tensor, torch.Tensor]:
batch_size = len(prompt)
prompt = [prompt_template["template"].format(p) for p in prompt]
crop_start = prompt_template.get("crop_start", None)
if crop_start is None:
prompt_template_input = tokenizer(
prompt_template["template"],
padding="max_length",
return_tensors="pt",
return_length=False,
return_overflowing_tokens=False,
return_attention_mask=False,
)
crop_start = prompt_template_input["input_ids"].shape[-1]
# Remove <|eot_id|> token and placeholder {}
crop_start -= 2
max_sequence_length += crop_start
text_inputs = tokenizer(
prompt,
max_length=max_sequence_length,
padding="max_length",
truncation=True,
return_tensors="pt",
return_length=False,
return_overflowing_tokens=False,
return_attention_mask=True,
)
text_input_ids = text_inputs.input_ids.to(device=device)
prompt_attention_mask = text_inputs.attention_mask.to(device=device)
prompt_embeds = text_encoder(
input_ids=text_input_ids,
attention_mask=prompt_attention_mask,
output_hidden_states=True,
).hidden_states[-(num_hidden_layers_to_skip + 1)]
prompt_embeds = prompt_embeds.to(dtype=dtype)
if crop_start is not None and crop_start > 0:
prompt_embeds = prompt_embeds[:, crop_start:]
prompt_attention_mask = prompt_attention_mask[:, crop_start:]
prompt_attention_mask = prompt_attention_mask.view(batch_size, -1)
return {"prompt_embeds": prompt_embeds, "prompt_attention_mask": prompt_attention_mask}
def _get_clip_prompt_embeds(
tokenizer_2: CLIPTokenizer,
text_encoder_2: CLIPTextModel,
prompt: Union[str, List[str]],
device: torch.device,
dtype: torch.dtype,
max_sequence_length: int = 77,
) -> torch.Tensor:
text_inputs = tokenizer_2(
prompt,
padding="max_length",
max_length=max_sequence_length,
truncation=True,
return_tensors="pt",
)
prompt_embeds = text_encoder_2(text_inputs.input_ids.to(device), output_hidden_states=False).pooler_output
prompt_embeds = prompt_embeds.to(dtype=dtype)
return {"pooled_prompt_embeds": prompt_embeds}
HUNYUAN_VIDEO_T2V_LORA_CONFIG = {
"pipeline_cls": HunyuanVideoPipeline,
"load_components": load_components,
"initialize_pipeline": initialize_pipeline,
"prepare_conditions": prepare_conditions,
"prepare_latents": prepare_latents,
"collate_fn": collate_fn_t2v,
"forward_pass": forward_pass,
"validation": validation,
}
+1 -1
View File
@@ -1 +1 @@
from .ltx_video import LTX_VIDEO_T2V_CONFIG
from .ltx_video_lora import LTX_VIDEO_T2V_LORA_CONFIG
@@ -17,17 +17,20 @@ def load_components(
text_encoder_dtype: torch.dtype = torch.bfloat16,
transformer_dtype: torch.dtype = torch.bfloat16,
vae_dtype: torch.dtype = torch.bfloat16,
revision: Optional[str] = None,
cache_dir: Optional[str] = None,
) -> Dict[str, nn.Module]:
tokenizer = T5Tokenizer.from_pretrained(model_id, subfolder="tokenizer", cache_dir=cache_dir)
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, cache_dir=cache_dir
model_id, subfolder="text_encoder", torch_dtype=text_encoder_dtype, revision=revision, cache_dir=cache_dir
)
transformer = LTXVideoTransformer3DModel.from_pretrained(
model_id, subfolder="transformer", torch_dtype=transformer_dtype, cache_dir=cache_dir
model_id, subfolder="transformer", torch_dtype=transformer_dtype, revision=revision, cache_dir=cache_dir
)
vae = AutoencoderKLLTXVideo.from_pretrained(model_id, subfolder="vae", torch_dtype=vae_dtype, cache_dir=cache_dir)
scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(model_id, subfolder="scheduler", cache_dir=cache_dir)
vae = AutoencoderKLLTXVideo.from_pretrained(
model_id, subfolder="vae", torch_dtype=vae_dtype, revision=revision, cache_dir=cache_dir
)
scheduler = FlowMatchEulerDiscreteScheduler()
return {
"tokenizer": tokenizer,
"text_encoder": text_encoder,
@@ -48,10 +51,12 @@ def initialize_pipeline(
vae: Optional[AutoencoderKLLTXVideo] = None,
scheduler: Optional[FlowMatchEulerDiscreteScheduler] = 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,
) -> LTXPipeline:
component_name_pairs = [
("tokenizer", tokenizer),
@@ -65,7 +70,7 @@ def initialize_pipeline(
if component is not None:
components[name] = component
pipe = LTXPipeline.from_pretrained(model_id, **components, cache_dir=cache_dir)
pipe = LTXPipeline.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)
@@ -90,6 +95,7 @@ def prepare_conditions(
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
max_sequence_length: int = 128,
**kwargs,
) -> torch.Tensor:
device = device or text_encoder.device
dtype = dtype or text_encoder.dtype
@@ -252,7 +258,7 @@ def _pack_latents(latents: torch.Tensor, patch_size: int = 1, patch_size_t: int
return latents
LTX_VIDEO_T2V_CONFIG = {
LTX_VIDEO_T2V_LORA_CONFIG = {
"pipeline_cls": LTXPipeline,
"load_components": load_components,
"initialize_pipeline": initialize_pipeline,
+14 -4
View File
@@ -1,16 +1,26 @@
from typing import Any, Dict
from .ltx_video import LTX_VIDEO_T2V_CONFIG
from .hunyuan_video import HUNYUAN_VIDEO_T2V_LORA_CONFIG
from .ltx_video import LTX_VIDEO_T2V_LORA_CONFIG
SUPPORTED_MODEL_CONFIGS = {
"ltx_video": LTX_VIDEO_T2V_CONFIG,
"hunyuan_video": {
"lora": HUNYUAN_VIDEO_T2V_LORA_CONFIG,
},
"ltx_video": {
"lora": LTX_VIDEO_T2V_LORA_CONFIG,
},
}
def get_config_from_model_name(model_name: str) -> Dict[str, Any]:
def get_config_from_model_name(model_name: str, training_type: str) -> Dict[str, Any]:
if model_name not in SUPPORTED_MODEL_CONFIGS:
raise ValueError(
f"Model {model_name} not supported. Supported models are: {list(SUPPORTED_MODEL_CONFIGS.keys())}"
)
return SUPPORTED_MODEL_CONFIGS[model_name]
if training_type not in SUPPORTED_MODEL_CONFIGS[model_name]:
raise ValueError(
f"Training type {training_type} not supported for model {model_name}. Supported training types are: {list(SUPPORTED_MODEL_CONFIGS[model_name].keys())}"
)
return SUPPORTED_MODEL_CONFIGS[model_name][training_type]
+28 -9
View File
@@ -78,7 +78,7 @@ class Trainer:
self._init_directories_and_repositories()
self.state.model_name = self.args.model_name
self.model_config = get_config_from_model_name(self.args.model_name)
self.model_config = get_config_from_model_name(self.args.model_name, self.args.training_type)
def prepare_models(self) -> None:
logger.info("Initializing models")
@@ -88,6 +88,7 @@ class Trainer:
"text_encoder_dtype": torch.bfloat16,
"transformer_dtype": torch.bfloat16,
"vae_dtype": torch.bfloat16,
"revision": self.args.revision,
"cache_dir": self.args.cache_dir,
}
if self.args.pretrained_model_name_or_path is not None:
@@ -96,10 +97,18 @@ class Trainer:
self.tokenizer = components.get("tokenizer", None)
self.text_encoder = components.get("text_encoder", None)
self.tokenizer_2 = components.get("tokenizer_2", None)
self.text_encoder_2 = components.get("text_encoder_2", None)
self.transformer = components.get("transformer", None)
self.vae = components.get("vae", None)
self.scheduler = components.get("scheduler", None)
if self.vae is not None:
if self.args.enable_slicing:
self.vae.enable_slicing()
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_dataset(self) -> None:
@@ -130,6 +139,9 @@ class Trainer:
self.transformer.requires_grad_(False)
self.vae.requires_grad_(False)
if self.text_encoder_2 is not None:
self.text_encoder_2.requires_grad_(False)
# For mixed precision training we cast all non-trainable weights (vae, text_encoder and transformer) to half-precision
# as these weights are only used for inference, keeping weights in full precision is not required.
weight_dtype = torch.float32
@@ -144,12 +156,15 @@ class Trainer:
"Mixed precision training with bfloat16 is not supported on MPS. Please use fp16 (recommended) or fp32 instead."
)
# TODO(aryan): handle torch dtype from accelerator vs model dtype
# TODO(aryan): handle torch dtype from accelerator vs model dtype; refactor
self.state.weight_dtype = weight_dtype
self.text_encoder.to(self.state.accelerator.device, dtype=weight_dtype)
self.transformer.to(self.state.accelerator.device, dtype=weight_dtype)
self.vae.to(self.state.accelerator.device, dtype=weight_dtype)
if self.text_encoder_2 is not None:
self.text_encoder_2.to(self.state.accelerator.device, dtype=weight_dtype)
if self.args.gradient_checkpointing:
self.transformer.enable_gradient_checkpointing()
@@ -320,7 +335,6 @@ class Trainer:
logger.info(f"Training configuration: {json.dumps(info, indent=4)}")
# TODO(aryan): handle resume from checkpoint
global_step = 0
first_epoch = 0
initial_global_step = 0
@@ -372,6 +386,8 @@ class Trainer:
other_conditions = self.model_config["prepare_conditions"](
tokenizer=self.tokenizer,
text_encoder=self.text_encoder,
tokenizer_2=self.tokenizer_2,
text_encoder_2=self.text_encoder_2,
prompt=prompts,
device=accelerator.device,
dtype=weight_dtype,
@@ -383,6 +399,10 @@ class Trainer:
other_conditions["prompt_embeds"].fill_(0)
other_conditions["prompt_attention_mask"].fill_(False)
# TODO(aryan): refactor later
if "pooled_prompt_embeds" in other_conditions:
other_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,
@@ -392,11 +412,7 @@ class Trainer:
mode_scale=self.args.flow_mode_scale,
)
indices = (weights * self.scheduler.config.num_train_timesteps).long()
sigmas = scheduler_sigmas[indices].flatten()
while sigmas.ndim < latent_conditions["latents"].ndim:
sigmas = sigmas.unsqueeze(-1)
sigmas = scheduler_sigmas[indices]
timesteps = (sigmas * 1000.0).long()
noise = torch.randn(
@@ -524,12 +540,15 @@ class Trainer:
pipeline = self.model_config["initialize_pipeline"](
model_id=self.args.pretrained_model_name_or_path,
cache_dir=self.args.cache_dir,
tokenizer=self.tokenizer,
text_encoder=self.text_encoder,
tokenizer_2=self.tokenizer_2,
text_encoder_2=self.text_encoder_2,
transformer=unwrap_model(accelerator, self.transformer),
vae=self.vae,
device=accelerator.device,
revision=self.args.revision,
cache_dir=self.args.cache_dir,
enable_slicing=self.args.enable_slicing,
enable_tiling=self.args.enable_tiling,
enable_model_cpu_offload=self.args.enable_model_cpu_offload,