From 1f304b36d4c4ccca785b77414ffe5aee0e615600 Mon Sep 17 00:00:00 2001 From: Aryan Date: Sat, 11 Jan 2025 01:28:47 +0100 Subject: [PATCH] handle saving of final model weights correctly --- finetrainers/trainer.py | 22 ++++++++++++---------- 1 file changed, 12 insertions(+), 10 deletions(-) diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index 6af4c2c..43d0dae 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -385,10 +385,10 @@ class Trainer: self._disable_grad_for_components(components) if self.args.training_type == "full-finetune": - logger.info("Finetuning transformer with no additional parameters.") + logger.info("Finetuning transformer with no additional parameters") self._enable_grad_for_components(components=[self.transformer]) else: - logger.info("Finetuning transformer with low-rank peft parameters.") + logger.info("Finetuning transformer with PEFT parameters") self._disable_grad_for_components(components=[self.transformer]) # For mixed precision training we cast all non-trainable weights (vae, text_encoder and transformer) to half-precision @@ -452,7 +452,6 @@ class Trainer: if weights: weights.pop() - # TODO: refactor later if needed. But for now, this is just a few LoC. if self.args.training_type == "lora": self.model_config["pipeline_cls"].save_lora_weights( output_dir, @@ -894,14 +893,17 @@ 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) - transformer_lora_layers = get_peft_model_state_dict(self.transformer) + transformer = unwrap_model(accelerator, self.transformer) - self.model_config["pipeline_cls"].save_lora_weights( - save_directory=self.args.output_dir, - transformer_lora_layers=transformer_lora_layers, - ) + if self.args.training_type == "lora": + transformer_lora_layers = get_peft_model_state_dict(transformer) + + self.model_config["pipeline_cls"].save_lora_weights( + save_directory=self.args.output_dir, + transformer_lora_layers=transformer_lora_layers, + ) + else: + transformer.save_pretrained(os.path.join(self.args.output_dir, "transformer")) self.validate(step=global_step, final_validation=True)