diff --git a/Makefile b/Makefile index be64981..334a3f2 100644 --- a/Makefile +++ b/Makefile @@ -1,6 +1,6 @@ .PHONY: quality style -check_dirs := training tests +check_dirs := finetrainers tests quality: ruff check $(check_dirs) diff --git a/README.md b/README.md index 60dfc7f..2dd6e67 100644 --- a/README.md +++ b/README.md @@ -30,6 +30,8 @@ huggingface-cli download \ Then launch LoRA fine-tuning. For CogVideoX and Mochi, refer to [this](./training/README.md) and [this](./training/mochi-1/README.md). +Note: It is recommended to use Pytorch 2.5.1 or above for training. Previous versions can lead to completely black videos, OOM errors, or other issues and are not tested. +
LTX Video @@ -52,6 +54,8 @@ CAPTION_COLUMN="prompts.txt" VIDEO_COLUMN="videos.txt" OUTPUT_DIR="/path/to/output/directory/ltx-video/ltxv_disney" +ID_TOKEN="BW_STYLE" + # Model arguments model_cmd="--model_name ltx_video \ --pretrained_model_name_or_path Lightricks/LTX-Video" @@ -60,7 +64,7 @@ model_cmd="--model_name ltx_video \ dataset_cmd="--data_root $DATA_ROOT \ --video_column $VIDEO_COLUMN \ --caption_column $CAPTION_COLUMN \ - --id_token BW_STYLE \ + --id_token $ID_TOKEN \ --video_resolution_buckets 49x512x768 \ --caption_dropout_p 0.05" @@ -99,7 +103,7 @@ optimizer_cmd="--optimizer adamw \ --max_grad_norm 1.0" # Validation arguments -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\" \ +validation_cmd="--validation_prompts \"$ID_TOKEN 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:::$ID_TOKEN 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" @@ -221,6 +225,8 @@ 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" +ID_TOKEN="afkx" + # Model arguments model_cmd="--model_name hunyuan_video \ --pretrained_model_name_or_path hunyuanvideo-community/HunyuanVideo" @@ -229,8 +235,8 @@ model_cmd="--model_name hunyuan_video \ dataset_cmd="--data_root $DATA_ROOT \ --video_column $VIDEO_COLUMN \ --caption_column $CAPTION_COLUMN \ - --id_token afkx \ - --video_resolution_buckets 17x512x768 49x512x768 61x512x768 129x512x768 \ + --id_token $ID_TOKEN \ + --video_resolution_buckets 17x512x768 49x512x768 61x512x768 \ --caption_dropout_p 0.05" # Dataloader arguments @@ -268,7 +274,7 @@ optimizer_cmd="--optimizer adamw \ --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\" \ +validation_cmd="--validation_prompts \"$ID_TOKEN 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:::$ID_TOKEN 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:::$ID_TOKEN 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:::$ID_TOKEN 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:::$ID_TOKEN 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:::$ID_TOKEN 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:::$ID_TOKEN 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:::$ID_TOKEN 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" @@ -350,7 +356,7 @@ Training configuration: { | after epoch 1 | 39.748 | 40.910 | | after training end | 25.288 | 40.910 | -Note: requires about `59` GB of VRAM without precomputation. +Note: requires about `59` GB of VRAM when validation is performed. LoRA with rank 128, batch size 1, gradient checkpointing, optimizer adamw, `49x512x768` resolutions, **with precomputation**: @@ -377,7 +383,7 @@ Training configuration: { | after validation end | 39.558 | 46.947 | | after training end | 24.842 | 41.039 | -Note: requires about `47` GB of VRAM with precomputation. If validation is not performed, the memory usage is reduced to about `42` GB. +Note: requires about `47` GB of VRAM with validation. If validation is not performed, the memory usage is reduced to about `42` GB.
@@ -385,6 +391,7 @@ If you would like to use a custom dataset, refer to the dataset preparation guid > [!NOTE] > To lower memory requirements: +> - Use a DeepSpeed config to launch training (refer to [`accelerate_configs/deepspeed.yaml`](./accelerate_configs/deepspeed.yaml) as an example). > - Pass `--precompute_conditions` when launching training. > - Pass `--gradient_checkpointing` when launching training. > - Do not perform validation/testing. This saves a significant amount of memory, which can be used to focus solely on training if you're on smaller VRAM GPUs. diff --git a/accelerate_configs/deepspeed.yaml b/accelerate_configs/deepspeed.yaml index 2827648..62db0b4 100644 --- a/accelerate_configs/deepspeed.yaml +++ b/accelerate_configs/deepspeed.yaml @@ -20,4 +20,4 @@ same_network: true tpu_env: [] tpu_use_cluster: false tpu_use_sudo: false -use_cpu: false +use_cpu: false \ No newline at end of file diff --git a/accelerate_configs/uncompiled_2.yaml b/accelerate_configs/uncompiled_2.yaml index c5216da..830b6e0 100644 --- a/accelerate_configs/uncompiled_2.yaml +++ b/accelerate_configs/uncompiled_2.yaml @@ -14,4 +14,4 @@ same_network: true tpu_env: [] tpu_use_cluster: false tpu_use_sudo: false -use_cpu: false +use_cpu: false \ No newline at end of file diff --git a/finetrainers/args.py b/finetrainers/args.py index 31c2a76..b32a7a5 100644 --- a/finetrainers/args.py +++ b/finetrainers/args.py @@ -60,7 +60,9 @@ class Args: # Training arguments training_type: str = None seed: int = 42 - mixed_precision: str = None + mixed_precision: str = ( + None # TODO: consider removing later https://github.com/a-r-r-o-w/finetrainers/pull/139#discussion_r1897438414 + ) batch_size: int = 1 train_epochs: int = 1 train_steps: int = None @@ -675,6 +677,7 @@ _DTYPE_MAP = { "fp16": torch.float16, "fp32": torch.float32, } +_INVERSE_DTYPE_MAP = {v: k for k, v in _DTYPE_MAP.items()} def _map_to_args_type(args: Dict[str, Any]) -> Args: diff --git a/finetrainers/state.py b/finetrainers/state.py index f30b8c2..15a92e2 100644 --- a/finetrainers/state.py +++ b/finetrainers/state.py @@ -15,6 +15,7 @@ class State: learning_rate: float = None train_batch_size: int = None generator: torch.Generator = None + num_update_steps_per_epoch: int = None # Hub state repo_id: str = None diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index 6f2348f..bbacc49 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -1,10 +1,8 @@ -import inspect import json import logging import math import os import random -import shutil from datetime import timedelta from typing import Any, Dict from pathlib import Path @@ -35,7 +33,7 @@ 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 -from .args import Args, validate_args +from .args import Args, validate_args, _INVERSE_DTYPE_MAP from .constants import ( FINETRAINERS_LOG_LEVEL, PRECOMPUTED_DIR_NAME, @@ -46,10 +44,11 @@ from .dataset import BucketSampler, PrecomputedDataset, VideoDatasetWithResizing from .models import get_config_from_model_name from .state import State from .utils.data_utils import should_perform_precomputation -from .utils.file_utils import find_files, delete_files, string_to_filename -from .utils.optimizer_utils import get_optimizer, gradient_norm +from .utils.file_utils import string_to_filename +from .utils.optimizer_utils import get_optimizer from .utils.memory_utils import get_memory_statistics, free_memory, make_contiguous -from .utils.torch_utils import unwrap_model, align_device_and_dtype +from .utils.torch_utils import unwrap_model, align_device_and_dtype, expand_tensor_to_dims +from .utils.checkpointing import get_latest_ckpt_path_to_resume_from, get_intermediate_ckpt_path logger = get_logger("finetrainers") @@ -361,11 +360,7 @@ class Trainer: # 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 - if self.state.accelerator.mixed_precision == "fp16": - weight_dtype = torch.float16 - elif self.state.accelerator.mixed_precision == "bf16": - weight_dtype = torch.bfloat16 + weight_dtype = self._get_training_dtype(accelerator=self.state.accelerator) if torch.backends.mps.is_available() and weight_dtype == torch.bfloat16: # due to pytorch#99272, MPS does not yet support bfloat16. @@ -375,6 +370,11 @@ class Trainer: # TODO(aryan): handle torch dtype from accelerator vs model dtype; refactor self.state.weight_dtype = weight_dtype + if self.args.mixed_precision != _INVERSE_DTYPE_MAP[weight_dtype]: + logger.warning( + f"`mixed_precision` was set to {_INVERSE_DTYPE_MAP[weight_dtype]} which is different from configured argument ({self.args.mixed_precision})." + ) + self.args.mixed_precision = _INVERSE_DTYPE_MAP[weight_dtype] self.transformer.to(dtype=weight_dtype) self._move_components_to_device() @@ -389,7 +389,13 @@ class Trainer: ) self.transformer.add_adapter(transformer_lora_config) - # TODO: refactor + # Enable TF32 for faster training on Ampere GPUs: https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices + if self.args.allow_tf32 and torch.cuda.is_available(): + torch.backends.cuda.matmul.allow_tf32 = True + + self.register_saving_loading_hooks(transformer_lora_config) + + def register_saving_loading_hooks(self, transformer_lora_config): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): if self.state.accelerator.is_main_process: @@ -415,13 +421,25 @@ class Trainer: ) def load_model_hook(models, input_dir): - transformer_ = self.model_config["pipeline_cls"].from_pretrained( - self.args.pretrained_model_name_or_path, subfolder="transformer" - ) - transformer_.add_adapter(transformer_lora_config) + if not self.state.accelerator.distributed_type == DistributedType.DEEPSPEED: + while len(models) > 0: + model = models.pop() + if isinstance( + unwrap_model(self.state.accelerator, model), + type(unwrap_model(self.state.accelerator, self.transformer)), + ): + transformer_ = unwrap_model(self.state.accelerator, model) + else: + raise ValueError( + f"Unexpected save model: {unwrap_model(self.state.accelerator, model).__class__}" + ) + else: + transformer_ = unwrap_model(self.state.accelerator, self.transformer).__class__.from_pretrained( + self.args.pretrained_model_name_or_path, subfolder="transformer" + ) + transformer_.add_adapter(transformer_lora_config) lora_state_dict = self.model_config["pipeline_cls"].lora_state_dict(input_dir) - transformer_state_dict = { f'{k.replace("transformer.", "")}': v for k, v in lora_state_dict.items() @@ -447,10 +465,6 @@ class Trainer: self.state.accelerator.register_save_state_pre_hook(save_model_hook) self.state.accelerator.register_load_state_pre_hook(load_model_hook) - # Enable TF32 for faster training on Ampere GPUs: https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices - if self.args.allow_tf32 and torch.cuda.is_available(): - torch.backends.cuda.matmul.allow_tf32 = True - def prepare_optimizer(self) -> None: logger.info("Initializing optimizer and lr scheduler") @@ -479,16 +493,20 @@ class Trainer: params_to_optimize = [transformer_parameters_with_lr] self.state.num_trainable_parameters = sum(p.numel() for p in transformer_lora_parameters) - # TODO(aryan): add deepspeed support + use_deepspeed_opt = ( + self.state.accelerator.state.deepspeed_plugin is not None + and "optimizer" in self.state.accelerator.state.deepspeed_plugin.deepspeed_config + ) optimizer = get_optimizer( params_to_optimize=params_to_optimize, optimizer_name=self.args.optimizer, - learning_rate=self.args.lr, + learning_rate=self.state.learning_rate, beta1=self.args.beta1, beta2=self.args.beta2, beta3=self.args.beta3, epsilon=self.args.epsilon, weight_decay=self.args.weight_decay, + use_deepspeed=use_deepspeed_opt, ) num_update_steps_per_epoch = math.ceil(len(self.dataloader) / self.args.gradient_accumulation_steps) @@ -496,14 +514,31 @@ class Trainer: self.state.train_steps = self.state.train_epochs * num_update_steps_per_epoch self.state.overwrote_max_train_steps = True - lr_scheduler = get_scheduler( - name=self.args.lr_scheduler, - optimizer=optimizer, - num_warmup_steps=self.args.lr_warmup_steps * self.state.accelerator.num_processes, - num_training_steps=self.state.train_steps * self.state.accelerator.num_processes, - num_cycles=self.args.lr_num_cycles, - power=self.args.lr_power, + use_deepspeed_lr_scheduler = ( + self.state.accelerator.state.deepspeed_plugin is not None + and "scheduler" in self.state.accelerator.state.deepspeed_plugin.deepspeed_config ) + total_training_steps = self.state.train_steps * self.state.accelerator.num_processes + num_warmup_steps = self.args.lr_warmup_steps * self.state.accelerator.num_processes + + if use_deepspeed_lr_scheduler: + from accelerate.utils import DummyScheduler + + lr_scheduler = DummyScheduler( + name=self.args.lr_scheduler, + optimizer=optimizer, + total_num_steps=total_training_steps, + num_warmup_steps=num_warmup_steps, + ) + else: + lr_scheduler = get_scheduler( + name=self.args.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=num_warmup_steps, + num_training_steps=total_training_steps, + num_cycles=self.args.lr_num_cycles, + power=self.args.lr_power, + ) self.optimizer = optimizer self.lr_scheduler = lr_scheduler @@ -519,6 +554,7 @@ class Trainer: self.state.train_steps = self.state.train_epochs * num_update_steps_per_epoch # Afterwards we recalculate our number of training epochs self.state.train_epochs = math.ceil(self.state.train_steps / num_update_steps_per_epoch) + self.state.num_update_steps_per_epoch = num_update_steps_per_epoch def prepare_trackers(self) -> None: logger.info("Initializing trackers") @@ -547,10 +583,24 @@ 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 + + # Potentially load in the weights and states from a previous save + ( + resume_from_checkpoint_path, + initial_global_step, + global_step, + first_epoch, + ) = get_latest_ckpt_path_to_resume_from( + resume_from_checkpoint=self.args.resume_from_checkpoint, + num_update_steps_per_epoch=self.state.num_update_steps_per_epoch, + output_dir=self.args.output_dir, + ) + if resume_from_checkpoint_path: + self.state.accelerator.load_state(resume_from_checkpoint_path) + progress_bar = tqdm( range(0, self.state.train_steps), initial=initial_global_step, @@ -646,6 +696,7 @@ class Trainer: 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 latent_conditions.update({"noisy_latents": noisy_latents}) @@ -666,8 +717,19 @@ class Trainer: loss = loss.mean() accelerator.backward(loss) - if accelerator.sync_gradients and accelerator.distributed_type != DistributedType.DEEPSPEED: - grad_norm = accelerator.clip_grad_norm_(self.transformer.parameters(), self.args.max_grad_norm) + if accelerator.sync_gradients: + if accelerator.distributed_type == DistributedType.DEEPSPEED: + grad_norm = self.transformer.get_global_grad_norm() + # In some cases the grad norm may not return a float + if torch.is_tensor(grad_norm): + grad_norm = grad_norm.item() + else: + grad_norm = accelerator.clip_grad_norm_( + self.transformer.parameters(), self.args.max_grad_norm + ) + if torch.is_tensor(grad_norm): + grad_norm = grad_norm.item() + logs["grad_norm"] = grad_norm self.optimizer.step() @@ -682,20 +744,12 @@ class Trainer: # Checkpointing if accelerator.distributed_type == DistributedType.DEEPSPEED or accelerator.is_main_process: if global_step % self.args.checkpointing_steps == 0: - # before saving state, check if this save would set us over the `checkpointing_limit` - if self.args.checkpointing_limit is not None: - checkpoints = find_files(self.args.output_dir, prefix="checkpoint") - - # before we save the new checkpoint, we need to have at_most `checkpoints_total_limit - 1` checkpoints - if len(checkpoints) >= self.args.checkpointing_limit: - num_to_remove = len(checkpoints) - self.args.checkpointing_limit + 1 - checkpoints_to_remove = checkpoints[0:num_to_remove] - delete_files(checkpoints_to_remove) - - logger.info(f"Checkpointing at step {global_step}") - save_path = os.path.join(self.args.output_dir, f"checkpoint-{global_step}") + save_path = get_intermediate_ckpt_path( + checkpointing_limit=self.args.checkpointing_limit, + step=global_step, + output_dir=self.args.output_dir, + ) accelerator.save_state(save_path) - logger.info(f"Saved state to {save_path}") # Maybe run validation should_run_validation = ( @@ -705,7 +759,8 @@ class Trainer: if should_run_validation: self.validate(global_step) - logs = {"loss": loss.detach().item(), "lr": self.lr_scheduler.get_last_lr()[0]} + logs["loss"] = loss.detach().item() + logs["lr"] = self.lr_scheduler.get_last_lr()[0] progress_bar.set_postfix(logs) accelerator.log(logs, step=global_step) @@ -725,15 +780,8 @@ class Trainer: accelerator.wait_for_everyone() if accelerator.is_main_process: + # TODO: consider factoring this out when supporting other types of training algos. self.transformer = unwrap_model(accelerator, self.transformer) - dtype = ( - torch.float16 - if self.args.mixed_precision == "fp16" - else torch.bfloat16 - if self.args.mixed_precision == "bf16" - else torch.float32 - ) - self.transformer = self.transformer.to(dtype) transformer_lora_layers = get_peft_model_state_dict(self.transformer) self.model_config["pipeline_cls"].save_lora_weights( @@ -741,6 +789,14 @@ class Trainer: transformer_lora_layers=transformer_lora_layers, ) + self.validate(step=global_step, final_validation=True) + + if accelerator.is_main_process: + if self.args.push_to_hub: + upload_folder( + repo_id=self.state.repo_id, folder_path=self.args.output_dir, ignore_patterns=["checkpoint-*"] + ) + del self.tokenizer, self.text_encoder, self.transformer, self.vae, self.scheduler free_memory() memory_statistics = get_memory_statistics() @@ -748,7 +804,7 @@ class Trainer: accelerator.end_training() - def validate(self, step: int) -> None: + def validate(self, step: int, final_validation: bool = False) -> None: logger.info("Starting validation") accelerator = self.state.accelerator @@ -763,21 +819,35 @@ class Trainer: memory_statistics = get_memory_statistics() logger.info(f"Memory before validation start: {json.dumps(memory_statistics, indent=4)}") - pipeline = self.model_config["initialize_pipeline"]( - model_id=self.args.pretrained_model_name_or_path, - 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, - ) + if not final_validation: + pipeline = self.model_config["initialize_pipeline"]( + model_id=self.args.pretrained_model_name_or_path, + 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, + ) + else: + # `torch_dtype` is manually set within `initialize_pipeline()`. + self._delete_components() + pipeline = self.model_config["initialize_pipeline"]( + model_id=self.args.pretrained_model_name_or_path, + 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, + ) + pipeline.load_lora_weights(self.args.output_dir) all_processes_artifacts = [] for i in range(num_validation_samples): @@ -811,6 +881,7 @@ class Trainer: num_frames=num_frames, num_videos_per_prompt=self.args.num_validation_videos_per_prompt, generator=self.state.generator, + # todo support passing `fps` for supported pipelines. ) prompt_filename = string_to_filename(prompt)[:25] @@ -841,6 +912,7 @@ class Trainer: artifact_value = wandb.Image(filename) elif artifact_type == "video": logger.debug(f"Saving video to {filename}") + # TODO: this should be configurable here as well as in validation runs where we call the pipeline that has `fps`. export_to_video(artifact_value, filename, fps=15) artifact_value = wandb.Video(filename, caption=prompt) @@ -849,9 +921,18 @@ class Trainer: all_artifacts = gather_object(all_processes_artifacts) if accelerator.is_main_process: + tracker_key = "final" if final_validation else "validation" for tracker in accelerator.trackers: if tracker.name == "wandb": - tracker.log({"validation": all_artifacts}, step=step) + image_artifacts = [artifact for artifact in all_artifacts if isinstance(artifact, wandb.Image)] + video_artifacts = [artifact for artifact in all_artifacts if isinstance(artifact, wandb.Video)] + tracker.log( + { + tracker_key: {"images": image_artifacts, "videos": video_artifacts}, + }, + step=step, + ) + # Remove all hooks that might have been added during pipeline initialization to the models pipeline.remove_all_hooks() @@ -864,11 +945,11 @@ class Trainer: logger.info(f"Memory after validation end: {json.dumps(memory_statistics, indent=4)}") torch.cuda.reset_peak_memory_stats(accelerator.device) - self.transformer.train() + if not final_validation: + self.transformer.train() def evaluate(self) -> None: - logger.info("Starting evaluation") - # TODO: implement metrics for evaluation + raise NotImplementedError def _init_distributed(self) -> None: logging_dir = Path(self.args.output_dir, self.args.logging_dir) @@ -922,7 +1003,7 @@ class Trainer: if self.args.push_to_hub: repo_id = self.args.hub_model_id or Path(self.args.output_dir).name - self.state.repo_id = create_repo(token=self.args.hub_token, name=repo_id).repo_id + self.state.repo_id = create_repo(token=self.args.hub_token, repo_id=repo_id).repo_id def _move_components_to_device(self): if self.text_encoder is not None: @@ -937,3 +1018,24 @@ class Trainer: self.unet = self.unet.to(self.state.accelerator.device) if self.vae is not None: self.vae = self.vae.to(self.state.accelerator.device) + + def _get_training_dtype(self, accelerator) -> torch.dtype: + weight_dtype = torch.float32 + if accelerator.state.deepspeed_plugin: + # DeepSpeed is handling precision, use what's in the DeepSpeed config + if ( + "fp16" in accelerator.state.deepspeed_plugin.deepspeed_config + and accelerator.state.deepspeed_plugin.deepspeed_config["fp16"]["enabled"] + ): + weight_dtype = torch.float16 + if ( + "bf16" in accelerator.state.deepspeed_plugin.deepspeed_config + and accelerator.state.deepspeed_plugin.deepspeed_config["bf16"]["enabled"] + ): + weight_dtype = torch.bfloat16 + else: + if self.state.accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + elif self.state.accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + return weight_dtype diff --git a/finetrainers/utils/checkpointing.py b/finetrainers/utils/checkpointing.py new file mode 100644 index 0000000..ba7f3d6 --- /dev/null +++ b/finetrainers/utils/checkpointing.py @@ -0,0 +1,61 @@ +import os +from typing import Tuple +from accelerate.logging import get_logger +from ..constants import FINETRAINERS_LOG_LEVEL +from ..utils.file_utils import find_files, delete_files + +logger = get_logger("finetrainers") +logger.setLevel(FINETRAINERS_LOG_LEVEL) + + +def get_latest_ckpt_path_to_resume_from( + resume_from_checkpoint: str, num_update_steps_per_epoch: int, output_dir: str +) -> Tuple[str, int, int, int]: + if not resume_from_checkpoint: + initial_global_step = 0 + global_step = 0 + first_epoch = 0 + resume_from_checkpoint_path = None + else: + if resume_from_checkpoint != "latest": + path = os.path.basename(resume_from_checkpoint) + else: + # Get the most recent checkpoint + dirs = os.listdir(output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + logger.info(f"Checkpoint '{resume_from_checkpoint}' does not exist. Starting a new training run.") + resume_from_checkpoint = None + initial_global_step = 0 + global_step = 0 + first_epoch = 0 + resume_from_checkpoint_path = None + else: + logger.info(f"Resuming from checkpoint {path}") + resume_from_checkpoint_path = os.path.join(output_dir, path) + global_step = int(path.split("-")[1]) + + initial_global_step = global_step + first_epoch = global_step // num_update_steps_per_epoch + + return resume_from_checkpoint_path, initial_global_step, global_step, first_epoch + + +def get_intermediate_ckpt_path(checkpointing_limit: int, step: int, output_dir: str) -> str: + # before saving state, check if this save would set us over the `checkpointing_limit` + if checkpointing_limit is not None: + checkpoints = find_files(output_dir, prefix="checkpoint") + + # before we save the new checkpoint, we need to have at_most `checkpoints_total_limit - 1` checkpoints + if len(checkpoints) >= checkpointing_limit: + num_to_remove = len(checkpoints) - checkpointing_limit + 1 + checkpoints_to_remove = checkpoints[0:num_to_remove] + delete_files(checkpoints_to_remove) + + logger.info(f"Checkpointing at step {step}") + save_path = os.path.join(output_dir, f"checkpoint-{step}") + logger.info(f"Saving state to {save_path}") + return save_path diff --git a/finetrainers/utils/optimizer_utils.py b/finetrainers/utils/optimizer_utils.py index d05d5d3..77842eb 100644 --- a/finetrainers/utils/optimizer_utils.py +++ b/finetrainers/utils/optimizer_utils.py @@ -1,5 +1,4 @@ import inspect -import logging from accelerate.logging import get_logger import torch diff --git a/finetrainers/utils/torch_utils.py b/finetrainers/utils/torch_utils.py index 2aad077..17c5f9c 100644 --- a/finetrainers/utils/torch_utils.py +++ b/finetrainers/utils/torch_utils.py @@ -27,3 +27,9 @@ def align_device_and_dtype( if dtype is not None: x = {k: align_device_and_dtype(v, device, dtype) for k, v in x.items()} return x + + +def expand_tensor_to_dims(tensor, ndim): + while len(tensor.shape) < ndim: + tensor = tensor.unsqueeze(-1) + return tensor \ No newline at end of file diff --git a/train.py b/train.py index e12d47c..088c061 100644 --- a/train.py +++ b/train.py @@ -33,7 +33,7 @@ def main(): trainer.prepare_for_training() trainer.prepare_trackers() trainer.train() - trainer.evaluate() + # trainer.evaluate() except KeyboardInterrupt: logger.info("Received keyboard interrupt. Exiting...")