From 971e60aa093a0fd0954a602d98a2b92a6045b31a Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Sat, 11 Nov 2023 04:15:59 +0100 Subject: [PATCH] Rearrange LCM code --- coreml_suite/lcm/nodes.py | 78 ----------------- coreml_suite/lcm/sampler.py | 167 ------------------------------------ coreml_suite/lcm/utils.py | 73 ++++++++++++++++ coreml_suite/nodes.py | 6 +- 4 files changed, 74 insertions(+), 250 deletions(-) delete mode 100644 coreml_suite/lcm/sampler.py create mode 100644 coreml_suite/lcm/utils.py diff --git a/coreml_suite/lcm/nodes.py b/coreml_suite/lcm/nodes.py index 5d320ef..2fc5b4b 100644 --- a/coreml_suite/lcm/nodes.py +++ b/coreml_suite/lcm/nodes.py @@ -1,11 +1,8 @@ import os -import torch from coremltools import ComputeUnit from python_coreml_stable_diffusion.coreml_model import CoreMLModel -from comfy.model_management import get_torch_device -from comfy_extras.nodes_model_advanced import ModelSamplingDiscreteLCM, LCM from coreml_suite.lcm import converter as lcm_converter @@ -70,78 +67,3 @@ class COREML_CONVERT_LCM: target_path = lcm_converter.compile_model(out_path=out_path, out_name=out_name) return (CoreMLModel(target_path, compute_unit, "compiled"),) - - -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 - - 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() - 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 deleted file mode 100644 index c43de97..0000000 --- a/coreml_suite/lcm/sampler.py +++ /dev/null @@ -1,167 +0,0 @@ -# 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 - -# 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, -# ) - -# 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/lcm/utils.py b/coreml_suite/lcm/utils.py new file mode 100644 index 0000000..fbd7c5b --- /dev/null +++ b/coreml_suite/lcm/utils.py @@ -0,0 +1,73 @@ +import torch + +from comfy.model_management import get_torch_device +from comfy_extras.nodes_model_advanced import ModelSamplingDiscreteLCM, LCM + + +def is_lcm(coreml_model): + return "timestep_cond" in coreml_model.expected_inputs + + +def get_w_embedding(w, embedding_dim=512, dtype=torch.float32): + 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() + 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/nodes.py b/coreml_suite/nodes.py index 1cde2e9..d11da55 100644 --- a/coreml_suite/nodes.py +++ b/coreml_suite/nodes.py @@ -8,7 +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.lcm.utils import add_lcm_model_options, lcm_patch, is_lcm from coreml_suite.logger import logger from nodes import KSampler @@ -155,7 +155,3 @@ 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