diff --git a/__init__.py b/__init__.py index 440fc76..b38e219 100644 --- a/__init__.py +++ b/__init__.py @@ -8,7 +8,7 @@ from coreml_suite.nodes import ( CoreMLSampler, CoreMLSamplerAdvanced, CoreMLModelAdapter, - COREML_CONVERT, + CoreMLConverter, COREML_LOAD_LORA, ) from coreml_suite.lcm import ( @@ -21,7 +21,7 @@ NODE_CLASS_MAPPINGS = { "CoreMLSamplerAdvanced": CoreMLSamplerAdvanced, "CoreMLModelAdapter": CoreMLModelAdapter, "Core ML LoRA Loader": COREML_LOAD_LORA, - "Core ML Converter": COREML_CONVERT, + "Core ML Converter": CoreMLConverter, "Core ML LCM Converter": COREML_CONVERT_LCM, } NODE_DISPLAY_NAME_MAPPINGS = { diff --git a/coreml_suite/config.py b/coreml_suite/config.py index 300e0c5..c3f25fb 100644 --- a/coreml_suite/config.py +++ b/coreml_suite/config.py @@ -11,6 +11,7 @@ class ModelVersion(Enum): SD15 = "sd15" SDXL = "sdxl" SDXL_REFINER = "sdxl_refiner" + LCM = "lcm" config_map = { diff --git a/coreml_suite/converter.py b/coreml_suite/converter.py index 92c9415..fd5edcf 100644 --- a/coreml_suite/converter.py +++ b/coreml_suite/converter.py @@ -2,44 +2,46 @@ import gc import os import shutil import time -from enum import Enum, auto import coremltools as ct import numpy as np import python_coreml_stable_diffusion.unet import torch -from diffusers import StableDiffusionPipeline, LatentConsistencyModelPipeline +from diffusers import ( + StableDiffusionPipeline, + LatentConsistencyModelPipeline, + StableDiffusionXLPipeline, +) from python_coreml_stable_diffusion.unet import ( UNet2DConditionModel, + UNet2DConditionModelXL, AttentionImplementations, ) +from coreml_suite.config import ModelVersion 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() - - class StableDiffusionLCMPipeline(LatentConsistencyModelPipeline): pass MODEL_TYPE_TO_UNET_CLS = { - ModelType.SD15: UNet2DConditionModel, - ModelType.LCM: UNet2DConditionModelLCM, + ModelVersion.SD15: UNet2DConditionModel, + ModelVersion.SDXL: UNet2DConditionModelXL, + ModelVersion.LCM: UNet2DConditionModelLCM, } MODEL_TYPE_TO_PIPE_CLS = { - ModelType.SD15: StableDiffusionPipeline, - ModelType.LCM: StableDiffusionLCMPipeline, + ModelVersion.SD15: StableDiffusionPipeline, + ModelVersion.SDXL: StableDiffusionXLPipeline, + ModelVersion.LCM: StableDiffusionLCMPipeline, } -def get_unet(model_type: ModelType, ref_pipe): +def get_unet(model_type: ModelVersion, ref_pipe): ref_unet = ref_pipe.unet unet_cls = MODEL_TYPE_TO_UNET_CLS[model_type] @@ -159,6 +161,25 @@ def lcm_inputs(sample_unet_inputs): return {"timestep_cond": torch.randn(batch_size, 256).to(torch.float32)} +def sdxl_inputs(sample_unet_inputs, ref_pipe): + sample_shape = sample_unet_inputs["sample"].shape + batch_size = sample_shape[0] + h = sample_shape[2] * 8 + w = sample_shape[3] * 8 + original_size = (h, w) # output_resolution + crops_coords_top_left = (0, 0) # topleft_crop_cond + target_size = (h, w) # resolution_cond + + time_ids_list = list(original_size + crops_coords_top_left + target_size) + time_ids = torch.tensor(time_ids_list).repeat(batch_size, 1).to(torch.int64) + text_embeds_shape = (batch_size, ref_pipe.text_encoder_2.config.hidden_size) + + return { + "time_ids": time_ids, + "text_embeds": torch.randn(*text_embeds_shape).to(torch.float32), + } + + def get_inputs_spec(inputs): inputs_spec = {k: (v.shape, v.dtype) for k, v in inputs.items()} return inputs_spec @@ -217,13 +238,13 @@ def add_cnet_support(sample_shape, reference_unet): def convert_unet( ref_pipe, - model_type: ModelType, + model_version: ModelVersion, 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) + coreml_unet = get_unet(model_version, ref_pipe) ref_unet = ref_pipe.unet sample_shape = ( @@ -242,9 +263,12 @@ def convert_unet( batch_size, encoder_hidden_states_shape, sample_shape, scheduler ) - if model_type == ModelType.LCM: + if model_version == ModelVersion.LCM: sample_inputs |= lcm_inputs(sample_inputs) + if model_version == ModelVersion.SDXL: + sample_inputs |= sdxl_inputs(sample_inputs, ref_pipe) + if controlnet_support: sample_inputs |= add_cnet_support(sample_shape, ref_unet) @@ -272,6 +296,7 @@ def convert_unet( def convert( ckpt_path: str, + model_version: ModelVersion, unet_out_path: str, batch_size: int = 1, sample_size: tuple[int, int] = (64, 64), @@ -288,10 +313,7 @@ def convert( AttentionImplementations(attn_impl) ) - model_type = ModelType.SD15 - - pipe_cls = MODEL_TYPE_TO_PIPE_CLS[model_type] - ref_pipe = pipe_cls.from_single_file(ckpt_path, original_config_file=config_path) + ref_pipe = get_pipeline(ckpt_path, config_path, model_version) for i, lora_weight in enumerate(lora_weights or []): lora_path, strength = lora_weight @@ -302,7 +324,7 @@ def convert( convert_unet( ref_pipe, - model_type, + model_version, unet_out_path, batch_size, sample_size, @@ -310,6 +332,12 @@ def convert( ) +def get_pipeline(ckpt_path, config_path, model_version): + pipe_cls = MODEL_TYPE_TO_PIPE_CLS[model_version] + ref_pipe = pipe_cls.from_single_file(ckpt_path, original_config_file=config_path) + return ref_pipe + + def compile_model(out_path, out_name, submodule_name): # Compile the model target_path = compile_coreml_model( diff --git a/coreml_suite/models.py b/coreml_suite/models.py index 53c2e67..c2905c4 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -132,7 +132,7 @@ class CoreMLInputs: chunked_ts_cond = chunk_batch(self.ts_cond, ts_cond_shape) chunked_time_ids = [None] * len(chunked_x) - if expected_inputs["time_ids"] is not None: + if expected_inputs.get("time_ids") is not None: time_ids_shape = expected_inputs["time_ids"]["shape"] if self.time_ids is None: self.time_ids = torch.zeros(len(chunked_x), *time_ids_shape[1:]).to( @@ -141,7 +141,7 @@ class CoreMLInputs: chunked_time_ids = chunk_batch(self.time_ids, time_ids_shape) chunked_text_embeds = [None] * len(chunked_x) - if expected_inputs["text_embeds"] is not None: + if expected_inputs.get("text_embeds") is not None: text_embeds_shape = expected_inputs["text_embeds"]["shape"] if self.text_embeds is None: self.text_embeds = torch.zeros( diff --git a/coreml_suite/nodes.py b/coreml_suite/nodes.py index 93524c0..cf1a3aa 100644 --- a/coreml_suite/nodes.py +++ b/coreml_suite/nodes.py @@ -7,6 +7,7 @@ from python_coreml_stable_diffusion.unet import AttentionImplementations import folder_paths from coreml_suite import COREML_NODE from coreml_suite import converter +from coreml_suite.config import ModelVersion from coreml_suite.lcm.utils import add_lcm_model_options, lcm_patch, is_lcm from coreml_suite.logger import logger from nodes import KSampler, LoraLoader, KSamplerAdvanced @@ -208,7 +209,7 @@ class CoreMLModelAdapter(COREML_NODE): return (model_patcher,) -class COREML_CONVERT(COREML_NODE): +class CoreMLConverter(COREML_NODE): """Converts a LCM model to Core ML.""" @classmethod @@ -216,6 +217,13 @@ class COREML_CONVERT(COREML_NODE): return { "required": { "ckpt_name": (folder_paths.get_filename_list("checkpoints"),), + "model_version": ( + [ + ModelVersion.SD15.name, + ModelVersion.SDXL.name, + ModelVersion.LCM.name, + ], + ), "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}), @@ -248,6 +256,7 @@ class COREML_CONVERT(COREML_NODE): def convert( self, ckpt_name, + model_version, height, width, batch_size, @@ -270,6 +279,8 @@ class COREML_CONVERT(COREML_NODE): The converted model is also saved to "models/unet" directory and can be loaded with the "LCMCoreMLLoaderUNet" node. """ + model_version = ModelVersion[model_version] + lora_params = lora_params or {} lora_params = [(k, v[0]) for k, v in lora_params.items()] lora_params = sorted(lora_params, key=lambda lora: lora[0]) @@ -317,6 +328,7 @@ class COREML_CONVERT(COREML_NODE): converter.convert( ckpt_path=ckpt_path, + model_version=model_version, unet_out_path=unet_out_path, sample_size=sample_size, batch_size=batch_size,