Remove dead code from LCM Sampler

This commit is contained in:
aszc-dev
2024-06-28 15:52:54 +02:00
parent 997c6a78ff
commit 8c9fbacb45
+6 -190
View File
@@ -1,18 +1,12 @@
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, samplers
from comfy.model_management import get_torch_device
from comfy.model_patcher import ModelPatcher
from comfy.sample import sample_custom, prepare_noise
from comfy.samplers import (
sampling_function,
Sampler,
KSamplerX0Inpaint,
)
from coreml_suite.lcm.lcm_scheduler import LCMScheduler
from coreml_suite.logger import logger
from coreml_suite.models import CoreMLModelWrapper
@@ -58,6 +52,8 @@ class CoreMLSamplerLCM(CoreMLSampler):
model.diffusion_model = wrapped_model
model_patcher = ModelPatcher(model, get_torch_device(), None)
positive[0][1]["control_apply_to_uncond"] = False
if latent_image is None:
logger.warning("No latent image provided, using empty tensor.")
expected = coreml_model.expected_inputs["sample"]["shape"]
@@ -87,11 +83,9 @@ class CoreMLSamplerLCM(CoreMLSampler):
noise = prepare_noise(latent, seed, batch_inds)
self.prepare_timesteps(denoise, device, steps)
# self.scheduler.set_timesteps(steps, 50, device)
# sampler = LCMSampler(self.scheduler)
sampler = samplers.ksampler("ddpm")()
all_sigmas, sigmas = self.get_sigmas(steps)
all_sigmas, sigmas = self.get_sigmas(steps, denoise)
model_patcher.model.model_sampling.set_sigmas(all_sigmas)
sigma_to_timestep = {
s.item(): t for s, t in zip(sigmas, self.scheduler.timesteps)
@@ -102,8 +96,7 @@ class CoreMLSamplerLCM(CoreMLSampler):
noise_mask = latent_image.get("noise_mask")
positive[0][1]["control_apply_to_uncond"] = False
sampler = samplers.ksampler("ddpm")()
samples = sample_custom(
model_patcher,
noise,
@@ -128,17 +121,6 @@ class CoreMLSamplerLCM(CoreMLSampler):
out_denoised = out
return (out, out_denoised)
#
# model, positive, _, _, _ = prepare_sampling(
# model_patcher, latent_image["samples"].shape, positive, (), None
# )
#
# pre_run_control(model, positive)
#
# return self._sample(
# model, steps, cfg, positive, latent_image, denoise, callback
# )
def get_sigmas(self, steps):
alphas_cumprod = self.scheduler.alphas_cumprod
sigmas = np.asarray(((1 - alphas_cumprod) / alphas_cumprod) ** 0.5)
@@ -149,48 +131,6 @@ class CoreMLSamplerLCM(CoreMLSampler):
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
):
device = get_torch_device()
batch_size = latent_image["samples"].shape[0]
timesteps = self.prepare_timesteps(denoise, device, steps)
latents = self.prepare_latents(latent_image, device)
# LCM MultiStep Sampling Loop:
iterator = tqdm(timesteps, desc="Core ML LCM Sampler", total=steps)
for i, t in enumerate(iterator):
ts = torch.full((batch_size,), t, device=device, dtype=torch.float16)
model_pred = sampling_function(
model.diffusion_model,
latents,
ts,
None,
positive,
denoise,
model_options,
)
# compute the previous noisy sample x_t -> x_t-1
latents, denoised = self.scheduler.step(
model_pred, i, t, latents, return_dict=False
)
if callback:
callback(i, denoised.float(), latents, steps)
return ({"samples": denoised / 0.1825},)
def prepare_prompt_embeds(self, batch_size, positive):
bs_embed, seq_len, _ = positive.shape
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = positive.repeat(1, batch_size, 1)
prompt_embeds = prompt_embeds.view(bs_embed * batch_size, seq_len, -1)
return prompt_embeds
def prepare_timesteps(self, denoise, device, steps):
lcm_origin_steps = 50
self.scheduler.num_inference_steps = steps
@@ -202,27 +142,6 @@ class CoreMLSamplerLCM(CoreMLSampler):
timesteps = lcm_origin_timesteps[::-skipping_step][:steps]
timesteps = torch.from_numpy(timesteps.copy()).to(device)
self.scheduler.timesteps = timesteps
timesteps = self.scheduler.timesteps
return timesteps
def prepare_latents(self, latent_image, device):
latent = latent_image["samples"].to(device) * 0.1825
latent = latent.to(torch.float16)
if not torch.any(latent):
latents = torch.randn(latent.shape, dtype=torch.float16).to(device)
latents *= self.scheduler.init_noise_sigma
return latents
batch_size = latent.shape[0]
burned = randn_tensor(latent.shape, device=device, dtype=torch.float16)
noise = randn_tensor(latent.shape, device=device, dtype=torch.float16)
latent_timestep = self.scheduler.timesteps[:1].repeat(batch_size)
latents = self.scheduler.add_noise(latent, noise, latent_timestep)
return latents
def get_w_embedding(self, w, embedding_dim=512, dtype=torch.float32):
"""
@@ -263,106 +182,3 @@ def model_function_wrapper(w_embedding):
return model_function(x, t, **c, timestep_cond=w_embedding)
return wrapper
class LCMSampler(Sampler):
def __init__(self, scheduler):
self.scheduler = scheduler
def sample(
self,
model_wrap,
sigmas,
extra_args,
callback,
noise,
latent_image=None,
denoise_mask=None,
disable_pbar=False,
):
extra_args["denoise_mask"] = denoise_mask
model_k = KSamplerX0Inpaint(model_wrap)
model_k.latent_image = latent_image
model_k.noise = noise
if self.max_denoise(model_wrap, sigmas):
noise = noise * torch.sqrt(1.0 + sigmas[0] ** 2.0)
else:
noise = noise * sigmas[0]
k_callback = None
total_steps = len(sigmas) - 1
if latent_image is not None:
noise += latent_image
samples = self._sample(model_k, noise, sigmas, extra_args, callback)
return samples
def _sample(self, model, noise, sigmas, extra_args, callback):
batch_size = noise.shape[0]
model_options = extra_args["model_options"]
positive = extra_args["cond"]
cond_scale = extra_args["cond_scale"]
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):
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
)
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))
return denoised