From c3038501eb439b455cd9638dcf2842e7fd2b8fa8 Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Mon, 30 Oct 2023 11:15:42 +0100 Subject: [PATCH] Rearrange stuff --- __init__.py | 2 +- coreml_suite/__init__.py | 4 -- coreml_suite/{utils.py => controlnet.py} | 0 coreml_suite/latents.py | 20 ++++++ coreml_suite/models.py | 2 +- coreml_suite/{loaders.py => nodes.py} | 85 ++++++++++++++++-------- coreml_suite/samplers.py | 74 --------------------- tests/test_latents.py | 2 +- 8 files changed, 79 insertions(+), 110 deletions(-) rename coreml_suite/{utils.py => controlnet.py} (100%) create mode 100644 coreml_suite/latents.py rename coreml_suite/{loaders.py => nodes.py} (51%) delete mode 100644 coreml_suite/samplers.py diff --git a/__init__.py b/__init__.py index 5811e2b..b259a08 100644 --- a/__init__.py +++ b/__init__.py @@ -3,7 +3,7 @@ import sys sys.path.append(os.path.dirname(__file__)) -from coreml_suite import CoreMLLoaderUNet, CoreMLSampler +from coreml_suite.nodes import CoreMLLoaderUNet, CoreMLSampler NODE_CLASS_MAPPINGS = { "CoreMLUNetLoader": CoreMLLoaderUNet, diff --git a/coreml_suite/__init__.py b/coreml_suite/__init__.py index d1075e8..e69de29 100644 --- a/coreml_suite/__init__.py +++ b/coreml_suite/__init__.py @@ -1,4 +0,0 @@ -from coreml_suite.loaders import CoreMLLoaderUNet -from coreml_suite.samplers import CoreMLSampler - -__all__ = ["CoreMLLoaderUNet", "CoreMLSampler"] diff --git a/coreml_suite/utils.py b/coreml_suite/controlnet.py similarity index 100% rename from coreml_suite/utils.py rename to coreml_suite/controlnet.py diff --git a/coreml_suite/latents.py b/coreml_suite/latents.py new file mode 100644 index 0000000..469074a --- /dev/null +++ b/coreml_suite/latents.py @@ -0,0 +1,20 @@ +import torch +from torchvision.transforms.functional import resize + +from coreml_suite.logger import logger + + +def reshape_latent_image(latent_image, target_shape): + if latent_image is None: + logger.warning("No latent image provided, using zeros.") + return {"samples": torch.zeros(target_shape)} + + if latent_image["samples"].shape == target_shape: + return latent_image + + logger.warning( + "Latent image shape does not match model input shape," + " resizing to match models expected input shape." + ) + resized = resize(latent_image["samples"], target_shape[-2:]) + return {"samples": resized} diff --git a/coreml_suite/models.py b/coreml_suite/models.py index 0b45145..8540f8d 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -6,7 +6,7 @@ from comfy import supported_models_base from comfy.latent_formats import SD15 from comfy.model_base import BaseModel -from coreml_suite.utils import expand_inputs, extract_residual_kwargs +from coreml_suite.controlnet import expand_inputs, extract_residual_kwargs def get_model_config(): diff --git a/coreml_suite/loaders.py b/coreml_suite/nodes.py similarity index 51% rename from coreml_suite/loaders.py rename to coreml_suite/nodes.py index 8d201c6..7246a6f 100644 --- a/coreml_suite/loaders.py +++ b/coreml_suite/nodes.py @@ -1,11 +1,65 @@ -import os.path +import os from coremltools import ComputeUnit from python_coreml_stable_diffusion.coreml_model import CoreMLModel import folder_paths - +from comfy.model_management import get_torch_device +from comfy.model_patcher import ModelPatcher from coreml_suite.logger import logger +from coreml_suite.latents import reshape_latent_image +from nodes import KSampler + +from coreml_suite.models import CoreMLModelWrapper, get_model_config + + +class CoreMLSampler(KSampler): + @classmethod + def INPUT_TYPES(s): + old_required = KSampler.INPUT_TYPES()["required"].copy() + old_required.pop("model") + old_required.pop("latent_image") + new_required = {"coreml_model": ("COREML_UNET",)} + return { + "required": new_required | old_required, + "optional": {"latent_image": ("LATENT",)}, + } + + CATEGORY = "Core ML Suite" + + def sample( + self, + coreml_model, + seed, + steps, + cfg, + sampler_name, + scheduler, + positive, + negative, + latent_image=None, + denoise=1.0, + ): + sample_shape = coreml_model.expected_inputs["sample"]["shape"] + latent_image = reshape_latent_image(latent_image, sample_shape) + latent_image["samples"] = latent_image["samples"][0:1] + + model_config = get_model_config() + wrapped_model = CoreMLModelWrapper(model_config, coreml_model) + model = ModelPatcher(wrapped_model, get_torch_device(), None) + + return super().sample( + model, + seed, + steps, + cfg, + sampler_name, + scheduler, + positive, + negative, + latent_image, + denoise, + ) class CoreMLLoader: @@ -51,34 +105,7 @@ class CoreMLLoader: return (CoreMLModel(coreml_path, compute_unit, sources),) -class CoreMLLoaderCkpt(CoreMLLoader): - PACKAGE_DIRNAME = "checkpoints" - RETURN_TYPES = ("MODEL", "CLIP", "VAE") - - def load(self, coreml_name, compute_unit): - # TODO: Implement this - pass - - -class CoreMLLoaderTextEncoder(CoreMLLoader): - PACKAGE_DIRNAME = "clip" - RETURN_TYPES = ("CLIP",) - - def load(self, coreml_name, compute_unit): - # TODO: Implement this - pass - - class CoreMLLoaderUNet(CoreMLLoader): PACKAGE_DIRNAME = "unet" RETURN_TYPES = ("COREML_UNET",) RETURN_NAMES = ("coreml_model",) - - -class CoreMLLoaderVAE(CoreMLLoader): - PACKAGE_DIRNAME = "vae" - RETURN_TYPES = ("VAE",) - - def load(self, coreml_name, compute_unit): - # TODO: Implement this - pass diff --git a/coreml_suite/samplers.py b/coreml_suite/samplers.py deleted file mode 100644 index 89f44db..0000000 --- a/coreml_suite/samplers.py +++ /dev/null @@ -1,74 +0,0 @@ -import torch -from torchvision.transforms.functional import resize - -from comfy.model_management import get_torch_device -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 - - -def reshape_latent_image(latent_image, target_shape): - if latent_image is None: - logger.warning("No latent image provided, using zeros.") - return {"samples": torch.zeros(target_shape)} - - if latent_image["samples"].shape == target_shape: - return latent_image - - logger.warning( - "Latent image shape does not match model input shape," - " resizing to match models expected input shape." - ) - resized = resize(latent_image["samples"], target_shape[-2:]) - return {"samples": resized} - - -class CoreMLSampler(KSampler): - @classmethod - def INPUT_TYPES(s): - old_required = KSampler.INPUT_TYPES()["required"].copy() - old_required.pop("model") - old_required.pop("latent_image") - new_required = {"coreml_model": ("COREML_UNET",)} - return { - "required": new_required | old_required, - "optional": {"latent_image": ("LATENT",)}, - } - - CATEGORY = "Core ML Suite" - - def sample( - self, - coreml_model, - seed, - steps, - cfg, - sampler_name, - scheduler, - positive, - negative, - latent_image=None, - denoise=1.0, - ): - sample_shape = coreml_model.expected_inputs["sample"]["shape"] - latent_image = reshape_latent_image(latent_image, sample_shape) - latent_image["samples"] = latent_image["samples"][0:1] - - model_config = get_model_config() - wrapped_model = CoreMLModelWrapper(model_config, coreml_model) - model = ModelPatcher(wrapped_model, get_torch_device(), None) - - return super().sample( - model, - seed, - steps, - cfg, - sampler_name, - scheduler, - positive, - negative, - latent_image, - denoise, - ) diff --git a/tests/test_latents.py b/tests/test_latents.py index aa62a7d..dd47294 100644 --- a/tests/test_latents.py +++ b/tests/test_latents.py @@ -2,7 +2,7 @@ import pytest import torch -from coreml_suite.samplers import reshape_latent_image +from coreml_suite.latents import reshape_latent_image def test_fix_latents_no_latent_image():