Handle SDXL config

This commit is contained in:
aszc-dev
2023-11-24 12:14:15 +01:00
committed by aszc
parent ae9a9874c5
commit 763ca3961b
3 changed files with 128 additions and 44 deletions
+82 -26
View File
@@ -1,38 +1,94 @@
from enum import Enum
import torch
from comfy import supported_models_base
from comfy import latent_formats
from comfy.model_detection import convert_config
SD15 = {
"use_checkpoint": False,
"image_size": 32,
"out_channels": 4,
"use_spatial_transformer": True,
"legacy": False,
"adm_in_channels": None,
"dtype": torch.float16,
"in_channels": 4,
"model_channels": 320,
"num_res_blocks": 2,
"attention_resolutions": [1, 2, 4],
"transformer_depth": [1, 1, 1, 0],
"channel_mult": [1, 2, 4, 4],
"transformer_depth_middle": 1,
"use_linear_in_transformer": False,
"context_dim": 768,
"num_heads": 8,
"disable_unet_model_creation": True,
class ModelVersion(Enum):
SD15 = "sd15"
SDXL = "sdxl"
SDXL_REFINER = "sdxl_refiner"
config_map = {
ModelVersion.SD15: {
"use_checkpoint": False,
"image_size": 32,
"out_channels": 4,
"use_spatial_transformer": True,
"legacy": False,
"adm_in_channels": None,
"dtype": torch.float16,
"in_channels": 4,
"model_channels": 320,
"num_res_blocks": 2,
"attention_resolutions": [1, 2, 4],
"transformer_depth": [1, 1, 1, 0],
"channel_mult": [1, 2, 4, 4],
"transformer_depth_middle": 1,
"use_linear_in_transformer": False,
"context_dim": 768,
"num_heads": 8,
"disable_unet_model_creation": True,
},
ModelVersion.SDXL: {
"use_checkpoint": False,
"image_size": 32,
"out_channels": 4,
"use_spatial_transformer": True,
"legacy": False,
"num_classes": "sequential",
"adm_in_channels": 2816,
"dtype": torch.float16,
"in_channels": 4,
"model_channels": 320,
"num_res_blocks": 2,
"attention_resolutions": [2, 4],
"transformer_depth": [0, 2, 10],
"channel_mult": [1, 2, 4],
"transformer_depth_middle": 10,
"use_linear_in_transformer": True,
"context_dim": 2048,
"num_head_channels": 64,
"disable_unet_model_creation": True,
},
ModelVersion.SDXL_REFINER: {
"use_checkpoint": False,
"image_size": 32,
"out_channels": 4,
"use_spatial_transformer": True,
"legacy": False,
"num_classes": "sequential",
"adm_in_channels": 2560,
"dtype": torch.float16,
"in_channels": 4,
"model_channels": 384,
"num_res_blocks": 2,
"attention_resolutions": [2, 4],
"transformer_depth": [0, 4, 4, 0],
"channel_mult": [1, 2, 4, 4],
"transformer_depth_middle": 4,
"use_linear_in_transformer": True,
"context_dim": 1280,
"num_head_channels": 64,
"disable_unet_model_creation": True,
},
}
latent_format_map = {
ModelVersion.SD15: latent_formats.SD15,
ModelVersion.SDXL: latent_formats.SDXL,
ModelVersion.SDXL_REFINER: latent_formats.SDXL,
}
def get_unet_config():
return convert_config(SD15)
def get_model_config():
config = supported_models_base.BASE(get_unet_config())
config.latent_format = latent_formats.SD15()
def get_model_config(model_version: ModelVersion):
unet_config = convert_config(config_map[model_version])
config = supported_models_base.BASE(unet_config)
config.latent_format = latent_format_map[model_version]()
return config
+45 -7
View File
@@ -4,9 +4,10 @@ import torch
from comfy import model_base
from comfy.model_management import get_torch_device
from comfy.model_patcher import ModelPatcher
from coreml_suite.config import get_model_config
from coreml_suite.config import get_model_config, ModelVersion
from coreml_suite.controlnet import extract_residual_kwargs, chunk_control
from coreml_suite.latents import chunk_batch, merge_chunks
from coreml_suite.lcm.utils import is_lcm
from coreml_suite.logger import logger
@@ -38,6 +39,28 @@ class CoreMLModelWrapper:
def expected_inputs(self):
return self.coreml_model.expected_inputs
@property
def is_lcm(self):
return is_lcm(self.coreml_model)
@property
def is_sdxl_base(self):
return is_sdxl_base(self.coreml_model)
@property
def is_sdxl_refiner(self):
return is_sdxl_refiner(self.coreml_model)
@property
def config(self):
if self.is_sdxl_base:
return get_model_config(ModelVersion.SDXL)
if self.is_sdxl_refiner:
return get_model_config(ModelVersion.SDXL_REFINER)
return get_model_config(ModelVersion.SD15)
class CoreMLModelWrapperLCM(CoreMLModelWrapper):
def __init__(self, coreml_model):
@@ -155,6 +178,20 @@ def is_sdxl(coreml_model):
)
def is_sdxl_base(coreml_model):
return (
is_sdxl(coreml_model)
and coreml_model.expected_inputs["time_ids"]["shape"][1] == 6
)
def is_sdxl_refiner(coreml_model):
return (
is_sdxl(coreml_model)
and coreml_model.expected_inputs["time_ids"]["shape"][1] == 5
)
def sdxl_model_function_wrapper(time_ids, text_embeds):
def wrapper(model_function, params):
x = params["input"]
@@ -191,7 +228,7 @@ def add_sdxl_model_options(model_patcher, positive, negative):
neg_dict.get("crop_w", 0),
]
if model_patcher.model.diffusion_model.expected_inputs["time_ids"]["shape"][1] == 6:
if model_patcher.model.diffusion_model.is_sdxl_base:
base_pos_time_ids = [
pos_dict.get("target_height", 768),
pos_dict.get("target_width", 768),
@@ -204,7 +241,7 @@ def add_sdxl_model_options(model_patcher, positive, negative):
]
neg_time_ids += base_neg_time_ids
else:
if model_patcher.model.diffusion_model.is_sdxl_refiner:
refiner_pos_time_ids = [
pos_dict.get("aesthetic_score", 6),
]
@@ -239,13 +276,14 @@ def get_latent_image(coreml_model, latent_image):
def get_model_patcher(coreml_model):
model_config = get_model_config()
wrapped_model = CoreMLModelWrapper(coreml_model)
if is_sdxl(coreml_model):
model = model_base.SDXL(model_config, device=get_torch_device())
if wrapped_model.is_sdxl_base:
model = model_base.SDXL(wrapped_model.config, device=get_torch_device())
elif wrapped_model.is_sdxl_refiner:
model = model_base.SDXLRefiner(wrapped_model.config, device=get_torch_device())
else:
model = model_base.BaseModel(model_config, device=get_torch_device())
model = model_base.BaseModel(wrapped_model.config, device=get_torch_device())
model.diffusion_model = wrapped_model
model_patcher = ModelPatcher(model, get_torch_device(), None)
+1 -11
View File
@@ -1,14 +1,10 @@
import os
import torch
from coremltools import ComputeUnit
from python_coreml_stable_diffusion.coreml_model import CoreMLModel
from python_coreml_stable_diffusion.unet import AttentionImplementations
import folder_paths
from comfy import model_base
from comfy.model_management import get_torch_device
from comfy.model_patcher import ModelPatcher
from coreml_suite import COREML_NODE
from coreml_suite import converter
from coreml_suite.lcm.utils import add_lcm_model_options, lcm_patch, is_lcm
@@ -16,13 +12,11 @@ from coreml_suite.logger import logger
from nodes import KSampler, LoraLoader, KSamplerAdvanced
from coreml_suite.models import (
CoreMLModelWrapper,
add_sdxl_model_options,
is_sdxl,
get_model_patcher,
get_latent_image,
)
from coreml_suite.config import get_model_config
class CoreMLSampler(COREML_NODE, KSampler):
@@ -210,11 +204,7 @@ class CoreMLModelAdapter(COREML_NODE):
CATEGORY = "Core ML Suite"
def wrap(self, coreml_model):
model_config = get_model_config()
wrapped_model = CoreMLModelWrapper(coreml_model)
model = model_base.BaseModel(model_config, device=get_torch_device())
model.diffusion_model = wrapped_model
model_patcher = ModelPatcher(model, get_torch_device(), None)
model_patcher = get_model_patcher(coreml_model)
return (model_patcher,)