Remove dead code from LCM Sampler
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user