From 3d4329c55bde4ae8e9223f0a2b43046428e39321 Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Sun, 29 Oct 2023 14:59:39 +0100 Subject: [PATCH] Rework UNet loader --- coreml_suite/loaders.py | 52 ++++++++++++++++++++++------------------- 1 file changed, 28 insertions(+), 24 deletions(-) diff --git a/coreml_suite/loaders.py b/coreml_suite/loaders.py index 90ec623..da5aed6 100644 --- a/coreml_suite/loaders.py +++ b/coreml_suite/loaders.py @@ -1,3 +1,5 @@ +import os.path + from coremltools import ComputeUnit from python_coreml_stable_diffusion.coreml_model import CoreMLModel @@ -9,6 +11,7 @@ from comfy.model_patcher import ModelPatcher from coreml_suite.logger import logger from coreml_suite.model import CoreMLModelWrapper + class CoreMLLoader: PACKAGE_DIRNAME = "" @@ -26,21 +29,22 @@ class CoreMLLoader: } } - RETURN_TYPES = ("COREML_MODEL",) FUNCTION = "load" CATEGORY = "CoreML Suite" @classmethod def coreml_filenames(cls): + extensions = [".mlmodelc", ".mlpackage"] + all_paths = folder_paths.get_filename_list_(cls.PACKAGE_DIRNAME)[1] + coreml_paths = folder_paths.filter_files_extensions(all_paths, extensions) + return { - p.split('/')[-1]: - p for p in - folder_paths.get_filename_list_(cls.PACKAGE_DIRNAME)[1] - if p.endswith((".mlpackage", ".mlmodelc")) + os.path.split(p)[-1]: p + for p in coreml_paths } def load(self, coreml_name, compute_unit): - logger.info(f"Loading {coreml_name}") + logger.info(f"Loading {coreml_name} to {compute_unit}") coreml_path = self.coreml_filenames()[coreml_name] @@ -73,25 +77,25 @@ class CoreMLLoaderTextEncoder(CoreMLLoader): class CoreMLLoaderUNet(CoreMLLoader): PACKAGE_DIRNAME = "unet" - RETURN_TYPES = ("MODEL",) + RETURN_TYPES = ("coreml_model",) - def _load(self, coreml_path, compute_unit, sources): - # 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], - } - coreml_model = CoreMLModelWrapper(model_config, coreml_path, - compute_unit, sources) - - return (ModelPatcher(coreml_model, model_management.get_torch_device(), - None),) + # def _load(self, coreml_path, compute_unit, sources): + # # 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], + # } + # coreml_model = CoreMLModelWrapper(model_config, coreml_path, + # compute_unit, sources) + # + # return (ModelPatcher(coreml_model, model_management.get_torch_device(), + # None),) class CoreMLLoaderVAE(CoreMLLoader):