diff --git a/coreml_suite/converter.py b/coreml_suite/converter.py index 52c8c7a..fdc0150 100644 --- a/coreml_suite/converter.py +++ b/coreml_suite/converter.py @@ -276,7 +276,7 @@ def convert( batch_size: int = 1, sample_size: tuple[int, int] = (64, 64), controlnet_support: bool = False, - lora_paths: list[str | os.PathLike] = None, + lora_weights: list[tuple[str | os.PathLike, float]] = None, attn_impl: str = AttentionImplementations.SPLIT_EINSUM.name, ): if os.path.exists(unet_out_path): @@ -292,8 +292,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) + for i, lora_weight in enumerate(lora_weights or []): + lora_path, strength = lora_weight + adapter_name = f"lora_{i}" + ref_pipe.load_lora_weights(lora_path, adapter_name=adapter_name) + ref_pipe.set_adapters([adapter_name], adapter_weights=[strength]) ref_pipe.fuse_lora() convert_unet( diff --git a/coreml_suite/nodes.py b/coreml_suite/nodes.py index aab7411..1bb6b8f 100644 --- a/coreml_suite/nodes.py +++ b/coreml_suite/nodes.py @@ -226,8 +226,9 @@ class COREML_CONVERT(COREML_NODE): can be loaded with the "LCMCoreMLLoaderUNet" node. """ lora_params = lora_params or {} - lora_params = [(k, v[0], v[1]) for k, v in lora_params.items()] + lora_params = [(k, v[0]) for k, v in lora_params.items()] lora_params = sorted(lora_params, key=lambda lora: lora[0]) + lora_weights = [(self.lora_path(lora[0]), lora[1]) for lora in lora_params] h = height w = width @@ -264,7 +265,6 @@ class COREML_CONVERT(COREML_NODE): unet_out_path = converter.get_out_path("unet", f"{out_name}") ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) - lora_params = lora_params or [] lora_paths = [ folder_paths.get_full_path("loras", lora[0]) for lora in lora_params ] @@ -275,7 +275,7 @@ class COREML_CONVERT(COREML_NODE): sample_size=sample_size, batch_size=batch_size, controlnet_support=controlnet_support, - lora_paths=lora_paths, + lora_weights=lora_weights, attn_impl=attention_implementation, ) unet_target_path = converter.compile_model( @@ -284,6 +284,10 @@ class COREML_CONVERT(COREML_NODE): return (CoreMLModel(unet_target_path, compute_unit, "compiled"),) + @staticmethod + def lora_path(lora_name): + return folder_paths.get_full_path("loras", lora_name) + class COREML_LOAD_LORA(COREML_NODE, LoraLoader): @classmethod