From 701443f59e566c57f03d991f15de71a1a7ad6462 Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Wed, 8 Nov 2023 17:51:27 +0100 Subject: [PATCH] Wrapped Core ML Model is now diffusion_model attribute of BaseModel --- coreml_suite/models.py | 23 ++++++++--------------- coreml_suite/nodes.py | 17 +++++++++++------ 2 files changed, 19 insertions(+), 21 deletions(-) diff --git a/coreml_suite/models.py b/coreml_suite/models.py index 3b456ae..0cea66e 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -24,18 +24,16 @@ def get_model_config(): return model_config -class CoreMLModelWrapper(BaseModel): - def __init__(self, model_config, coreml_model): - super().__init__(model_config) +class CoreMLModelWrapper: + def __init__(self, coreml_model): self.diffusion_model = coreml_model + self.dtype = torch.float16 def apply_model( self, x, t, - c_concat=None, c_crossattn=None, - c_adm=None, control=None, transformer_options={}, **kwargs, @@ -108,18 +106,13 @@ class CoreMLModelWrapper(BaseModel): def expected_inputs(self): return self.diffusion_model.expected_inputs - def __call__(self, latents, ts, encoder_hidden_states, **kwargs): - return ( - self.apply_model(latents, ts, c_crossattn=encoder_hidden_states, **kwargs), + def __call__(self, latents, ts, context, control, transformer_options, **kwargs): + return self.apply_model( + latents, ts, context, control, transformer_options, **kwargs ) class CoreMLModelWrapperLCM(CoreMLModelWrapper): - def __init__(self, model_config, coreml_model): - super().__init__(model_config, coreml_model) + def __init__(self, coreml_model): + super().__init__(coreml_model) self.config = None - - def __call__(self, latents, ts, encoder_hidden_states, **kwargs): - return ( - self.apply_model(latents, ts, c_crossattn=encoder_hidden_states, **kwargs), - ) diff --git a/coreml_suite/nodes.py b/coreml_suite/nodes.py index f73c79e..48f1695 100644 --- a/coreml_suite/nodes.py +++ b/coreml_suite/nodes.py @@ -5,6 +5,7 @@ from coremltools import ComputeUnit from python_coreml_stable_diffusion.coreml_model import CoreMLModel 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.logger import logger @@ -41,8 +42,10 @@ class CoreMLSampler(KSampler): denoise=1.0, ): model_config = get_model_config() - wrapped_model = CoreMLModelWrapper(model_config, coreml_model) - 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.") @@ -50,7 +53,7 @@ class CoreMLSampler(KSampler): latent_image = {"samples": torch.zeros(expected[0] // 2, *expected[1:])} return super().sample( - model, + model_patcher, seed, steps, cfg, @@ -130,6 +133,8 @@ class CoreMLModelAdapter: def wrap(self, coreml_model): model_config = get_model_config() - wrapped_model = CoreMLModelWrapper(model_config, coreml_model) - patched_model = ModelPatcher(wrapped_model, get_torch_device(), None) - return (patched_model,) + 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) + return (model_patcher,)