From 6907d4910f5e1723c8ff8deddec15075047e1021 Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Thu, 16 Nov 2023 00:08:19 +0100 Subject: [PATCH] Remove ckpt loading when loading lora clip --- coreml_suite/converter.py | 8 ++++---- coreml_suite/nodes.py | 20 ++++++++++---------- 2 files changed, 14 insertions(+), 14 deletions(-) diff --git a/coreml_suite/converter.py b/coreml_suite/converter.py index 5f39563..30593ae 100644 --- a/coreml_suite/converter.py +++ b/coreml_suite/converter.py @@ -277,11 +277,11 @@ def convert( pipe_cls = MODEL_TYPE_TO_PIPE_CLS[model_type] ref_pipe = pipe_cls.from_single_file(ckpt_path) - for lora_path in lora_paths: - ref_pipe.load_lora_weights(lora_path) - ref_pipe.fuse_lora() - if not os.path.exists(unet_out_path): + for lora_path in lora_paths: + ref_pipe.load_lora_weights(lora_path) + ref_pipe.fuse_lora() + convert_unet( ref_pipe, model_type, diff --git a/coreml_suite/nodes.py b/coreml_suite/nodes.py index 0db554e..d2f6db7 100644 --- a/coreml_suite/nodes.py +++ b/coreml_suite/nodes.py @@ -11,7 +11,6 @@ from comfy.model_patcher import ModelPatcher from coreml_suite import COREML_NODE from comfy.sd import CLIP from coreml_suite import converter -from coreml_suite.clip import CoreMLCLIP, SDClipModelCoreML from coreml_suite.converter import ModelType from coreml_suite.lcm.utils import add_lcm_model_options, lcm_patch, is_lcm from coreml_suite.logger import logger @@ -165,6 +164,7 @@ class CoreMLModelAdapter: return (model_patcher,) +<<<<<<< HEAD class COREML_LOAD_CLIP(CoreMLLoader): PACKAGE_DIRNAME = "clip" RETURN_TYPES = ("CLIP",) @@ -275,7 +275,7 @@ class COREML_CONVERT(COREML_NODE): out_path=unet_out_path, out_name=out_name, submodule_name="unet" ) - _, clip = load_lora(lora_stack, ckpt_name) + clip = load_lora(lora_stack, ckpt_name) return (CoreMLModel(unet_target_path, compute_unit, "compiled"), clip) @@ -290,9 +290,9 @@ def load_lora(lora_params, ckpt_name): ckpt_name.copy() if isinstance(ckpt_name, (list, dict, set)) else ckpt_name ) - def recursive_load_lora(lora_params, ckpt, clip): + def recursive_load_lora(lora_params, clip): if len(lora_params) == 0: - return ckpt, clip + return clip lora_name, strength_model, strength_clip = lora_params[0] if os.path.isabs(lora_name): @@ -300,21 +300,21 @@ def load_lora(lora_params, ckpt_name): else: lora_path = folder_paths.get_full_path("loras", lora_name) - lora_model, lora_clip = sd.load_lora_for_models( - ckpt, clip, utils.load_torch_file(lora_path), strength_model, strength_clip + lora_clip = sd.load_lora_for_models( + None, clip, utils.load_torch_file(lora_path), strength_model, strength_clip ) # Call the function again with the new lora_model and lora_clip and the remaining tuples - return recursive_load_lora(lora_params[1:], lora_model, lora_clip) + return recursive_load_lora(lora_params[1:], lora_clip) ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) - ckpt, clip, _, _ = sd.load_checkpoint_guess_config( + _, clip, _, _ = sd.load_checkpoint_guess_config( ckpt_path, output_vae=False, output_clip=True, embedding_directory=folder_paths.get_folder_paths("embeddings"), ) - lora_model, lora_clip = recursive_load_lora(lora_params, ckpt, clip) + lora_clip = recursive_load_lora(lora_params, clip) - return lora_model, lora_clip + return lora_clip