diff --git a/coreml_suite/lcm/lcm_sampler.py b/coreml_suite/lcm/lcm_sampler.py index e002ec7..c217db1 100644 --- a/coreml_suite/lcm/lcm_sampler.py +++ b/coreml_suite/lcm/lcm_sampler.py @@ -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