Move lora related code around, remove clip stuff
This commit is contained in:
@@ -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
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user