Setting LoRA model weights works

This commit is contained in:
aszc-dev
2024-06-28 15:52:54 +02:00
parent a3cf825d79
commit 5de5722474
2 changed files with 13 additions and 6 deletions
+6 -3
View File
@@ -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(
+7 -3
View File
@@ -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