diff --git a/coreml_suite/lora.py b/coreml_suite/lora.py new file mode 100644 index 0000000..d396a61 --- /dev/null +++ b/coreml_suite/lora.py @@ -0,0 +1,44 @@ +import os + +import folder_paths +from comfy import sd, utils + + +def load_lora(lora_params, ckpt_name): + lora_params = ( + lora_params.copy() + if isinstance(lora_params, (list, dict, set)) + else lora_params + ) + ckpt_name = ( + ckpt_name.copy() if isinstance(ckpt_name, (list, dict, set)) else ckpt_name + ) + + def recursive_load_lora(lora_params, clip): + if len(lora_params) == 0: + return clip + + lora_name, strength_model, strength_clip = lora_params[0] + if os.path.isabs(lora_name): + lora_path = lora_name + else: + lora_path = folder_paths.get_full_path("loras", lora_name) + + 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_clip) + + ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) + _, clip, _, _ = sd.load_checkpoint_guess_config( + ckpt_path, + output_vae=False, + output_clip=True, + embedding_directory=folder_paths.get_folder_paths("embeddings"), + ) + + lora_clip = recursive_load_lora(lora_params, clip) + + return lora_clip diff --git a/coreml_suite/nodes.py b/coreml_suite/nodes.py index d2f6db7..a19ebcb 100644 --- a/coreml_suite/nodes.py +++ b/coreml_suite/nodes.py @@ -5,7 +5,7 @@ from coremltools import ComputeUnit from python_coreml_stable_diffusion.coreml_model import CoreMLModel import folder_paths -from comfy import model_base, sd1_clip, sd, utils +from comfy import model_base from comfy.model_management import get_torch_device from comfy.model_patcher import ModelPatcher from coreml_suite import COREML_NODE @@ -14,6 +14,7 @@ from coreml_suite import converter 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 +from coreml_suite.lora import load_lora from nodes import KSampler from coreml_suite.models import CoreMLModelWrapper @@ -164,7 +165,6 @@ class CoreMLModelAdapter: return (model_patcher,) -<<<<<<< HEAD class COREML_LOAD_CLIP(CoreMLLoader): PACKAGE_DIRNAME = "clip" RETURN_TYPES = ("CLIP",) @@ -278,43 +278,3 @@ class COREML_CONVERT(COREML_NODE): clip = load_lora(lora_stack, ckpt_name) return (CoreMLModel(unet_target_path, compute_unit, "compiled"), clip) - - -def load_lora(lora_params, ckpt_name): - lora_params = ( - lora_params.copy() - if isinstance(lora_params, (list, dict, set)) - else lora_params - ) - ckpt_name = ( - ckpt_name.copy() if isinstance(ckpt_name, (list, dict, set)) else ckpt_name - ) - - def recursive_load_lora(lora_params, clip): - if len(lora_params) == 0: - return clip - - lora_name, strength_model, strength_clip = lora_params[0] - if os.path.isabs(lora_name): - lora_path = lora_name - else: - lora_path = folder_paths.get_full_path("loras", lora_name) - - 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_clip) - - ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) - _, clip, _, _ = sd.load_checkpoint_guess_config( - ckpt_path, - output_vae=False, - output_clip=True, - embedding_directory=folder_paths.get_folder_paths("embeddings"), - ) - - lora_clip = recursive_load_lora(lora_params, clip) - - return lora_clip