From c8d961dedfb111ee534bbe449dc281684d8f652c Mon Sep 17 00:00:00 2001 From: Xander Steenbrugge Date: Thu, 8 Aug 2024 18:57:20 +0200 Subject: [PATCH] optionally set unet_optimizer to none --- main.py | 23 +++++++++++++---------- 1 file changed, 13 insertions(+), 10 deletions(-) diff --git a/main.py b/main.py index 09c8e6e..5451d8a 100755 --- a/main.py +++ b/main.py @@ -148,15 +148,18 @@ def train(config: TrainingConfig): pipe=pipe ) - optimizer_unet = get_unet_optimizer( - prodigy_d_coef=config.prodigy_d_coef, - prodigy_growth_factor=config.unet_prodigy_growth_factor, - lora_weight_decay=config.lora_weight_decay, - use_dora=config.use_dora, - unet_trainable_params=unet_trainable_params, - optimizer_name=config.unet_optimizer_type - ) - + if config.unet_lr > 0.0: + optimizer_unet = get_unet_optimizer( + prodigy_d_coef=config.prodigy_d_coef, + prodigy_growth_factor=config.unet_prodigy_growth_factor, + lora_weight_decay=config.lora_weight_decay, + use_dora=config.use_dora, + unet_trainable_params=unet_trainable_params, + optimizer_name=config.unet_optimizer_type + ) + else: + optimizer_unet = None + print_trainable_parameters(unet, model_name = 'unet') for i, text_encoder in enumerate(text_encoders): if text_encoder is not None: @@ -548,4 +551,4 @@ if __name__ == "__main__": for progress in train(config=config): print(f"Progress: {(100*progress):.2f}%", end="\r") - print("Training done :)") \ No newline at end of file + print("Training done :)")