From 4c438e1ee60b3eeee0566f33ad8f7984e75e674f Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Fri, 3 Nov 2023 17:40:28 +0100 Subject: [PATCH] Leverage Comfy's mechanisms to enable LCM ControlNet support --- coreml_suite/lcm/lcm_sampler.py | 27 +++++++++++++++------------ coreml_suite/lcm/nodes.py | 6 ++---- coreml_suite/nodes.py | 3 ++- 3 files changed, 19 insertions(+), 17 deletions(-) diff --git a/coreml_suite/lcm/lcm_sampler.py b/coreml_suite/lcm/lcm_sampler.py index 7488c3d..a46ac2c 100644 --- a/coreml_suite/lcm/lcm_sampler.py +++ b/coreml_suite/lcm/lcm_sampler.py @@ -9,6 +9,8 @@ import latent_preview from comfy import model_base from comfy.model_management import get_torch_device from comfy.model_patcher import ModelPatcher +from comfy.sample import prepare_sampling +from comfy.samplers import sampling_function from coreml_suite.lcm.lcm_scheduler import LCMScheduler from coreml_suite.logger import logger from coreml_suite.models import CoreMLModelWrapper @@ -59,11 +61,13 @@ class CoreMLSamplerLCM(CoreMLSampler): expected = coreml_model.expected_inputs["sample"]["shape"] latent_image = {"samples": torch.zeros(*expected).to(get_torch_device())} - positive = positive[0][0] - callback = latent_preview.prepare_callback(model_patcher, steps, None) torch.manual_seed(seed) + model, positive, _, _, _ = prepare_sampling( + model_patcher, latent_image["samples"].shape, positive, (), None + ) + return self._sample( model_patcher, steps, cfg, positive, latent_image, denoise, callback ) @@ -74,7 +78,7 @@ class CoreMLSamplerLCM(CoreMLSampler): device = get_torch_device() batch_size = latent_image["samples"].shape[0] - prompt_embeds = self.prepare_prompt_embeds(batch_size, positive) + # prompt_embeds = self.prepare_prompt_embeds(batch_size, positive) timesteps = self.prepare_timesteps(denoise, device, steps) @@ -90,12 +94,13 @@ class CoreMLSamplerLCM(CoreMLSampler): for i, t in enumerate(iterator): ts = torch.full((batch_size,), t, device=device, dtype=torch.float16) - model_pred = model.model( - latents, - ts, - encoder_hidden_states=prompt_embeds, - timestep_cond=w_embedding, - )[0] + model_options = { + "transformer_options": {"timestep_cond": w_embedding}, + "sampler_cfg_function": lambda x: x["cond"].to(device), + } + model_pred = sampling_function( + model.apply_model, latents, ts, None, positive, denoise, model_options + ) # compute the previous noisy sample x_t -> x_t-1 latents, denoised = self.scheduler.step( @@ -103,9 +108,7 @@ class CoreMLSamplerLCM(CoreMLSampler): ) if callback: - callback(i, denoised, latents, steps) - - denoised = denoised.to(get_torch_device()) + callback(i, denoised.float(), latents, steps) return ({"samples": denoised / 0.1825},) diff --git a/coreml_suite/lcm/nodes.py b/coreml_suite/lcm/nodes.py index 7e84339..e320d6a 100644 --- a/coreml_suite/lcm/nodes.py +++ b/coreml_suite/lcm/nodes.py @@ -24,7 +24,7 @@ class CoreMLConverterLCM: ComputeUnit.CPU_ONLY.name, ], ), - # "controlnet_support": ("BOOLEAN", {"default": False}), + "controlnet_support": ("BOOLEAN", {"default": False}), } } @@ -32,9 +32,7 @@ class CoreMLConverterLCM: RETURN_NAMES = ("coreml_model",) FUNCTION = "convert" - def convert( - self, height, width, batch_size, compute_unit, controlnet_support=False - ): + def convert(self, height, width, batch_size, compute_unit, controlnet_support): """Converts a LCM model to Core ML. Args: diff --git a/coreml_suite/nodes.py b/coreml_suite/nodes.py index 899523e..b224f9c 100644 --- a/coreml_suite/nodes.py +++ b/coreml_suite/nodes.py @@ -51,7 +51,8 @@ class CoreMLSampler(KSampler): if latent_image is None: logger.warning("No latent image provided, using empty tensor.") expected = coreml_model.expected_inputs["sample"]["shape"] - latent_image = {"samples": torch.zeros(expected[0] // 2, *expected[1:])} + batch_size = max(expected[0] // 2, 1) + latent_image = {"samples": torch.zeros(batch_size, *expected[1:])} return super().sample( model_patcher,