Basic conversion + LoRA support works

This commit is contained in:
aszc-dev
2023-11-17 22:55:13 +01:00
committed by aszc
parent fc1132a5d5
commit 45be6761d1
4 changed files with 494 additions and 2 deletions
+11 -1
View File
@@ -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",
}
+22
View File
@@ -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
+301
View File
@@ -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
View File
@@ -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