From c09bbeabe2f8099b72d9734360df6c7875b7412b Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Wed, 1 Nov 2023 21:46:24 +0100 Subject: [PATCH] Add more advanced LCM Sampler --- __init__.py | 8 ++- coreml_suite/lcm/__init__.py | 4 +- coreml_suite/lcm/lcm_sampler.py | 122 ++++++++++++++++++++++++++++++++ coreml_suite/models.py | 3 +- 4 files changed, 132 insertions(+), 5 deletions(-) diff --git a/__init__.py b/__init__.py index 8410c43..6897878 100644 --- a/__init__.py +++ b/__init__.py @@ -4,12 +4,17 @@ import sys sys.path.append(os.path.dirname(__file__)) from coreml_suite.nodes import CoreMLLoaderUNet, CoreMLSampler, CoreMLModelAdapter -from coreml_suite.lcm import CoreMLConverterLCM, CoreMLSamplerLCM_Simple +from coreml_suite.lcm import ( + CoreMLSamplerLCM, + CoreMLConverterLCM, + CoreMLSamplerLCM_Simple, +) NODE_CLASS_MAPPINGS = { "CoreMLUNetLoader": CoreMLLoaderUNet, "CoreMLSampler": CoreMLSampler, "CoreMLModelAdapter": CoreMLModelAdapter, + "Core ML LCM Sampler": CoreMLSamplerLCM, "Core ML LCM Sampler (Simple)": CoreMLSamplerLCM_Simple, "CoreMLConverterLCM": CoreMLConverterLCM, } @@ -17,6 +22,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "CoreMLUNetLoader": "Load Core ML UNet", "CoreMLSampler": "Core ML Sampler", "CoreMLModelAdapter": "Core ML Adapter (Experimental)", + "Core ML LCM Sampler": "Core ML LCM Sampler", "Core ML LCM Sampler (Simple)": "Core ML LCM Sampler (Simple)", "CoreMLConverterLCM": "Convert LCM to Core ML", } diff --git a/coreml_suite/lcm/__init__.py b/coreml_suite/lcm/__init__.py index 2df93e0..3886985 100644 --- a/coreml_suite/lcm/__init__.py +++ b/coreml_suite/lcm/__init__.py @@ -1,4 +1,4 @@ -from .lcm_sampler import CoreMLSamplerLCM_Simple +from .lcm_sampler import CoreMLSamplerLCM, CoreMLSamplerLCM_Simple from .nodes import CoreMLConverterLCM -__all__ = ["CoreMLSamplerLCM_Simple", "CoreMLConverterLCM"] +__all__ = ["CoreMLSamplerLCM", "CoreMLSamplerLCM_Simple", "CoreMLConverterLCM"] diff --git a/coreml_suite/lcm/lcm_sampler.py b/coreml_suite/lcm/lcm_sampler.py index d9b18d7..d9caece 100644 --- a/coreml_suite/lcm/lcm_sampler.py +++ b/coreml_suite/lcm/lcm_sampler.py @@ -1,11 +1,17 @@ import os +import numpy as np import torch +import comfy.utils +import latent_preview from comfy.model_management import get_torch_device +from comfy.model_patcher import ModelPatcher from coreml_suite.lcm.lcm_pipeline import LatentConsistencyModelPipeline from coreml_suite.lcm.lcm_scheduler import LCMScheduler +from coreml_suite.logger import logger from coreml_suite.models import get_model_config, CoreMLModelWrapperLCM +from coreml_suite.nodes import CoreMLSampler class CoreMLSamplerLCM_Simple: @@ -86,3 +92,119 @@ class CoreMLSamplerLCM_Simple: images_tensor = torch.from_numpy(result) return (images_tensor,) + + +class CoreMLSamplerLCM(CoreMLSampler): + @classmethod + def INPUT_TYPES(s): + old_required = CoreMLSampler.INPUT_TYPES()["required"].copy() + old_required.pop("negative") + old_required.pop("sampler_name") + old_required.pop("scheduler") + new_required = {"coreml_model": ("COREML_UNET",)} + return { + "required": new_required | old_required, + "optional": {"latent_image": ("LATENT",)}, + } + + CATEGORY = "Core ML Suite" + + def __init__(self): + self.scheduler = LCMScheduler.from_pretrained( + os.path.join(os.path.dirname(__file__), "scheduler_config.json") + ) + + def sample( + self, + coreml_model, + seed, + steps, + cfg, + positive, + latent_image=None, + denoise=1.0, + **kwargs, + ): + model_config = get_model_config() + wrapped_model = CoreMLModelWrapperLCM(model_config, coreml_model) + patched_model = ModelPatcher(wrapped_model, get_torch_device(), None) + + 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).to(get_torch_device())} + + positive = positive[0][0] + + torch.manual_seed(seed) + + return self._sample(patched_model, steps, cfg, positive, latent_image, denoise) + + def _sample(self, model, steps, cfg, positive, latent_image, denoise): + batch_size = latent_image["samples"].shape[0] + + 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) + + device = get_torch_device() + # callback = latent_preview.prepare_callback(model, steps, None) + + # Prepare timesteps + lcm_origin_steps = 50 + self.scheduler.num_inference_steps = steps + c = self.scheduler.config.num_train_timesteps // lcm_origin_steps + lcm_origin_timesteps = ( + np.asarray(list(range(1, int(lcm_origin_steps * denoise) + 1))) * c - 1 + ) + skipping_step = len(lcm_origin_timesteps) // steps + timesteps = lcm_origin_timesteps[::-skipping_step][:steps] + timesteps = torch.from_numpy(timesteps.copy()).to(device) + self.scheduler.timesteps = timesteps + + timesteps = self.scheduler.timesteps + + # Prepare latent variable + latents = self.prepare_latents(latent_image, device) + + # LCM MultiStep Sampling Loop: + progress_bar = comfy.utils.ProgressBar(total=steps) + for i, t in enumerate(timesteps): + ts = torch.full((batch_size,), t, device=device, dtype=torch.float16) + + # model prediction (v-prediction, eps, x) + model_pred = model.model( + latents, + ts, + encoder_hidden_states=prompt_embeds, + )[0] + + # model_pred *= cfg + + # compute the previous noisy sample x_t -> x_t-1 + latents, denoised = self.scheduler.step( + model_pred, i, t, latents, return_dict=False + ) + + # # call the callback, if provided + # if i == len(timesteps) - 1: + # callback(i, t, latents, steps) + + denoised = denoised.to(get_torch_device()) + + return ({"samples": denoised / 0.1825},) + + def prepare_latents(self, latent_image, device): + latent = latent_image["samples"] + 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] + noise = torch.randn(latent.shape, dtype=torch.float16).cpu() + latent_timestep = self.scheduler.timesteps[:1].repeat(batch_size) + latents = self.scheduler.add_noise(latent, noise, latent_timestep) + + return latents diff --git a/coreml_suite/models.py b/coreml_suite/models.py index a7d4b32..ae4067e 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -48,8 +48,7 @@ class CoreMLModelWrapper(BaseModel): merged_out = merge_chunks(chunked_out, x.shape) return merged_out - def _apply_model(self, x, t, c_concat=None, c_crossattn=None, c_adm=None, - control=None, transformer_options={}): + def _apply_model(self, x, t, c_crossattn, control=None): model_input_kwargs = self.prepare_inputs(x, t, c_crossattn, control) np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"]