Rearrange LCM code

This commit is contained in:
aszc-dev
2023-11-11 04:15:59 +01:00
parent 8bcdeab234
commit 971e60aa09
4 changed files with 74 additions and 250 deletions
-78
View File
@@ -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
-167
View File
@@ -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
+73
View File
@@ -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
+1 -5
View File
@@ -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