Base SDXL conversion works
This commit is contained in:
+2
-2
@@ -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 = {
|
||||
|
||||
@@ -11,6 +11,7 @@ class ModelVersion(Enum):
|
||||
SD15 = "sd15"
|
||||
SDXL = "sdxl"
|
||||
SDXL_REFINER = "sdxl_refiner"
|
||||
LCM = "lcm"
|
||||
|
||||
|
||||
config_map = {
|
||||
|
||||
+48
-20
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
+13
-1
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user