Core ML Sampler supports LCM
This commit is contained in:
@@ -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"]
|
||||
|
||||
+67
-72
@@ -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
|
||||
|
||||
+165
-134
@@ -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
|
||||
|
||||
+29
-10
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user