From a8d2d6ec46a7e9986f12822a7db4ea2918e25e74 Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Thu, 16 Nov 2023 19:40:53 +0100 Subject: [PATCH] Move lora related code around, remove clip stuff --- __init__.py | 3 --- coreml_suite/converter.py | 39 ++++++++++++++++++++--------------- coreml_suite/lcm/converter.py | 8 +++++++ coreml_suite/lora.py | 2 +- coreml_suite/nodes.py | 13 ++++++------ 5 files changed, 38 insertions(+), 27 deletions(-) diff --git a/__init__.py b/__init__.py index aadaeac..7fde9f4 100644 --- a/__init__.py +++ b/__init__.py @@ -8,7 +8,6 @@ from coreml_suite.nodes import ( CoreMLSampler, CoreMLModelAdapter, COREML_CONVERT, - COREML_LOAD_CLIP, ) from coreml_suite.lcm import ( COREML_CONVERT_LCM, @@ -16,7 +15,6 @@ from coreml_suite.lcm import ( NODE_CLASS_MAPPINGS = { "CoreMLUNetLoader": CoreMLLoaderUNet, - "Core ML CLIP Loader": COREML_LOAD_CLIP, "CoreMLSampler": CoreMLSampler, "CoreMLModelAdapter": CoreMLModelAdapter, "Core ML Converter": COREML_CONVERT, @@ -25,7 +23,6 @@ NODE_CLASS_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = { "CoreMLUNetLoader": "Load Core ML UNet", "CoreMLSampler": "Core ML Sampler", - "Core ML CLIP Loader": "Load Core ML CLIP", "CoreMLModelAdapter": "Core ML Adapter (Experimental)", "Core ML Converter": "Convert Checkpoint to Core ML", "Core ML LCM Converter": "Convert LCM to Core ML", diff --git a/coreml_suite/converter.py b/coreml_suite/converter.py index 30593ae..1d196a2 100644 --- a/coreml_suite/converter.py +++ b/coreml_suite/converter.py @@ -1,4 +1,3 @@ -import abc import gc import os import shutil @@ -9,9 +8,7 @@ import coremltools as ct import numpy as np import torch from diffusers import StableDiffusionPipeline, LatentConsistencyModelPipeline -from python_coreml_stable_diffusion.coreml_model import CoreMLModel from python_coreml_stable_diffusion.unet import UNet2DConditionModel -from torch import nn from coreml_suite.lcm.unet import UNet2DConditionModelLCM from coreml_suite.logger import logger @@ -23,6 +20,10 @@ class ModelType(Enum): LCM = auto() +class StableDiffusionLCMPipeline(LatentConsistencyModelPipeline): + pass + + MODEL_TYPE_TO_UNET_CLS = { ModelType.SD15: UNet2DConditionModel, ModelType.LCM: UNet2DConditionModelLCM, @@ -30,7 +31,7 @@ MODEL_TYPE_TO_UNET_CLS = { MODEL_TYPE_TO_PIPE_CLS = { ModelType.SD15: StableDiffusionPipeline, - ModelType.LCM: LatentConsistencyModelPipeline, + ModelType.LCM: StableDiffusionLCMPipeline, } @@ -266,7 +267,6 @@ def convert_unet( def convert( - model_type: ModelType, ckpt_path: str, unet_out_path: str, batch_size: int = 1, @@ -274,22 +274,27 @@ def convert( controlnet_support: bool = False, lora_paths: list[str | os.PathLike] = None, ): + if os.path.exists(unet_out_path): + logger.info(f"Found existing model at {unet_out_path}! Skipping..") + return + + model_type = ModelType.SD15 + pipe_cls = MODEL_TYPE_TO_PIPE_CLS[model_type] ref_pipe = pipe_cls.from_single_file(ckpt_path) - 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() + for lora_path in lora_paths: + ref_pipe.load_lora_weights(lora_path) + ref_pipe.fuse_lora() - convert_unet( - ref_pipe, - model_type, - unet_out_path, - batch_size, - sample_size, - controlnet_support, - ) + convert_unet( + ref_pipe, + model_type, + unet_out_path, + batch_size, + sample_size, + controlnet_support, + ) def compile_model(out_path, out_name, submodule_name): diff --git a/coreml_suite/lcm/converter.py b/coreml_suite/lcm/converter.py index eb60c99..e79e761 100644 --- a/coreml_suite/lcm/converter.py +++ b/coreml_suite/lcm/converter.py @@ -7,6 +7,7 @@ import gc import numpy as np import torch from diffusers import UNet2DConditionModel, LCMScheduler +from diffusers.loaders import LoraLoaderMixin from comfy.model_management import get_torch_device from coreml_suite.lcm.unet import UNet2DConditionModelLCM @@ -219,9 +220,16 @@ def convert( batch_size: int = 1, sample_size: tuple[int, int] = (64, 64), controlnet_support: bool = False, + lora_paths: list[str] = None, ): + lora_paths = lora_paths or [] coreml_unet, ref_unet = get_unets() + for lora_path in lora_paths: + lora_sd, network_alphas = LoraLoaderMixin.lora_state_dict(lora_path) + LoraLoaderMixin.load_lora_into_unet(lora_sd, network_alphas, ref_unet) + ref_unet.fuse_lora() + sample_shape = ( batch_size, # B ref_unet.config.in_channels, # C diff --git a/coreml_suite/lora.py b/coreml_suite/lora.py index d396a61..8eeba51 100644 --- a/coreml_suite/lora.py +++ b/coreml_suite/lora.py @@ -24,7 +24,7 @@ def load_lora(lora_params, ckpt_name): else: lora_path = folder_paths.get_full_path("loras", lora_name) - lora_clip = sd.load_lora_for_models( + _, lora_clip = sd.load_lora_for_models( None, clip, utils.load_torch_file(lora_path), strength_model, strength_clip ) diff --git a/coreml_suite/nodes.py b/coreml_suite/nodes.py index a19ebcb..36de5f4 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.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 @@ -194,7 +193,6 @@ class COREML_CONVERT(COREML_NODE): return { "required": { "ckpt_name": (folder_paths.get_filename_list("checkpoints"),), - "model_type": (["SD15", "LCM"], {"default": "SD15"}), "height": ("INT", {"default": 512, "min": 512, "max": 2048, "step": 8}), "width": ("INT", {"default": 512, "min": 512, "max": 2048, "step": 8}), "batch_size": ("INT", {"default": 1, "min": 1, "max": 64}), @@ -220,7 +218,6 @@ class COREML_CONVERT(COREML_NODE): def convert( self, ckpt_name, - model_type, height, width, batch_size, @@ -247,13 +244,18 @@ class COREML_CONVERT(COREML_NODE): sample_size = (h // 8, w // 8) batch_size = batch_size cn_support_str = "_cn" if controlnet_support else "" - lcm_str = "_lcm" if model_type == "LCM" else "" + lora_str = ( + "_" + "_".join(lora_param[0].split(".")[0] for lora_param in lora_stack) + if lora_stack + else "" + ) out_name = ( - f"{ckpt_name.split('.')[0]}_{batch_size}x{w}x{h}{cn_support_str}{lcm_str}" + f"{ckpt_name.split('.')[0]}{lora_str}_{batch_size}x{w}x{h}{cn_support_str}" ) unet_out_path = converter.get_out_path("unet", f"{out_name}") + unet_out_path = unet_out_path.replace(" ", "_") ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) @@ -263,7 +265,6 @@ class COREML_CONVERT(COREML_NODE): ] converter.convert( - model_type=ModelType[model_type], ckpt_path=ckpt_path, unet_out_path=unet_out_path, sample_size=sample_size,