diff --git a/coreml_suite/converter.py b/coreml_suite/converter.py index fdc0150..92c9415 100644 --- a/coreml_suite/converter.py +++ b/coreml_suite/converter.py @@ -278,6 +278,7 @@ def convert( controlnet_support: bool = False, lora_weights: list[tuple[str | os.PathLike, float]] = None, attn_impl: str = AttentionImplementations.SPLIT_EINSUM.name, + config_path: str = None, ): if os.path.exists(unet_out_path): logger.info(f"Found existing model at {unet_out_path}! Skipping..") @@ -290,7 +291,7 @@ def convert( model_type = ModelType.SD15 pipe_cls = MODEL_TYPE_TO_PIPE_CLS[model_type] - ref_pipe = pipe_cls.from_single_file(ckpt_path) + ref_pipe = pipe_cls.from_single_file(ckpt_path, original_config_file=config_path) for i, lora_weight in enumerate(lora_weights or []): lora_path, strength = lora_weight diff --git a/coreml_suite/nodes.py b/coreml_suite/nodes.py index 1bb6b8f..80988d1 100644 --- a/coreml_suite/nodes.py +++ b/coreml_suite/nodes.py @@ -265,9 +265,10 @@ 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_paths = [ - folder_paths.get_full_path("loras", lora[0]) for lora in lora_params - ] + config_filename = ckpt_name.split(".")[0] + ".yaml" + config_path = folder_paths.get_full_path("configs", config_filename) + if config_path: + logger.info(f"Using config file {config_path}") converter.convert( ckpt_path=ckpt_path, @@ -277,6 +278,7 @@ class COREML_CONVERT(COREML_NODE): controlnet_support=controlnet_support, lora_weights=lora_weights, attn_impl=attention_implementation, + config_path=config_path, ) unet_target_path = converter.compile_model( out_path=unet_out_path, out_name=out_name, submodule_name="unet"