Base SDXL conversion works

This commit is contained in:
aszc-dev
2023-11-24 12:14:15 +01:00
committed by aszc
parent 763ca3961b
commit 0c78803b25
5 changed files with 66 additions and 25 deletions
+2 -2
View File
@@ -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 = {
+1
View File
@@ -11,6 +11,7 @@ class ModelVersion(Enum):
SD15 = "sd15"
SDXL = "sdxl"
SDXL_REFINER = "sdxl_refiner"
LCM = "lcm"
config_map = {
+48 -20
View File
@@ -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(
+2 -2
View File
@@ -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
View File
@@ -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,