diff --git a/coreml_suite/config.py b/coreml_suite/config.py new file mode 100644 index 0000000..92293e8 --- /dev/null +++ b/coreml_suite/config.py @@ -0,0 +1,36 @@ +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, +} + + +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() + return config diff --git a/coreml_suite/lcm/lcm_sampler.py b/coreml_suite/lcm/lcm_sampler.py index 3f0796e..7488c3d 100644 --- a/coreml_suite/lcm/lcm_sampler.py +++ b/coreml_suite/lcm/lcm_sampler.py @@ -6,12 +6,14 @@ from diffusers.utils.torch_utils import randn_tensor from tqdm import tqdm import latent_preview +from comfy import model_base from comfy.model_management import get_torch_device from comfy.model_patcher import ModelPatcher from coreml_suite.lcm.lcm_scheduler import LCMScheduler from coreml_suite.logger import logger -from coreml_suite.models import get_model_config, CoreMLModelWrapperLCM +from coreml_suite.models import CoreMLModelWrapper from coreml_suite.nodes import CoreMLSampler +from coreml_suite.config import get_model_config class CoreMLSamplerLCM(CoreMLSampler): @@ -47,8 +49,10 @@ class CoreMLSamplerLCM(CoreMLSampler): **kwargs, ): model_config = get_model_config() - wrapped_model = CoreMLModelWrapperLCM(model_config, coreml_model) - patched_model = ModelPatcher(wrapped_model, get_torch_device(), None) + 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) if latent_image is None: logger.warning("No latent image provided, using empty tensor.") @@ -57,11 +61,11 @@ class CoreMLSamplerLCM(CoreMLSampler): positive = positive[0][0] - callback = latent_preview.prepare_callback(patched_model, steps, None) + callback = latent_preview.prepare_callback(model_patcher, steps, None) torch.manual_seed(seed) return self._sample( - patched_model, steps, cfg, positive, latent_image, denoise, callback + model_patcher, steps, cfg, positive, latent_image, denoise, callback ) def _sample( diff --git a/coreml_suite/models.py b/coreml_suite/models.py index 70c0ba7..e7227a6 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -1,28 +1,10 @@ import numpy as np import torch -from comfy import supported_models_base -from comfy.latent_formats import SD15 - from coreml_suite.controlnet import extract_residual_kwargs, chunk_control from coreml_suite.latents import chunk_batch, merge_chunks -def get_model_config(): - # TODO: This is a dummy model config, but it should be enough to - # get the model to load - implement a proper model config - model_config = supported_models_base.BASE({}) - model_config.latent_format = SD15() - model_config.unet_config = { - "disable_unet_model_creation": True, - "num_res_blocks": 2, - "attention_resolutions": [1, 2, 4], - "channel_mult": [1, 2, 4, 4], - "transformer_depth": [1, 1, 1, 0], - } - return model_config - - class CoreMLModelWrapper: def __init__(self, coreml_model): self.coreml_model = coreml_model diff --git a/coreml_suite/nodes.py b/coreml_suite/nodes.py index 48f1695..899523e 100644 --- a/coreml_suite/nodes.py +++ b/coreml_suite/nodes.py @@ -11,7 +11,8 @@ from comfy.model_patcher import ModelPatcher from coreml_suite.logger import logger from nodes import KSampler -from coreml_suite.models import CoreMLModelWrapper, get_model_config +from coreml_suite.models import CoreMLModelWrapper +from coreml_suite.config import get_model_config class CoreMLSampler(KSampler): diff --git a/tests/test_chunks.py b/tests/test_chunks.py index afa8821..edb919f 100644 --- a/tests/test_chunks.py +++ b/tests/test_chunks.py @@ -1,5 +1,3 @@ -from unittest import mock - import pytest import torch @@ -8,11 +6,9 @@ from comfy.model_management import get_torch_device from coreml_suite.latents import chunk_batch, merge_chunks from coreml_suite.controlnet import chunk_control from coreml_suite.models import ( - CoreMLModelWrapper, - get_model_config, - CoreMLModelWrapperLCM, CoreMLInputs, ) +from coreml_suite.config import get_model_config @pytest.fixture