Basic conversion + LoRA support works
This commit is contained in:
+11
-1
@@ -3,20 +3,30 @@ import sys
|
||||
|
||||
sys.path.append(os.path.dirname(__file__))
|
||||
|
||||
from coreml_suite.nodes import CoreMLLoaderUNet, CoreMLSampler, CoreMLModelAdapter
|
||||
from coreml_suite.nodes import (
|
||||
CoreMLLoaderUNet,
|
||||
CoreMLSampler,
|
||||
CoreMLModelAdapter,
|
||||
COREML_CONVERT,
|
||||
COREML_LOAD_CLIP,
|
||||
)
|
||||
from coreml_suite.lcm import (
|
||||
COREML_CONVERT_LCM,
|
||||
)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"CoreMLUNetLoader": CoreMLLoaderUNet,
|
||||
"Core ML CLIP Loader": COREML_LOAD_CLIP,
|
||||
"CoreMLSampler": CoreMLSampler,
|
||||
"CoreMLModelAdapter": CoreMLModelAdapter,
|
||||
"Core ML Converter": COREML_CONVERT,
|
||||
"Core ML LCM Converter": COREML_CONVERT_LCM,
|
||||
}
|
||||
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",
|
||||
}
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from comfy.sd import CLIP
|
||||
from comfy.sd1_clip import SDClipModel
|
||||
|
||||
|
||||
class CoreMLCLIP(CLIP):
|
||||
pass
|
||||
|
||||
|
||||
class SDClipModelCoreML(SDClipModel):
|
||||
def __init__(self, **kwargs):
|
||||
self.model = kwargs.pop("coreml_model")
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def encode_token_weights(self, token_weight_pairs):
|
||||
tokens = np.array(token_weight_pairs["l"]).astype(np.float16)[:, :, 0]
|
||||
encoded = self.model(input_ids=tokens)
|
||||
cond = torch.from_numpy(encoded["last_hidden_state"])
|
||||
pooled = torch.from_numpy(encoded["pooled_outputs"])
|
||||
return cond, pooled
|
||||
@@ -0,0 +1,301 @@
|
||||
import abc
|
||||
import gc
|
||||
import os
|
||||
import shutil
|
||||
import time
|
||||
from enum import Enum, auto
|
||||
|
||||
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
|
||||
from folder_paths import get_folder_paths
|
||||
|
||||
|
||||
class ModelType(Enum):
|
||||
SD15 = auto()
|
||||
LCM = auto()
|
||||
|
||||
|
||||
MODEL_TYPE_TO_UNET_CLS = {
|
||||
ModelType.SD15: UNet2DConditionModel,
|
||||
ModelType.LCM: UNet2DConditionModelLCM,
|
||||
}
|
||||
|
||||
MODEL_TYPE_TO_PIPE_CLS = {
|
||||
ModelType.SD15: StableDiffusionPipeline,
|
||||
ModelType.LCM: LatentConsistencyModelPipeline,
|
||||
}
|
||||
|
||||
|
||||
def get_unet(model_type: ModelType, ref_pipe):
|
||||
ref_unet = ref_pipe.unet
|
||||
|
||||
unet_cls = MODEL_TYPE_TO_UNET_CLS[model_type]
|
||||
cml_unet = unet_cls.from_config(ref_unet.config).eval()
|
||||
cml_unet.load_state_dict(ref_unet.state_dict(), strict=False)
|
||||
|
||||
return cml_unet
|
||||
|
||||
|
||||
def get_encoder_hidden_states_shape(ref_pipe, batch_size):
|
||||
text_encoder = ref_pipe.text_encoder
|
||||
|
||||
text_token_sequence_length = text_encoder.config.max_position_embeddings
|
||||
hidden_size = (text_encoder.config.hidden_size,)
|
||||
|
||||
encoder_hidden_states_shape = (
|
||||
batch_size,
|
||||
ref_pipe.unet.config.cross_attention_dim or hidden_size,
|
||||
1,
|
||||
text_token_sequence_length,
|
||||
)
|
||||
|
||||
return encoder_hidden_states_shape
|
||||
|
||||
|
||||
def get_coreml_inputs(sample_inputs):
|
||||
coreml_sample_unet_inputs = {
|
||||
k: v.numpy().astype(np.float16) for k, v in sample_inputs.items()
|
||||
}
|
||||
return [
|
||||
ct.TensorType(
|
||||
name=k,
|
||||
shape=v.shape,
|
||||
dtype=v.numpy().dtype if isinstance(v, torch.Tensor) else v.dtype,
|
||||
)
|
||||
for k, v in coreml_sample_unet_inputs.items()
|
||||
]
|
||||
|
||||
|
||||
def load_coreml_model(out_path):
|
||||
logger.info(f"Loading model from {out_path}")
|
||||
|
||||
start = time.time()
|
||||
coreml_model = ct.models.MLModel(out_path)
|
||||
logger.info(f"Loading {out_path} took {time.time() - start:.1f} seconds")
|
||||
|
||||
return coreml_model
|
||||
|
||||
|
||||
def convert_to_coreml(
|
||||
submodule_name, torchscript_module, sample_inputs, output_names, out_path
|
||||
):
|
||||
if os.path.exists(out_path):
|
||||
logger.info(f"Skipping export because {out_path} already exists")
|
||||
coreml_model = load_coreml_model(out_path)
|
||||
else:
|
||||
logger.info(f"Converting {submodule_name} to CoreML..")
|
||||
coreml_model = ct.convert(
|
||||
torchscript_module,
|
||||
convert_to="mlprogram",
|
||||
minimum_deployment_target=ct.target.macOS13,
|
||||
inputs=sample_inputs,
|
||||
outputs=[
|
||||
ct.TensorType(name=name, dtype=np.float32) for name in output_names
|
||||
],
|
||||
skip_model_load=True,
|
||||
)
|
||||
|
||||
del torchscript_module
|
||||
gc.collect()
|
||||
|
||||
return coreml_model
|
||||
|
||||
|
||||
def get_out_path(submodule_name, model_name):
|
||||
fname = f"{model_name}_{submodule_name}.mlpackage"
|
||||
unet_path = get_folder_paths(submodule_name)[0]
|
||||
out_path = os.path.join(unet_path, fname)
|
||||
return out_path
|
||||
|
||||
|
||||
def compile_coreml_model(source_model_path, output_dir, final_name):
|
||||
"""Compiles Core ML models using the coremlcompiler utility from Xcode toolchain"""
|
||||
target_path = os.path.join(output_dir, f"{final_name}.mlmodelc")
|
||||
if os.path.exists(target_path):
|
||||
logger.warning(f"Found existing compiled model at {target_path}! Skipping..")
|
||||
return target_path
|
||||
|
||||
logger.info(f"Compiling {source_model_path}")
|
||||
source_model_name = os.path.basename(os.path.splitext(source_model_path)[0])
|
||||
|
||||
os.system(f"xcrun coremlcompiler compile {source_model_path} {output_dir}")
|
||||
compiled_output = os.path.join(output_dir, f"{source_model_name}.mlmodelc")
|
||||
shutil.move(compiled_output, target_path)
|
||||
|
||||
return target_path
|
||||
|
||||
|
||||
def get_sample_input(batch_size, encoder_hidden_states_shape, sample_shape, scheduler):
|
||||
sample_unet_inputs = dict(
|
||||
[
|
||||
("sample", torch.rand(*sample_shape)),
|
||||
(
|
||||
"timestep",
|
||||
torch.tensor([scheduler.timesteps[0].item()] * batch_size).to(
|
||||
torch.float32
|
||||
),
|
||||
),
|
||||
("encoder_hidden_states", torch.rand(*encoder_hidden_states_shape)),
|
||||
]
|
||||
)
|
||||
return sample_unet_inputs
|
||||
|
||||
|
||||
def lcm_inputs(sample_unet_inputs):
|
||||
batch_size = sample_unet_inputs["sample"].shape[0]
|
||||
return {"timestep_cond": torch.randn(batch_size, 256).to(torch.float32)}
|
||||
|
||||
|
||||
def get_inputs_spec(inputs):
|
||||
inputs_spec = {k: (v.shape, v.dtype) for k, v in inputs.items()}
|
||||
return inputs_spec
|
||||
|
||||
|
||||
def add_cnet_support(sample_shape, reference_unet):
|
||||
from python_coreml_stable_diffusion.unet import calculate_conv2d_output_shape
|
||||
|
||||
additional_residuals_shapes = []
|
||||
|
||||
batch_size = sample_shape[0]
|
||||
h, w = sample_shape[2:]
|
||||
|
||||
# conv_in
|
||||
out_h, out_w = calculate_conv2d_output_shape(
|
||||
h,
|
||||
w,
|
||||
reference_unet.conv_in,
|
||||
)
|
||||
additional_residuals_shapes.append(
|
||||
(batch_size, reference_unet.conv_in.out_channels, out_h, out_w)
|
||||
)
|
||||
|
||||
# down_blocks
|
||||
for down_block in reference_unet.down_blocks:
|
||||
additional_residuals_shapes += [
|
||||
(batch_size, resnet.out_channels, out_h, out_w)
|
||||
for resnet in down_block.resnets
|
||||
]
|
||||
if hasattr(down_block, "downsamplers") and down_block.downsamplers is not None:
|
||||
for downsampler in down_block.downsamplers:
|
||||
out_h, out_w = calculate_conv2d_output_shape(
|
||||
out_h, out_w, downsampler.conv
|
||||
)
|
||||
additional_residuals_shapes.append(
|
||||
(
|
||||
batch_size,
|
||||
down_block.downsamplers[-1].conv.out_channels,
|
||||
out_h,
|
||||
out_w,
|
||||
)
|
||||
)
|
||||
|
||||
# mid_block
|
||||
additional_residuals_shapes.append(
|
||||
(batch_size, reference_unet.mid_block.resnets[-1].out_channels, out_h, out_w)
|
||||
)
|
||||
|
||||
additional_inputs = {}
|
||||
for i, shape in enumerate(additional_residuals_shapes):
|
||||
sample_residual_input = torch.rand(*shape)
|
||||
additional_inputs[f"additional_residual_{i}"] = sample_residual_input
|
||||
|
||||
return additional_inputs
|
||||
|
||||
|
||||
def convert_unet(
|
||||
ref_pipe,
|
||||
model_type: ModelType,
|
||||
unet_out_path: str,
|
||||
batch_size: int = 1,
|
||||
sample_size: tuple[int, int] = (64, 64),
|
||||
controlnet_support: bool = False,
|
||||
):
|
||||
coreml_unet = get_unet(model_type, ref_pipe)
|
||||
ref_unet = ref_pipe.unet
|
||||
|
||||
sample_shape = (
|
||||
batch_size, # B
|
||||
ref_unet.config.in_channels, # C
|
||||
sample_size[0], # H
|
||||
sample_size[1], # W
|
||||
)
|
||||
|
||||
encoder_hidden_states_shape = get_encoder_hidden_states_shape(ref_pipe, batch_size)
|
||||
|
||||
scheduler = ref_pipe.scheduler
|
||||
scheduler.set_timesteps(50)
|
||||
|
||||
sample_inputs = get_sample_input(
|
||||
batch_size, encoder_hidden_states_shape, sample_shape, scheduler
|
||||
)
|
||||
|
||||
if model_type == ModelType.LCM:
|
||||
sample_inputs |= lcm_inputs(sample_inputs)
|
||||
|
||||
if controlnet_support:
|
||||
sample_inputs |= add_cnet_support(sample_shape, ref_unet)
|
||||
|
||||
sample_inputs_spec = get_inputs_spec(sample_inputs)
|
||||
|
||||
logger.info(f"Sample UNet inputs spec: {sample_inputs_spec}")
|
||||
logger.info("JIT tracing..")
|
||||
traced_unet = torch.jit.trace(
|
||||
coreml_unet, example_inputs=list(sample_inputs.values())
|
||||
)
|
||||
logger.info("Done.")
|
||||
|
||||
coreml_sample_inputs = get_coreml_inputs(sample_inputs)
|
||||
|
||||
coreml_unet = convert_to_coreml(
|
||||
"unet", traced_unet, coreml_sample_inputs, ["noise_pred"], unet_out_path
|
||||
)
|
||||
|
||||
del traced_unet
|
||||
gc.collect()
|
||||
|
||||
coreml_unet.save(unet_out_path)
|
||||
logger.info(f"Saved unet into {unet_out_path}")
|
||||
|
||||
|
||||
def convert(
|
||||
model_type: ModelType,
|
||||
ckpt_path: str,
|
||||
unet_out_path: str,
|
||||
batch_size: int = 1,
|
||||
sample_size: tuple[int, int] = (64, 64),
|
||||
controlnet_support: bool = False,
|
||||
lora_paths: list[str | os.PathLike] = None,
|
||||
):
|
||||
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)
|
||||
ref_pipe.fuse_lora()
|
||||
|
||||
if not os.path.exists(unet_out_path):
|
||||
convert_unet(
|
||||
ref_pipe,
|
||||
model_type,
|
||||
unet_out_path,
|
||||
batch_size,
|
||||
sample_size,
|
||||
controlnet_support,
|
||||
)
|
||||
|
||||
|
||||
def compile_model(out_path, out_name, submodule_name):
|
||||
# Compile the model
|
||||
target_path = compile_coreml_model(
|
||||
out_path, get_folder_paths(submodule_name)[0], f"{out_name}_{submodule_name}"
|
||||
)
|
||||
logger.info(f"Compiled {out_path} to {target_path}")
|
||||
return target_path
|
||||
+160
-1
@@ -5,10 +5,14 @@ from coremltools import ComputeUnit
|
||||
from python_coreml_stable_diffusion.coreml_model import CoreMLModel
|
||||
|
||||
import folder_paths
|
||||
from comfy import model_base
|
||||
from comfy import model_base, sd1_clip, sd, utils
|
||||
from comfy.model_management import get_torch_device
|
||||
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.clip import CoreMLCLIP, SDClipModelCoreML
|
||||
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 nodes import KSampler
|
||||
@@ -159,3 +163,158 @@ class CoreMLModelAdapter:
|
||||
model.diffusion_model = wrapped_model
|
||||
model_patcher = ModelPatcher(model, get_torch_device(), None)
|
||||
return (model_patcher,)
|
||||
|
||||
|
||||
class COREML_LOAD_CLIP(CoreMLLoader):
|
||||
PACKAGE_DIRNAME = "clip"
|
||||
RETURN_TYPES = ("CLIP",)
|
||||
|
||||
FUNCTION = "load_clip"
|
||||
|
||||
def load_clip(self, coreml_name, compute_unit):
|
||||
coreml_model = super().load(coreml_name, compute_unit)[0]
|
||||
|
||||
class EmptyClass:
|
||||
pass
|
||||
|
||||
clip_target = EmptyClass()
|
||||
clip_target.params = {"coreml_model": coreml_model}
|
||||
clip_target.clip = SDClipModelCoreML
|
||||
clip_target.tokenizer = sd1_clip.SD1Tokenizer
|
||||
embedding_directory = folder_paths.get_folder_paths("embeddings")[0]
|
||||
|
||||
return (CLIP(clip_target, embedding_directory),)
|
||||
|
||||
|
||||
class COREML_CONVERT(COREML_NODE):
|
||||
"""Converts a LCM model to Core ML."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
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}),
|
||||
"compute_unit": (
|
||||
[
|
||||
ComputeUnit.CPU_AND_NE.name,
|
||||
ComputeUnit.CPU_AND_GPU.name,
|
||||
ComputeUnit.ALL.name,
|
||||
ComputeUnit.CPU_ONLY.name,
|
||||
],
|
||||
),
|
||||
"controlnet_support": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"optional": {
|
||||
"lora_stack": ("LORA_STACK",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("COREML_UNET", "CLIP")
|
||||
RETURN_NAMES = ("coreml_model", "CLIP")
|
||||
FUNCTION = "convert"
|
||||
|
||||
def convert(
|
||||
self,
|
||||
ckpt_name,
|
||||
model_type,
|
||||
height,
|
||||
width,
|
||||
batch_size,
|
||||
compute_unit,
|
||||
controlnet_support,
|
||||
lora_stack=None,
|
||||
):
|
||||
"""Converts a LCM model to Core ML.
|
||||
|
||||
Args:
|
||||
height (int): Height of the target image.
|
||||
width (int): Width of the target image.
|
||||
batch_size (int): Batch size.
|
||||
compute_unit (str): Compute unit to use when loading the model.
|
||||
|
||||
Returns:
|
||||
coreml_model: The converted Core ML model.
|
||||
|
||||
The converted model is also saved to "models/unet" directory and
|
||||
can be loaded with the "LCMCoreMLLoaderUNet" node.
|
||||
"""
|
||||
h = height
|
||||
w = width
|
||||
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 ""
|
||||
|
||||
out_name = (
|
||||
f"{ckpt_name.split('.')[0]}_{batch_size}x{w}x{h}{cn_support_str}{lcm_str}"
|
||||
)
|
||||
|
||||
unet_out_path = converter.get_out_path("unet", f"{out_name}")
|
||||
|
||||
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
|
||||
|
||||
lora_stack = lora_stack or []
|
||||
lora_paths = [
|
||||
folder_paths.get_full_path("loras", lora[0]) for lora in lora_stack
|
||||
]
|
||||
|
||||
converter.convert(
|
||||
model_type=ModelType[model_type],
|
||||
ckpt_path=ckpt_path,
|
||||
unet_out_path=unet_out_path,
|
||||
sample_size=sample_size,
|
||||
batch_size=batch_size,
|
||||
controlnet_support=controlnet_support,
|
||||
lora_paths=lora_paths,
|
||||
)
|
||||
unet_target_path = converter.compile_model(
|
||||
out_path=unet_out_path, out_name=out_name, submodule_name="unet"
|
||||
)
|
||||
|
||||
_, clip = load_lora(lora_stack, ckpt_name)
|
||||
|
||||
return (CoreMLModel(unet_target_path, compute_unit, "compiled"), clip)
|
||||
|
||||
|
||||
def load_lora(lora_params, ckpt_name):
|
||||
lora_params = (
|
||||
lora_params.copy()
|
||||
if isinstance(lora_params, (list, dict, set))
|
||||
else lora_params
|
||||
)
|
||||
ckpt_name = (
|
||||
ckpt_name.copy() if isinstance(ckpt_name, (list, dict, set)) else ckpt_name
|
||||
)
|
||||
|
||||
def recursive_load_lora(lora_params, ckpt, clip):
|
||||
if len(lora_params) == 0:
|
||||
return ckpt, clip
|
||||
|
||||
lora_name, strength_model, strength_clip = lora_params[0]
|
||||
if os.path.isabs(lora_name):
|
||||
lora_path = lora_name
|
||||
else:
|
||||
lora_path = folder_paths.get_full_path("loras", lora_name)
|
||||
|
||||
lora_model, lora_clip = sd.load_lora_for_models(
|
||||
ckpt, clip, utils.load_torch_file(lora_path), strength_model, strength_clip
|
||||
)
|
||||
|
||||
# Call the function again with the new lora_model and lora_clip and the remaining tuples
|
||||
return recursive_load_lora(lora_params[1:], lora_model, lora_clip)
|
||||
|
||||
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
|
||||
ckpt, clip, _, _ = sd.load_checkpoint_guess_config(
|
||||
ckpt_path,
|
||||
output_vae=False,
|
||||
output_clip=True,
|
||||
embedding_directory=folder_paths.get_folder_paths("embeddings"),
|
||||
)
|
||||
|
||||
lora_model, lora_clip = recursive_load_lora(lora_params, ckpt, clip)
|
||||
|
||||
return lora_model, lora_clip
|
||||
|
||||
Reference in New Issue
Block a user