Handle SDXL config
This commit is contained in:
+82
-26
@@ -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
@@ -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
@@ -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,)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user