diff --git a/README.md b/README.md index 782d625..11b2738 100644 --- a/README.md +++ b/README.md @@ -376,6 +376,111 @@ output = pipe( export_to_video(output, "output.mp4", fps=15) ``` +### Memory Usage + +LoRA with rank 128, batch size 1, gradient checkpointing, optimizer adamw, `49x512x768` resolutions, **without precomputation**: + +``` +Memory before training start: { + "memory_allocated": 38.889, + "memory_reserved": 39.02, + "max_memory_allocated": 38.889, + "max_memory_reserved": 39.02 +} +Training configuration: { + "trainable parameters": 163577856, + "total samples": 69, + "train epochs": 1, + "train steps": 10, + "batches per device": 1, + "total batches observed per epoch": 69, + "train batch size": 1, + "gradient accumulation steps": 1 +} +Memory before validation start: { + "memory_allocated": 39.747, + "memory_reserved": 56.266, + "max_memory_allocated": 51.867, + "max_memory_reserved": 56.266 +} +Memory after validation end: { + "memory_allocated": 39.748, + "memory_reserved": 41.445, + "max_memory_allocated": 51.867, + "max_memory_reserved": 58.385 +} +Memory after epoch 1: { + "memory_allocated": 39.748, + "memory_reserved": 40.91, + "max_memory_allocated": 39.748, + "max_memory_reserved": 40.91 +} +Memory after training end: { + "memory_allocated": 25.288, + "memory_reserved": 27.783, + "max_memory_allocated": 39.748, + "max_memory_reserved": 40.91 +} +``` + +LoRA with rank 128, batch size 1, gradient checkpointing, optimizer adamw, `49x512x768` resolutions, **with precomputation**: + +``` +Memory after precomputing conditions: { + "memory_allocated": 14.232, + "memory_reserved": 14.336, + "max_memory_allocated": 14.395, + "max_memory_reserved": 14.461 +} +Memory after precomputing latents: { + "memory_allocated": 14.717, + "memory_reserved": 14.762, + "max_memory_allocated": 16.759, + "max_memory_reserved": 17.244 +} +Memory before training start: { + "memory_allocated": 24.195, + "memory_reserved": 26.039, + "max_memory_allocated": 24.195, + "max_memory_reserved": 26.039 +} +Training configuration: { + "trainable parameters": 163577856, + "total samples": 1, + "train epochs": 10, + "train steps": 10, + "batches per device": 1, + "total batches observed per epoch": 1, + "train batch size": 1, + "gradient accumulation steps": 1 +} +Memory after epoch 1: { + "memory_allocated": 24.83, + "memory_reserved": 42.387, + "max_memory_allocated": 36.357, + "max_memory_reserved": 42.387 +} +Memory before validation start: { + "memory_allocated": 24.842, + "memory_reserved": 42.387, + "max_memory_allocated": 36.977, + "max_memory_reserved": 42.387 +} + +Memory after validation end: { + "memory_allocated": 39.558, + "memory_reserved": 41.039, + "max_memory_allocated": 43.226, + "max_memory_reserved": 46.947 +} +Memory after training end: { + "memory_allocated": 24.842, + "memory_reserved": 26.82, + "max_memory_allocated": 39.558, + "max_memory_reserved": 41.039 +} +``` + If you would like to use a custom dataset, refer to the dataset preparation guide [here](./assets/dataset.md). diff --git a/finetrainers/hunyuan_video/hunyuan_video_lora.py b/finetrainers/hunyuan_video/hunyuan_video_lora.py index 609f6fc..5fefb9c 100644 --- a/finetrainers/hunyuan_video/hunyuan_video_lora.py +++ b/finetrainers/hunyuan_video/hunyuan_video_lora.py @@ -16,41 +16,44 @@ from PIL import Image logger = get_logger("finetrainers") # pylint: disable=invalid-name -def load_components( +def load_condition_models( 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, + revision: Optional[str] = None, + cache_dir: Optional[str] = None, + **kwargs, +) -> 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) + return {"tokenizer": tokenizer, "text_encoder": text_encoder, "tokenizer_2": tokenizer_2, "text_encoder_2": text_encoder_2} + + +def load_latent_models( + model_id: str = "tencent/HunyuanVideo", vae_dtype: torch.dtype = torch.float16, revision: Optional[str] = None, cache_dir: Optional[str] = None, + **kwargs, +) -> Dict[str, nn.Module]: + vae = AutoencoderKLHunyuanVideo.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 = "tencent/HunyuanVideo", + transformer_dtype: torch.dtype = torch.bfloat16, + revision: Optional[str] = None, + cache_dir: Optional[str] = None, + **kwargs, ) -> 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, - } + return {"transformer": transformer, "scheduler": scheduler} def initialize_pipeline( @@ -72,6 +75,7 @@ def initialize_pipeline( enable_slicing: bool = False, enable_tiling: bool = False, enable_model_cpu_offload: bool = False, + **kwargs, ) -> HunyuanVideoPipeline: component_name_pairs = [ ("tokenizer", tokenizer), @@ -115,7 +119,7 @@ def prepare_conditions( guidance: float = 1.0, device: Optional[torch.device] = None, dtype: Optional[torch.dtype] = None, - max_sequence_length: int = 128, + max_sequence_length: int = 256, # TODO(aryan): make configurable prompt_template: Dict[str, Any] = { "template": ( @@ -129,6 +133,7 @@ def prepare_conditions( ), "crop_start": 95, }, + **kwargs, ) -> torch.Tensor: device = device or text_encoder.device dtype = dtype or text_encoder.dtype @@ -154,6 +159,7 @@ def prepare_latents( 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 @@ -165,9 +171,24 @@ def prepare_latents( 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) + if not precompute: + 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} + else: + 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( + latents: torch.Tensor, + **kwargs, +) -> torch.Tensor: return {"latents": latents} @@ -309,10 +330,13 @@ def _get_clip_prompt_embeds( HUNYUAN_VIDEO_T2V_LORA_CONFIG = { "pipeline_cls": HunyuanVideoPipeline, - "load_components": load_components, + "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, "forward_pass": forward_pass, "validation": validation, diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index 87c8b4c..970f697 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -319,7 +319,7 @@ class Trainer: generator=self.state.generator, precompute=True, ) - filename = latents_dir / f"latents-{i}-{index}.pt" + filename = latents_dir / f"latents-{self.state.accelerator.process_index}-{index}.pt" torch.save(latent_conditions, filename.as_posix()) index += 1 progress_bar.update(1)