From 284bb6f112aa7358d6ff7073efb676f35638acd9 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 8 Sep 2024 16:20:40 +0300 Subject: [PATCH] add lora path to load network_weights in the model select node --- nodes.py | 21 ++++++++++++++------- 1 file changed, 14 insertions(+), 7 deletions(-) diff --git a/nodes.py b/nodes.py index bb652cc..e57ff56 100644 --- a/nodes.py +++ b/nodes.py @@ -37,11 +37,14 @@ class FluxTrainModelSelect: @classmethod def INPUT_TYPES(s): return {"required": { - "transformer": (folder_paths.get_filename_list("unet"), ), - "vae": (folder_paths.get_filename_list("vae"), ), - "clip_l": (folder_paths.get_filename_list("clip"), ), - "t5": (folder_paths.get_filename_list("clip"), ), - }, + "transformer": (folder_paths.get_filename_list("unet"), ), + "vae": (folder_paths.get_filename_list("vae"), ), + "clip_l": (folder_paths.get_filename_list("clip"), ), + "t5": (folder_paths.get_filename_list("clip"), ), + }, + "optional": { + "lora_path": ("STRING",{"multiline": True, "forceInput": True, "default": "", "tooltip": "pre-trained LoRA path to load (network_weights)"}), + } } RETURN_TYPES = ("TRAIN_FLUX_MODELS",) @@ -49,7 +52,7 @@ class FluxTrainModelSelect: FUNCTION = "loadmodel" CATEGORY = "FluxTrainer" - def loadmodel(self, transformer, vae, clip_l, t5): + def loadmodel(self, transformer, vae, clip_l, t5, lora_path=None): transformer_path = folder_paths.get_full_path("unet", transformer) vae_path = folder_paths.get_full_path("vae", vae) @@ -60,7 +63,8 @@ class FluxTrainModelSelect: "transformer": transformer_path, "vae": vae_path, "clip_l": clip_path, - "t5": t5_path + "t5": t5_path, + "lora_path": lora_path if lora_path is not None else None } return (flux_models,) @@ -446,6 +450,9 @@ class InitFluxLoRATraining: else: config_dict["gradient_checkpointing"] = True + if flux_models["lora_path"]: + config_dict["network_weights"] = flux_models["lora_path"] + config_dict.update(kwargs) config_dict.update(optimizer_settings)