Move lora related code around, remove clip stuff

This commit is contained in:
aszc-dev
2024-06-28 15:52:54 +02:00
parent 7643211d8d
commit 4c9195bbc0
5 changed files with 38 additions and 27 deletions
-3
View File
@@ -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",
+22 -17
View File
@@ -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):
+8
View File
@@ -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
+1 -1
View File
@@ -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
)
+7 -6
View File
@@ -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,