Leverage Comfy's mechanisms to enable LCM ControlNet support

This commit is contained in:
aszc-dev
2023-11-09 00:07:30 +01:00
parent fa0735746c
commit 4c438e1ee6
3 changed files with 19 additions and 17 deletions
+15 -12
View File
@@ -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},)
+2 -4
View File
@@ -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:
+2 -1
View File
@@ -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,