diff --git a/__init__.py b/__init__.py index 4466edb..aadaeac 100644 --- a/__init__.py +++ b/__init__.py @@ -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", } diff --git a/coreml_suite/clip.py b/coreml_suite/clip.py new file mode 100644 index 0000000..8a46b51 --- /dev/null +++ b/coreml_suite/clip.py @@ -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 diff --git a/coreml_suite/converter.py b/coreml_suite/converter.py new file mode 100644 index 0000000..5f39563 --- /dev/null +++ b/coreml_suite/converter.py @@ -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 diff --git a/coreml_suite/nodes.py b/coreml_suite/nodes.py index 04b66a0..0db554e 100644 --- a/coreml_suite/nodes.py +++ b/coreml_suite/nodes.py @@ -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