From 1bc728d0ea7f228465e7cb07332eed2d937abb1c Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Sat, 11 Nov 2023 03:11:35 +0100 Subject: [PATCH] Core ML Sampler supports LCM --- __init__.py | 3 - coreml_suite/lcm/__init__.py | 4 +- coreml_suite/lcm/nodes.py | 139 ++++++++-------- coreml_suite/lcm/sampler.py | 299 +++++++++++++++++++---------------- coreml_suite/nodes.py | 39 +++-- 5 files changed, 263 insertions(+), 221 deletions(-) diff --git a/__init__.py b/__init__.py index 3783b87..4466edb 100644 --- a/__init__.py +++ b/__init__.py @@ -5,7 +5,6 @@ sys.path.append(os.path.dirname(__file__)) from coreml_suite.nodes import CoreMLLoaderUNet, CoreMLSampler, CoreMLModelAdapter from coreml_suite.lcm import ( - COREML_SAMPLER_LCM, COREML_CONVERT_LCM, ) @@ -13,13 +12,11 @@ NODE_CLASS_MAPPINGS = { "CoreMLUNetLoader": CoreMLLoaderUNet, "CoreMLSampler": CoreMLSampler, "CoreMLModelAdapter": CoreMLModelAdapter, - "Core ML LCM Sampler": COREML_SAMPLER_LCM, "Core ML LCM Converter": COREML_CONVERT_LCM, } 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 Converter": "Convert LCM to Core ML", } diff --git a/coreml_suite/lcm/__init__.py b/coreml_suite/lcm/__init__.py index f0829fd..4432285 100644 --- a/coreml_suite/lcm/__init__.py +++ b/coreml_suite/lcm/__init__.py @@ -1,3 +1,3 @@ -from .nodes import COREML_SAMPLER_LCM, COREML_CONVERT_LCM +from .nodes import COREML_CONVERT_LCM -__all__ = ["COREML_SAMPLER_LCM", "COREML_CONVERT_LCM"] +__all__ = ["COREML_CONVERT_LCM"] diff --git a/coreml_suite/lcm/nodes.py b/coreml_suite/lcm/nodes.py index 1387717..5d320ef 100644 --- a/coreml_suite/lcm/nodes.py +++ b/coreml_suite/lcm/nodes.py @@ -2,20 +2,11 @@ import os import torch from coremltools import ComputeUnit -from diffusers import LCMScheduler from python_coreml_stable_diffusion.coreml_model import CoreMLModel -import comfy.utils -import latent_preview -from comfy import model_base from comfy.model_management import get_torch_device -from comfy.model_patcher import ModelPatcher -from coreml_suite.config import get_model_config +from comfy_extras.nodes_model_advanced import ModelSamplingDiscreteLCM, LCM from coreml_suite.lcm import converter as lcm_converter -from coreml_suite.lcm.sampler import CoreMLSamplerLCM -from coreml_suite.logger import logger -from coreml_suite.models import CoreMLModelWrapper -from coreml_suite.nodes import CoreMLSampler class COREML_CONVERT_LCM: @@ -81,72 +72,76 @@ class COREML_CONVERT_LCM: return (CoreMLModel(target_path, compute_unit, "compiled"),) -class COREML_SAMPLER_LCM(CoreMLSampler): - @classmethod - def INPUT_TYPES(s): - old_required = CoreMLSampler.INPUT_TYPES()["required"].copy() - old_required["steps"][1]["default"] = 4 - 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",)}, - } +def get_w_embedding(w, embedding_dim=512, dtype=torch.float32): + """ + see https://github.com/google-research/vdm/blob/dc27b98a554f65cdc654b800da5aa1846545d41b/model_vdm.py#L298 + Args: + timesteps: torch.Tensor: generate embedding vectors at these timesteps + embedding_dim: int: dimension of the embeddings to generate + dtype: data type of the generated embeddings - CATEGORY = "Core ML Suite" + Returns: + embedding vectors with shape `(len(timesteps), embedding_dim)` + """ + assert len(w.shape) == 1 + w = w * 1000.0 - def sample( - self, - coreml_model, - seed, - steps, - cfg, - positive, - latent_image=None, - denoise=1.0, - **kwargs, - ): - scheduler = LCMScheduler.from_pretrained( - "SimianLuo/LCM_Dreamshaper_v7", subfolder="scheduler" - ) + half_dim = embedding_dim // 2 + emb = torch.log(torch.tensor(10000.0)) / (half_dim - 1) + emb = torch.exp(torch.arange(half_dim, dtype=dtype) * -emb) + emb = w.to(dtype)[:, None] * emb[None, :] + emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1) + if embedding_dim % 2 == 1: # zero pad + emb = torch.nn.functional.pad(emb, (0, 1)) + assert emb.shape == (w.shape[0], embedding_dim) + return emb - model_config = get_model_config() - wrapped_model = CoreMLModelWrapper(coreml_model) - model = model_base.BaseModel(model_config, device=get_torch_device()) - model.diffusion_model = wrapped_model - model_patcher = ModelPatcher(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())} +def model_function_wrapper(w_embedding): + def wrapper(model_function, params): + x = params["input"] + t = params["timestep"] + c = params["c"] - x0_output = {} - callback = latent_preview.prepare_callback(model_patcher, steps, x0_output) - disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED + context = c.get("c_crossattn") - sampler = CoreMLSamplerLCM(scheduler) - samples = sampler.sample( - model_patcher, - seed, - steps, - cfg, - positive, - latent_image=latent_image, - denoise=denoise, - callback=callback, - disable_pbar=disable_pbar, - ) + if context is None: + return torch.zeros_like(x) - out = latent_image.copy() - out["samples"] = samples - if "x0" in x0_output: - out_denoised = latent_image.copy() - out_denoised["samples"] = model.process_latent_out( - x0_output["x0"].to(get_torch_device()) - ) - else: - out_denoised = out - return out, out_denoised + return model_function(x, t, **c, timestep_cond=w_embedding) + + return wrapper + + +def lcm_patch(model): + m = model.clone() + sampling_type = LCM + sampling_base = ModelSamplingDiscreteLCM + + class ModelSamplingAdvanced(sampling_base, sampling_type): + pass + + model_sampling = ModelSamplingAdvanced() + m.add_object_patch("model_sampling", model_sampling) + + return m + + +def add_lcm_model_options(model_patcher, cfg, latent_image): + mp = model_patcher.clone() + + latent = latent_image["samples"].to(get_torch_device()) + batch_size = latent.shape[0] + dtype = latent.dtype + device = get_torch_device() + + w = torch.tensor(cfg).repeat(batch_size) + w_embedding = get_w_embedding(w, embedding_dim=256).to(device=device, dtype=dtype) + + model_options = { + "model_function_wrapper": model_function_wrapper(w_embedding), + "sampler_cfg_function": lambda x: x["cond"].to(device), + } + mp.model_options |= model_options + + return mp diff --git a/coreml_suite/lcm/sampler.py b/coreml_suite/lcm/sampler.py index 1f3355a..c43de97 100644 --- a/coreml_suite/lcm/sampler.py +++ b/coreml_suite/lcm/sampler.py @@ -1,136 +1,167 @@ -import numpy as np -import torch +# import numpy as np +# import torch +# +# from comfy.model_management import get_torch_device +# from comfy_extras.nodes_model_advanced import rescale_zero_terminal_snr_sigmas, \ +# ModelSamplingDiscreteLCM, LCM +# from nodes import KSampler +# +# from coreml_suite.nodes import CoreMLSampler +# +# class CoreMLSamplerLCM(CoreMLSampler): +# def sample( +# self, +# model_patcher, +# seed, +# steps, +# cfg, +# positive, +# latent_image, +# denoise=1.0, +# callback=None, +# disable_pbar=False, +# **kwargs +# ): +# positive[0][1]["control_apply_to_uncond"] = False +# +# latent = latent_image["samples"].to(get_torch_device()) +# +# batch_size = latent.shape[0] +# dtype = latent.dtype +# device = get_torch_device() +# +# w = torch.tensor(cfg).repeat(batch_size) +# w_embedding = self.get_w_embedding(w, embedding_dim=256).to( +# device=device, dtype=dtype +# ) +# +# model_options = { +# "model_function_wrapper": model_function_wrapper(w_embedding), +# "sampler_cfg_function": lambda x: x["cond"].to(device), +# } +# model_patcher.model_options |= model_options -from comfy import samplers -from comfy.model_management import get_torch_device -from comfy.sample import sample_custom, prepare_noise +# self.prepare_timesteps(denoise, device, 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) +# } +# model_patcher.model.model_sampling.timestep = lambda x: sigma_to_timestep[ +# x[0].item() +# ].expand(1) +# +# noise_mask = latent_image.get("noise_mask") +# batch_inds = latent_image.get("batch_index") +# noise = prepare_noise(latent, seed, batch_inds) +# +# sampler = samplers.ksampler("ddpm")() +# samples = sample_custom( +# model_patcher, +# noise, +# cfg, +# sampler, +# sigmas, +# positive, +# (), +# latent, +# noise_mask, +# callback, +# disable_pbar, +# seed, +# ) - -class CoreMLSamplerLCM: - def __init__(self, scheduler): - self.scheduler = scheduler - - def sample( - self, - model_patcher, - seed, - steps, - cfg, - positive, - latent_image, - denoise=1.0, - callback=None, - disable_pbar=False, - ): - positive[0][1]["control_apply_to_uncond"] = False - - latent = latent_image["samples"].to(get_torch_device()) - - batch_size = latent.shape[0] - dtype = latent.dtype - device = get_torch_device() - - w = torch.tensor(cfg).repeat(batch_size) - w_embedding = self.get_w_embedding(w, embedding_dim=256).to( - device=device, dtype=dtype - ) - - model_options = { - "model_function_wrapper": model_function_wrapper(w_embedding), - "sampler_cfg_function": lambda x: x["cond"].to(device), - } - model_patcher.model_options |= model_options - - self.prepare_timesteps(denoise, device, 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) - } - model_patcher.model.model_sampling.timestep = lambda x: sigma_to_timestep[ - x[0].item() - ].expand(1) - - noise_mask = latent_image.get("noise_mask") - batch_inds = latent_image.get("batch_index") - noise = prepare_noise(latent, seed, batch_inds) - - sampler = samplers.ksampler("ddpm")() - samples = sample_custom( - model_patcher, - noise, - cfg, - sampler, - sigmas, - positive, - (), - latent, - noise_mask, - callback, - disable_pbar, - seed, - ) - return samples - - def get_sigmas(self, steps, denoise): - alphas_cumprod = self.scheduler.alphas_cumprod - 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, 0.0).astype(np.float32) - sigmas = torch.from_numpy(sigmas).to(get_torch_device()) - return (sigmas, torch.from_numpy(s.copy()).to(get_torch_device())) - - def prepare_timesteps(self, denoise, device, steps): - 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 - - def get_w_embedding(self, w, embedding_dim=512, dtype=torch.float32): - """ - see https://github.com/google-research/vdm/blob/dc27b98a554f65cdc654b800da5aa1846545d41b/model_vdm.py#L298 - Args: - timesteps: torch.Tensor: generate embedding vectors at these timesteps - embedding_dim: int: dimension of the embeddings to generate - dtype: data type of the generated embeddings - - Returns: - embedding vectors with shape `(len(timesteps), embedding_dim)` - """ - assert len(w.shape) == 1 - w = w * 1000.0 - - half_dim = embedding_dim // 2 - emb = torch.log(torch.tensor(10000.0)) / (half_dim - 1) - emb = torch.exp(torch.arange(half_dim, dtype=dtype) * -emb) - emb = w.to(dtype)[:, None] * emb[None, :] - emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1) - if embedding_dim % 2 == 1: # zero pad - emb = torch.nn.functional.pad(emb, (0, 1)) - assert emb.shape == (w.shape[0], embedding_dim) - return emb - - -def model_function_wrapper(w_embedding): - def wrapper(model_function, params): - x = params["input"] - t = params["timestep"] - c = params["c"] - - context = c.get("c_crossattn") - - if context is None: - return torch.zeros_like(x) - - return model_function(x, t, **c, timestep_cond=w_embedding) - - return wrapper +# model_patcher = lcm_patch(model_patcher) +# +# return KSampler.sample( +# self, +# model_patcher, +# seed, +# steps, +# cfg, +# "lcm", +# "sgm_uniform", +# positive, +# None, +# latent_image, +# denoise, +# ) +# +# def get_sigmas(self, steps, denoise): +# alphas_cumprod = self.scheduler.alphas_cumprod +# 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, 0.0).astype(np.float32) +# sigmas = torch.from_numpy(sigmas).to(get_torch_device()) +# return (sigmas, torch.from_numpy(s.copy()).to(get_torch_device())) +# +# def prepare_timesteps(self, denoise, device, steps): +# 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 +# +# def get_w_embedding(self, w, embedding_dim=512, dtype=torch.float32): +# """ +# see https://github.com/google-research/vdm/blob/dc27b98a554f65cdc654b800da5aa1846545d41b/model_vdm.py#L298 +# Args: +# timesteps: torch.Tensor: generate embedding vectors at these timesteps +# embedding_dim: int: dimension of the embeddings to generate +# dtype: data type of the generated embeddings +# +# Returns: +# embedding vectors with shape `(len(timesteps), embedding_dim)` +# """ +# assert len(w.shape) == 1 +# w = w * 1000.0 +# +# half_dim = embedding_dim // 2 +# emb = torch.log(torch.tensor(10000.0)) / (half_dim - 1) +# emb = torch.exp(torch.arange(half_dim, dtype=dtype) * -emb) +# emb = w.to(dtype)[:, None] * emb[None, :] +# emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1) +# if embedding_dim % 2 == 1: # zero pad +# emb = torch.nn.functional.pad(emb, (0, 1)) +# assert emb.shape == (w.shape[0], embedding_dim) +# return emb +# +# +# def model_function_wrapper(w_embedding): +# def wrapper(model_function, params): +# x = params["input"] +# t = params["timestep"] +# c = params["c"] +# +# context = c.get("c_crossattn") +# +# if context is None: +# return torch.zeros_like(x) +# +# return model_function(x, t, **c, timestep_cond=w_embedding) +# +# return wrapper +# +# +# def lcm_patch(model): +# m = model.clone() +# sampling_type = LCM +# sampling_base = ModelSamplingDiscreteLCM +# +# class ModelSamplingAdvanced(sampling_base, sampling_type): +# pass +# +# model_sampling = ModelSamplingAdvanced() +# model_sampling.set_sigmas(rescale_zero_terminal_snr_sigmas(model_sampling.sigmas)) +# +# m.add_object_patch("model_sampling", model_sampling) +# +# return m diff --git a/coreml_suite/nodes.py b/coreml_suite/nodes.py index b224f9c..1cde2e9 100644 --- a/coreml_suite/nodes.py +++ b/coreml_suite/nodes.py @@ -8,6 +8,7 @@ import folder_paths from comfy import model_base from comfy.model_management import get_torch_device from comfy.model_patcher import ModelPatcher +from coreml_suite.lcm.nodes import add_lcm_model_options, lcm_patch from coreml_suite.logger import logger from nodes import KSampler @@ -42,17 +43,13 @@ class CoreMLSampler(KSampler): latent_image=None, denoise=1.0, ): - model_config = get_model_config() - wrapped_model = CoreMLModelWrapper(coreml_model) - model = model_base.BaseModel(model_config, device=get_torch_device()) - model.diffusion_model = wrapped_model - model_patcher = ModelPatcher(model, get_torch_device(), None) + model_patcher = self.get_model_patcher(coreml_model) + latent_image = self.get_latent_image(coreml_model, latent_image) - if latent_image is None: - logger.warning("No latent image provided, using empty tensor.") - expected = coreml_model.expected_inputs["sample"]["shape"] - batch_size = max(expected[0] // 2, 1) - latent_image = {"samples": torch.zeros(batch_size, *expected[1:])} + if is_lcm(coreml_model): + positive[0][1]["control_apply_to_uncond"] = False + model_patcher = add_lcm_model_options(model_patcher, cfg, latent_image) + model_patcher = lcm_patch(model_patcher) return super().sample( model_patcher, @@ -67,6 +64,24 @@ class CoreMLSampler(KSampler): denoise, ) + def get_latent_image(self, coreml_model, latent_image): + if latent_image is not None: + return latent_image + + logger.warning("No latent image provided, using empty tensor.") + expected = coreml_model.expected_inputs["sample"]["shape"] + batch_size = max(expected[0] // 2, 1) + latent_image = {"samples": torch.zeros(batch_size, *expected[1:])} + return latent_image + + def get_model_patcher(self, coreml_model): + model_config = get_model_config() + wrapped_model = CoreMLModelWrapper(coreml_model) + model = model_base.BaseModel(model_config, device=get_torch_device()) + model.diffusion_model = wrapped_model + model_patcher = ModelPatcher(model, get_torch_device(), None) + return model_patcher + class CoreMLLoader: PACKAGE_DIRNAME = "" @@ -140,3 +155,7 @@ class CoreMLModelAdapter: model.diffusion_model = wrapped_model model_patcher = ModelPatcher(model, get_torch_device(), None) return (model_patcher,) + + +def is_lcm(coreml_model): + return "timestep_cond" in coreml_model.expected_inputs