Refactor LCM sampling

This commit is contained in:
aszc-dev
2024-06-28 15:52:54 +02:00
parent 31a6ac6d2f
commit d9be9c13e2
3 changed files with 188 additions and 24 deletions
+29
View File
@@ -34,3 +34,32 @@ def get_model_config():
config = supported_models_base.BASE(get_unet_config())
config.latent_format = latent_formats.SD15()
return config
def unet_config_from_diffusers_unet(state_dict):
match = {}
attention_resolutions = []
attn_res = 1
for i in range(5):
k = "down_blocks.{}.attentions.1.transformer_blocks.0.attn2.to_k.weight".format(
i
)
if k in state_dict:
match["context_dim"] = state_dict[k].shape[1]
attention_resolutions.append(attn_res)
attn_res *= 2
match["attention_resolutions"] = attention_resolutions
match["model_channels"] = state_dict["conv_in.weight"].shape[0]
match["in_channels"] = state_dict["conv_in.weight"].shape[1]
match["adm_in_channels"] = None
if "class_embedding.linear_1.weight" in state_dict:
match["adm_in_channels"] = state_dict["class_embedding.linear_1.weight"].shape[
1
]
elif "add_embedding.linear_1.weight" in state_dict:
match["adm_in_channels"] = state_dict["add_embedding.linear_1.weight"].shape[1]
print(match)
+158 -23
View File
@@ -6,11 +6,17 @@ from diffusers.utils.torch_utils import randn_tensor
from tqdm import tqdm
import latent_preview
from comfy import model_base
from comfy import model_base, sample
from comfy.k_diffusion.sampling import generic_step_sampler
from comfy.model_management import get_torch_device
from comfy.model_patcher import ModelPatcher
from comfy.sample import prepare_sampling
from comfy.samplers import sampling_function
from comfy.sample import prepare_sampling, sample_custom, prepare_noise
from comfy.samplers import (
sampling_function,
pre_run_control,
Sampler,
KSamplerX0Inpaint,
)
from coreml_suite.lcm.lcm_scheduler import LCMScheduler
from coreml_suite.logger import logger
from coreml_suite.models import CoreMLModelWrapper
@@ -36,8 +42,7 @@ class CoreMLSamplerLCM(CoreMLSampler):
def __init__(self):
self.scheduler = LCMScheduler.from_pretrained(
"SimianLuo/LCM_Dreamshaper_v7",
subfolder="scheduler"
"SimianLuo/LCM_Dreamshaper_v7", subfolder="scheduler"
)
def sample(
@@ -62,45 +67,99 @@ class CoreMLSamplerLCM(CoreMLSampler):
expected = coreml_model.expected_inputs["sample"]["shape"]
latent_image = {"samples": torch.zeros(*expected).to(get_torch_device())}
callback = latent_preview.prepare_callback(model_patcher, steps, None)
torch.manual_seed(seed)
latent = latent_image["samples"].to(get_torch_device())
model, positive, _, _, _ = prepare_sampling(
model_patcher, latent_image["samples"].shape, positive, (), None
x0_output = {}
callback = latent_preview.prepare_callback(model_patcher, steps, x0_output)
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
)
return self._sample(
model_patcher, steps, cfg, positive, latent_image, denoise, callback
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
batch_inds = latent_image.get("batch_index")
noise = prepare_noise(latent, seed, batch_inds)
self.prepare_timesteps(denoise, device, steps)
sampler = LCMSampler(self.scheduler)
sigmas = self.get_sigmas()
noise_mask = latent_image.get("noise_mask")
samples = sample_custom(
model_patcher,
noise,
cfg,
sampler,
sigmas,
positive,
(),
latent,
noise_mask,
callback,
)
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(device)
)
else:
out_denoised = out
return (out, out_denoised)
#
# model, positive, _, _, _ = prepare_sampling(
# model_patcher, latent_image["samples"].shape, positive, (), None
# )
#
# pre_run_control(model, positive)
#
# return self._sample(
# model, steps, cfg, positive, latent_image, denoise, callback
# )
def get_sigmas(self):
alphas_cumprod = self.scheduler.alphas_cumprod
return ((1 - alphas_cumprod) / alphas_cumprod) ** 0.5
def _sample(
self, model, steps, cfg, positive, latent_image, denoise, callback=None
):
device = get_torch_device()
batch_size = latent_image["samples"].shape[0]
# prompt_embeds = self.prepare_prompt_embeds(batch_size, positive)
timesteps = self.prepare_timesteps(denoise, device, steps)
latents = self.prepare_latents(latent_image, device)
w = torch.tensor(cfg).repeat(batch_size)
w_embedding = self.get_w_embedding(w, embedding_dim=256).to(
device=device, dtype=latents.dtype
)
# LCM MultiStep Sampling Loop:
iterator = tqdm(timesteps, desc="Core ML LCM Sampler", total=steps)
for i, t in enumerate(iterator):
ts = torch.full((batch_size,), t, device=device, dtype=torch.float16)
model_options = {
"transformer_options": {"timestep_cond": w_embedding},
"sampler_cfg_function": lambda x: x["cond"].to(device),
}
model_pred = sampling_function(
model.apply_model, latents, ts, None, positive, denoise, model_options
model.diffusion_model,
latents,
ts,
None,
positive,
denoise,
model_options,
)
# compute the previous noisy sample x_t -> x_t-1
@@ -176,3 +235,79 @@ class CoreMLSamplerLCM(CoreMLSampler):
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")
control = c.get("control")
return model_function(x, t, **c, timestep_cond=w_embedding)
return wrapper
class LCMSampler(Sampler):
def __init__(self, scheduler):
self.scheduler = scheduler
def sample(
self,
model_wrap,
sigmas,
extra_args,
callback,
noise,
latent_image=None,
denoise_mask=None,
disable_pbar=False,
):
extra_args["denoise_mask"] = denoise_mask
model_k = KSamplerX0Inpaint(model_wrap)
model_k.latent_image = latent_image
model_k.noise = noise
if self.max_denoise(model_wrap, sigmas):
noise = noise * torch.sqrt(1.0 + sigmas[0] ** 2.0)
else:
noise = noise * sigmas[0]
k_callback = None
total_steps = len(sigmas) - 1
if latent_image is not None:
noise += latent_image
samples = self._sample(model_k, noise, extra_args, callback)
return samples
def _sample(self, model, noise, extra_args, callback):
batch_size = noise.shape[0]
model_options = extra_args["model_options"]
positive = extra_args["cond"]
cond_scale = extra_args["cond_scale"]
denoise_mask = extra_args["denoise_mask"]
timesteps = self.scheduler.timesteps
# LCM MultiStep Sampling Loop:
iterator = tqdm(timesteps, desc="Core ML LCM Sampler", total=len(timesteps))
for i, t in enumerate(iterator):
ts = torch.full((batch_size,), t, device=t.device, dtype=torch.float16)
model_pred = model(
noise, ts, None, positive, cond_scale, denoise_mask, model_options
)
# compute the previous noisy sample x_t -> x_t-1
noise, denoised = self.scheduler.step(
model_pred, i, t, noise, return_dict=False
)
if callback:
callback(i, denoised.float(), noise, len(timesteps))
return denoised
+1 -1
View File
@@ -10,7 +10,7 @@ class CoreMLModelWrapper:
self.coreml_model = coreml_model
self.dtype = torch.float16
def __call__(self, x, t, context, control, transformer_options, **kwargs):
def __call__(self, x, t, context, control, transformer_options=None, **kwargs):
inputs = CoreMLInputs(x, t, context, control, **kwargs)
input_list = inputs.chunks(self.expected_inputs)