diff --git a/training/mochi-1/dataset_mochi.py b/training/mochi-1/dataset_mochi.py index c40a747..2e58f3a 100644 --- a/training/mochi-1/dataset_mochi.py +++ b/training/mochi-1/dataset_mochi.py @@ -5,7 +5,6 @@ import numpy as np import torch import torchvision.transforms as TT from accelerate.logging import get_logger -from torchvision import transforms from torchvision.transforms import InterpolationMode from torchvision.transforms.functional import resize diff --git a/training/mochi-1/text_to_video_lora.py b/training/mochi-1/text_to_video_lora.py index 9fc4ac6..4ffe908 100644 --- a/training/mochi-1/text_to_video_lora.py +++ b/training/mochi-1/text_to_video_lora.py @@ -76,7 +76,7 @@ def save_model_card( fps=8, ): widget_dict = [] - if videos is not None: + if videos is not None and len(videos) > 0: for i, video in enumerate(videos): export_to_video(video, os.path.join(repo_folder, f"final_video_{i}.mp4", fps=fps)) widget_dict.append( @@ -876,34 +876,34 @@ def main(args): torch.cuda.empty_cache() torch.cuda.synchronize(accelerator.device) - accelerator.print("===== Memory before testing =====") - print_memory(accelerator.device) - reset_memory(accelerator.device) - # Final test inference - pipe = MochiPipeline.from_pretrained( - args.pretrained_model_name_or_path, - revision=args.revision, - variant=args.variant, - # torch_dtype=weight_dtype, - ) - pipe.scheduler = FlowMatchEulerDiscreteScheduler.from_config(pipe.scheduler.config) - - if args.enable_slicing: - pipe.vae.enable_slicing() - if args.enable_tiling: - pipe.vae.enable_tiling() - if args.enable_model_cpu_offload: - pipe.enable_model_cpu_offload() - - # Load LoRA weights - lora_scaling = args.lora_alpha / args.rank - pipe.load_lora_weights(args.output_dir, adapter_name="mochi-lora") - pipe.set_adapters(["mochi-lora"], [lora_scaling]) - - # Run inference validation_outputs = [] if args.validation_prompt and args.num_validation_videos > 0: + accelerator.print("===== Memory before testing =====") + print_memory(accelerator.device) + reset_memory(accelerator.device) + + pipe = MochiPipeline.from_pretrained( + args.pretrained_model_name_or_path, + revision=args.revision, + variant=args.variant, + # torch_dtype=weight_dtype, + ) + pipe.scheduler = FlowMatchEulerDiscreteScheduler.from_config(pipe.scheduler.config) + + if args.enable_slicing: + pipe.vae.enable_slicing() + if args.enable_tiling: + pipe.vae.enable_tiling() + if args.enable_model_cpu_offload: + pipe.enable_model_cpu_offload() + + # Load LoRA weights + lora_scaling = args.lora_alpha / args.rank + pipe.load_lora_weights(args.output_dir, adapter_name="mochi-lora") + pipe.set_adapters(["mochi-lora"], [lora_scaling]) + + # Run inference validation_prompts = args.validation_prompt.split(args.validation_prompt_separator) for validation_prompt in validation_prompts: pipeline_args = { @@ -924,10 +924,10 @@ def main(args): ) validation_outputs.extend(video) - accelerator.print("===== Memory after testing =====") - print_memory(accelerator.device) - reset_memory(accelerator.device) - torch.cuda.synchronize(accelerator.device) + accelerator.print("===== Memory after testing =====") + print_memory(accelerator.device) + reset_memory(accelerator.device) + torch.cuda.synchronize(accelerator.device) if args.push_to_hub: save_model_card( @@ -942,7 +942,7 @@ def main(args): repo_id=repo_id, folder_path=args.output_dir, commit_message="End of training", - ignore_patterns=["step_*", "epoch_*", "*.bin"], + ignore_patterns=["step_*", "epoch_*", "*.bin", "*.pt"], ) accelerator.end_training()