ControlNet works for LCM

This commit is contained in:
aszc-dev
2023-11-10 02:07:11 +01:00
parent 6ccf41e5c9
commit 6864c233e3
+74 -19
View File
@@ -1,19 +1,15 @@
import os
import numpy as np
import torch
from diffusers.utils.torch_utils import randn_tensor
from tqdm import tqdm
import latent_preview
from comfy import model_base, sample
from comfy.k_diffusion.sampling import generic_step_sampler
from comfy import model_base, samplers
from comfy.model_management import get_torch_device
from comfy.model_patcher import ModelPatcher
from comfy.sample import prepare_sampling, sample_custom, prepare_noise
from comfy.sample import sample_custom, prepare_noise
from comfy.samplers import (
sampling_function,
pre_run_control,
Sampler,
KSamplerX0Inpaint,
)
@@ -92,12 +88,22 @@ class CoreMLSamplerLCM(CoreMLSampler):
self.prepare_timesteps(denoise, device, steps)
sampler = LCMSampler(self.scheduler)
# sampler = LCMSampler(self.scheduler)
sampler = samplers.ksampler("ddpm")()
sigmas = self.get_sigmas()
all_sigmas, sigmas = self.get_sigmas(steps)
model_patcher.model.model_sampling.set_sigmas(all_sigmas)
sigma_to_timestep = {
s.item(): t for s, t in zip(sigmas, self.scheduler.timesteps)
}
model_patcher.model.model_sampling.timestep = lambda x: sigma_to_timestep[
x[0].item()
].expand(1)
noise_mask = latent_image.get("noise_mask")
positive[0][1]["control_apply_to_uncond"] = False
samples = sample_custom(
model_patcher,
noise,
@@ -133,9 +139,15 @@ class CoreMLSamplerLCM(CoreMLSampler):
# model, steps, cfg, positive, latent_image, denoise, callback
# )
def get_sigmas(self):
def get_sigmas(self, steps):
alphas_cumprod = self.scheduler.alphas_cumprod
return ((1 - alphas_cumprod) / alphas_cumprod) ** 0.5
sigmas = np.asarray(((1 - alphas_cumprod) / alphas_cumprod) ** 0.5)
skipping_step = len(sigmas) // steps
s = sigmas[::-skipping_step][:steps]
if len(s) == steps:
s = np.append(s, sigmas[0])
sigmas = torch.from_numpy(sigmas).to(get_torch_device())
return (sigmas, torch.from_numpy(s.copy()).to(get_torch_device()))
def _sample(
self, model, steps, cfg, positive, latent_image, denoise, callback=None
@@ -244,7 +256,9 @@ def model_function_wrapper(w_embedding):
c = params["c"]
context = c.get("c_crossattn")
control = c.get("control")
if context is None:
return torch.zeros_like(x)
return model_function(x, t, **c, timestep_cond=w_embedding)
@@ -282,10 +296,10 @@ class LCMSampler(Sampler):
if latent_image is not None:
noise += latent_image
samples = self._sample(model_k, noise, extra_args, callback)
samples = self._sample(model_k, noise, sigmas, extra_args, callback)
return samples
def _sample(self, model, noise, extra_args, callback):
def _sample(self, model, noise, sigmas, extra_args, callback):
batch_size = noise.shape[0]
model_options = extra_args["model_options"]
positive = extra_args["cond"]
@@ -293,20 +307,61 @@ class LCMSampler(Sampler):
denoise_mask = extra_args["denoise_mask"]
timesteps = self.scheduler.timesteps
sample = noise
# LCM MultiStep Sampling Loop:
iterator = tqdm(timesteps, desc="Core ML LCM Sampler", total=len(timesteps))
for i, t in enumerate(iterator):
ts = torch.full((batch_size,), t, device=t.device, dtype=torch.float16)
model_pred = model(
noise, ts, None, positive, cond_scale, denoise_mask, model_options
ss = torch.full(
(batch_size,), sigmas[i], device=t.device, dtype=sigmas.dtype
)
# 1. get previous step value
prev_timeindex = i + 1
if prev_timeindex < len(timesteps):
prev_timestep = timesteps[prev_timeindex]
else:
prev_timestep = t
# 2. compute alphas, betas
alpha_prod_t = self.scheduler.alphas_cumprod[t]
alpha_prod_t_prev = (
self.scheduler.alphas_cumprod[prev_timestep]
if prev_timestep >= 0
else self.scheduler.final_alpha_cumprod
)
# compute the previous noisy sample x_t -> x_t-1
noise, denoised = self.scheduler.step(
model_pred, i, t, noise, return_dict=False
beta_prod_t = 1 - alpha_prod_t
beta_prod_t_prev = 1 - alpha_prod_t_prev
# 3. Get scalings for boundary conditions
c_skip, c_out = self.scheduler.get_scalings_for_boundary_condition_discrete(
t
)
model_output = model(
sample, ss, None, positive, cond_scale, denoise_mask, model_options
)
pred_x0 = (sample - beta_prod_t.sqrt() * model_output) / alpha_prod_t.sqrt()
# 4. Denoise model output using boundary conditions
denoised = c_out * pred_x0 + c_skip * sample
# 5. Sample z ~ N(0, I), For MultiStep Inference
# Noise is not used for one-step sampling.
if len(timesteps) > 1:
noise = torch.randn(sample.shape).to(sample.device)
sample = (
alpha_prod_t_prev.sqrt() * denoised
+ beta_prod_t_prev.sqrt() * noise
)
else:
sample = denoised
# # # compute the previous noisy sample x_t -> x_t-1
# noise, denoised = self.scheduler.step(
# model_pred, i, t, noise, return_dict=False
# )
if callback:
callback(i, denoised.float(), noise, len(timesteps))