diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index 0f3e5b9..04f5737 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -48,7 +48,7 @@ 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.checkpointing import sort_out_and_load_latest_ckpt_states, save_intermediate_ckpt_states +from .utils.checkpointing import get_latest_ckpt_path_to_resume_from, get_intermediate_ckpt_path logger = get_logger("finetrainers") @@ -561,12 +561,13 @@ class Trainer: initial_global_step = 0 # Potentially load in the weights and states from a previous save - initial_global_step, global_step, first_epoch = sort_out_and_load_latest_ckpt_states( - accelerator=self.state.accelerator, + 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), @@ -699,9 +700,10 @@ class Trainer: # Checkpointing if accelerator.distributed_type == DistributedType.DEEPSPEED or accelerator.is_main_process: if global_step % self.args.checkpointing_steps == 0: - save_intermediate_ckpt_states( - accelerator=accelerator, checkpointing_limit=self.args.checkpointing_limit, step=global_step, output_dir=self.args.output_dir + 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) # Maybe run validation should_run_validation = ( @@ -748,12 +750,11 @@ class Trainer: transformer_lora_layers=transformer_lora_layers, ) + self.validate(step=global_step, final_validation=True) if self.args.push_to_hub: upload_folder( repo_id=self.state.repo_id, folder_path=self.args.output_dir, ignore_patterns=["checkpoint-*"] ) - - self.validate(step=global_step, final_validation=True) del self.tokenizer, self.text_encoder, self.transformer, self.vae, self.scheduler free_memory() @@ -795,6 +796,7 @@ class Trainer: ) 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, diff --git a/finetrainers/utils/checkpointing.py b/finetrainers/utils/checkpointing.py index 1a69a1f..11b94fc 100644 --- a/finetrainers/utils/checkpointing.py +++ b/finetrainers/utils/checkpointing.py @@ -1,4 +1,5 @@ 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 @@ -6,13 +7,14 @@ from ..utils.file_utils import find_files, delete_files logger = get_logger("finetrainers") logger.setLevel(FINETRAINERS_LOG_LEVEL) -def sort_out_and_load_latest_ckpt_states( - accelerator, resume_from_checkpoint, num_update_steps_per_epoch, output_dir - ): +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) @@ -29,18 +31,19 @@ def sort_out_and_load_latest_ckpt_states( ) resume_from_checkpoint = None initial_global_step = 0 + resume_from_checkpoint_path = None else: logger.info(f"Resuming from checkpoint {path}") - accelerator.load_state(os.path.join(output_dir, 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 initial_global_step, global_step, first_epoch + return resume_from_checkpoint_path, initial_global_step, global_step, first_epoch -def save_intermediate_ckpt_states(accelerator, checkpointing_limit, step, output_dir): +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") @@ -53,5 +56,5 @@ def save_intermediate_ckpt_states(accelerator, checkpointing_limit, step, output logger.info(f"Checkpointing at step {step}") save_path = os.path.join(output_dir, f"checkpoint-{step}") - accelerator.save_state(save_path) - logger.info(f"Saved state to {save_path}") \ No newline at end of file + logger.info(f"Saving state to {save_path}") + return save_path \ No newline at end of file