+1
-1
@@ -1,3 +1,3 @@
|
||||
playground/
|
||||
experiments/
|
||||
__pycache__/
|
||||
__pycache__/
|
||||
@@ -4,14 +4,22 @@ import sys
|
||||
sys.path.append(os.path.dirname(__file__))
|
||||
|
||||
from coreml_suite.nodes import CoreMLLoaderUNet, CoreMLSampler, CoreMLModelAdapter
|
||||
from coreml_suite.lcm import (
|
||||
CoreMLSamplerLCM,
|
||||
CoreMLConverterLCM,
|
||||
)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"CoreMLUNetLoader": CoreMLLoaderUNet,
|
||||
"CoreMLSampler": CoreMLSampler,
|
||||
"CoreMLModelAdapter": CoreMLModelAdapter,
|
||||
"Core ML LCM Sampler": CoreMLSamplerLCM,
|
||||
"Core ML LCM Converter": CoreMLConverterLCM,
|
||||
}
|
||||
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",
|
||||
}
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
import torch
|
||||
|
||||
from comfy.model_management import get_torch_device
|
||||
|
||||
|
||||
def chunk_batch(input_tensor, target_shape):
|
||||
if input_tensor.shape == target_shape:
|
||||
@@ -13,7 +11,7 @@ def chunk_batch(input_tensor, target_shape):
|
||||
num_chunks = batch_size // target_batch_size
|
||||
if num_chunks == 0:
|
||||
padding = torch.zeros(target_batch_size - batch_size, *target_shape[1:]).to(
|
||||
get_torch_device()
|
||||
input_tensor.device
|
||||
)
|
||||
return [torch.cat((input_tensor, padding), dim=0)]
|
||||
|
||||
@@ -21,7 +19,7 @@ def chunk_batch(input_tensor, target_shape):
|
||||
if mod != 0:
|
||||
chunks = list(torch.chunk(input_tensor[:-mod], num_chunks))
|
||||
padding = torch.zeros(target_batch_size - mod, *target_shape[1:]).to(
|
||||
get_torch_device()
|
||||
input_tensor.device
|
||||
)
|
||||
padded = torch.cat((input_tensor[-mod:], padding), dim=0)
|
||||
chunks.append(padded)
|
||||
|
||||
@@ -0,0 +1,4 @@
|
||||
from .lcm_sampler import CoreMLSamplerLCM
|
||||
from .nodes import CoreMLConverterLCM
|
||||
|
||||
__all__ = ["CoreMLSamplerLCM", "CoreMLConverterLCM"]
|
||||
@@ -0,0 +1,293 @@
|
||||
import os
|
||||
import shutil
|
||||
import logging
|
||||
import time
|
||||
import gc
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers import UNet2DConditionModel
|
||||
from coreml_suite.lcm.unet import UNet2DConditionModelLCM
|
||||
|
||||
from transformers import CLIPTextModel
|
||||
import coremltools as ct
|
||||
|
||||
from folder_paths import get_folder_paths
|
||||
from coreml_suite.lcm.lcm_scheduler import LCMScheduler
|
||||
|
||||
logging.basicConfig()
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.setLevel(logging.DEBUG)
|
||||
|
||||
MODEL_VERSION = "SimianLuo/LCM_Dreamshaper_v7"
|
||||
MODEL_NAME = MODEL_VERSION.split("/")[-1] + "_4k"
|
||||
|
||||
import python_coreml_stable_diffusion.unet as unet
|
||||
|
||||
unet.ATTENTION_IMPLEMENTATION_IN_EFFECT = unet.AttentionImplementations.SPLIT_EINSUM
|
||||
|
||||
|
||||
def get_unets():
|
||||
ref_unet = UNet2DConditionModel.from_pretrained(
|
||||
MODEL_VERSION,
|
||||
subfolder="unet",
|
||||
device_map=None,
|
||||
low_cpu_mem_usage=False,
|
||||
)
|
||||
|
||||
cml_unet = UNet2DConditionModelLCM.from_config(ref_unet.config).eval()
|
||||
cml_unet.load_state_dict(ref_unet.state_dict(), strict=False)
|
||||
|
||||
return cml_unet, ref_unet
|
||||
|
||||
|
||||
def get_encoder_hidden_states_shape(unet_config, batch_size):
|
||||
text_encoder = CLIPTextModel.from_pretrained(
|
||||
MODEL_VERSION, subfolder="text_encoder"
|
||||
)
|
||||
|
||||
text_token_sequence_length = text_encoder.config.max_position_embeddings
|
||||
hidden_size = (text_encoder.config.hidden_size,)
|
||||
|
||||
encoder_hidden_states_shape = (
|
||||
batch_size,
|
||||
unet_config.cross_attention_dim or hidden_size,
|
||||
1,
|
||||
text_token_sequence_length,
|
||||
)
|
||||
|
||||
return encoder_hidden_states_shape
|
||||
|
||||
|
||||
def get_scheduler():
|
||||
scheduler = LCMScheduler(
|
||||
beta_start=0.00085,
|
||||
beta_end=0.0120,
|
||||
beta_schedule="scaled_linear",
|
||||
prediction_type="epsilon",
|
||||
)
|
||||
scheduler.set_timesteps(50, 50)
|
||||
return scheduler
|
||||
|
||||
|
||||
def get_coreml_inputs(sample_inputs):
|
||||
coreml_sample_unet_inputs = {
|
||||
k: v.numpy().astype(np.float16) for k, v in sample_inputs.items()
|
||||
}
|
||||
return [
|
||||
ct.TensorType(
|
||||
name=k,
|
||||
shape=v.shape,
|
||||
dtype=v.numpy().dtype if isinstance(v, torch.Tensor) else v.dtype,
|
||||
)
|
||||
for k, v in coreml_sample_unet_inputs.items()
|
||||
]
|
||||
|
||||
|
||||
def load_coreml_model(out_path):
|
||||
logger.info(f"Loading model from {out_path}")
|
||||
|
||||
start = time.time()
|
||||
coreml_model = ct.models.MLModel(out_path)
|
||||
logger.info(f"Loading {out_path} took {time.time() - start:.1f} seconds")
|
||||
|
||||
return coreml_model
|
||||
|
||||
|
||||
def convert_to_coreml(
|
||||
submodule_name, torchscript_module, sample_inputs, output_names, out_path
|
||||
):
|
||||
if os.path.exists(out_path):
|
||||
logger.info(f"Skipping export because {out_path} already exists")
|
||||
coreml_model = load_coreml_model(out_path)
|
||||
else:
|
||||
logger.info(f"Converting {submodule_name} to CoreML..")
|
||||
coreml_model = ct.convert(
|
||||
torchscript_module,
|
||||
convert_to="mlprogram",
|
||||
minimum_deployment_target=ct.target.macOS13,
|
||||
inputs=sample_inputs,
|
||||
outputs=[
|
||||
ct.TensorType(name=name, dtype=np.float32) for name in output_names
|
||||
],
|
||||
skip_model_load=True,
|
||||
)
|
||||
|
||||
del torchscript_module
|
||||
gc.collect()
|
||||
|
||||
return coreml_model
|
||||
|
||||
|
||||
def get_out_path(submodule_name, model_name):
|
||||
fname = f"{model_name}_{submodule_name}.mlpackage"
|
||||
unet_path = get_folder_paths(submodule_name)[0]
|
||||
out_path = os.path.join(unet_path, fname)
|
||||
return out_path
|
||||
|
||||
|
||||
def compile_coreml_model(source_model_path, output_dir, final_name):
|
||||
"""Compiles Core ML models using the coremlcompiler utility from Xcode toolchain"""
|
||||
target_path = os.path.join(output_dir, f"{final_name}.mlmodelc")
|
||||
if os.path.exists(target_path):
|
||||
logger.warning(f"Found existing compiled model at {target_path}! Skipping..")
|
||||
return target_path
|
||||
|
||||
logger.info(f"Compiling {source_model_path}")
|
||||
source_model_name = os.path.basename(os.path.splitext(source_model_path)[0])
|
||||
|
||||
os.system(f"xcrun coremlcompiler compile {source_model_path} {output_dir}")
|
||||
compiled_output = os.path.join(output_dir, f"{source_model_name}.mlmodelc")
|
||||
shutil.move(compiled_output, target_path)
|
||||
|
||||
return target_path
|
||||
|
||||
|
||||
def get_sample_input(batch_size, encoder_hidden_states_shape, sample_shape, scheduler):
|
||||
sample_unet_inputs = dict(
|
||||
[
|
||||
("sample", torch.rand(*sample_shape)),
|
||||
(
|
||||
"timestep",
|
||||
torch.tensor([scheduler.timesteps[0].item()] * batch_size).to(
|
||||
torch.float32
|
||||
),
|
||||
),
|
||||
("encoder_hidden_states", torch.rand(*encoder_hidden_states_shape)),
|
||||
("timestep_cond", torch.randn(batch_size, 256).to(torch.float32)),
|
||||
]
|
||||
)
|
||||
return sample_unet_inputs
|
||||
|
||||
|
||||
def get_unet_inputs_spec(sample_unet_inputs):
|
||||
sample_unet_inputs_spec = {
|
||||
k: (v.shape, v.dtype) for k, v in sample_unet_inputs.items()
|
||||
}
|
||||
return sample_unet_inputs_spec
|
||||
|
||||
|
||||
def add_cnet_support(sample_shape, reference_unet):
|
||||
from python_coreml_stable_diffusion.unet import calculate_conv2d_output_shape
|
||||
|
||||
additional_residuals_shapes = []
|
||||
|
||||
batch_size = sample_shape[0]
|
||||
h, w = sample_shape[2:]
|
||||
|
||||
# conv_in
|
||||
out_h, out_w = calculate_conv2d_output_shape(
|
||||
h,
|
||||
w,
|
||||
reference_unet.conv_in,
|
||||
)
|
||||
additional_residuals_shapes.append(
|
||||
(batch_size, reference_unet.conv_in.out_channels, out_h, out_w)
|
||||
)
|
||||
|
||||
# down_blocks
|
||||
for down_block in reference_unet.down_blocks:
|
||||
additional_residuals_shapes += [
|
||||
(batch_size, resnet.out_channels, out_h, out_w)
|
||||
for resnet in down_block.resnets
|
||||
]
|
||||
if hasattr(down_block, "downsamplers") and down_block.downsamplers is not None:
|
||||
for downsampler in down_block.downsamplers:
|
||||
out_h, out_w = calculate_conv2d_output_shape(
|
||||
out_h, out_w, downsampler.conv
|
||||
)
|
||||
additional_residuals_shapes.append(
|
||||
(
|
||||
batch_size,
|
||||
down_block.downsamplers[-1].conv.out_channels,
|
||||
out_h,
|
||||
out_w,
|
||||
)
|
||||
)
|
||||
|
||||
# mid_block
|
||||
additional_residuals_shapes.append(
|
||||
(batch_size, reference_unet.mid_block.resnets[-1].out_channels, out_h, out_w)
|
||||
)
|
||||
|
||||
additional_inputs = {}
|
||||
for i, shape in enumerate(additional_residuals_shapes):
|
||||
sample_residual_input = torch.rand(*shape)
|
||||
additional_inputs[f"additional_residual_{i}"] = sample_residual_input
|
||||
|
||||
return additional_inputs
|
||||
|
||||
|
||||
def convert(
|
||||
out_path: str,
|
||||
batch_size: int = 1,
|
||||
sample_size: tuple[int, int] = (64, 64),
|
||||
controlnet_support: bool = False,
|
||||
):
|
||||
coreml_unet, ref_unet = get_unets()
|
||||
|
||||
sample_shape = (
|
||||
batch_size, # B
|
||||
ref_unet.config.in_channels, # C
|
||||
sample_size[0], # H
|
||||
sample_size[1], # W
|
||||
)
|
||||
|
||||
encoder_hidden_states_shape = get_encoder_hidden_states_shape(
|
||||
ref_unet.config, batch_size
|
||||
)
|
||||
|
||||
scheduler = get_scheduler()
|
||||
|
||||
sample_inputs = get_sample_input(
|
||||
batch_size, encoder_hidden_states_shape, sample_shape, scheduler
|
||||
)
|
||||
|
||||
if controlnet_support:
|
||||
sample_inputs |= add_cnet_support(sample_shape, ref_unet)
|
||||
|
||||
sample_inputs_spec = get_unet_inputs_spec(sample_inputs)
|
||||
|
||||
logger.info(f"Sample UNet inputs spec: {sample_inputs_spec}")
|
||||
logger.info("JIT tracing..")
|
||||
traced_unet = torch.jit.trace(
|
||||
coreml_unet, example_inputs=list(sample_inputs.values())
|
||||
)
|
||||
logger.info("Done.")
|
||||
|
||||
coreml_sample_inputs = get_coreml_inputs(sample_inputs)
|
||||
|
||||
coreml_unet = convert_to_coreml(
|
||||
"unet", traced_unet, coreml_sample_inputs, ["noise_pred"], out_path
|
||||
)
|
||||
|
||||
del traced_unet
|
||||
gc.collect()
|
||||
|
||||
coreml_unet.save(out_path)
|
||||
logger.info(f"Saved unet into {out_path}")
|
||||
|
||||
|
||||
def compile_model(out_path, out_name):
|
||||
# Compile the model
|
||||
target_path = compile_coreml_model(
|
||||
out_path, get_folder_paths("unet")[0], f"{out_name}_unet"
|
||||
)
|
||||
logger.info(f"Compiled {out_path} to {target_path}")
|
||||
return target_path
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
h = 512
|
||||
w = 512
|
||||
sample_size = (h // 8, w // 8)
|
||||
batch_size = 4
|
||||
|
||||
cn_support_str = "_cn" if True else ""
|
||||
|
||||
out_name = f"{MODEL_NAME}_{batch_size}x{w}x{h}{cn_support_str}"
|
||||
|
||||
out_path = get_out_path("unet", f"{out_name}")
|
||||
if not os.path.exists(out_path):
|
||||
convert(out_path=out_path, sample_size=sample_size, batch_size=batch_size)
|
||||
compile_model(out_path=out_path, out_name=out_name)
|
||||
@@ -0,0 +1,170 @@
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from tqdm import tqdm
|
||||
|
||||
import latent_preview
|
||||
from comfy.model_management import get_torch_device
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
from coreml_suite.lcm.lcm_scheduler import LCMScheduler
|
||||
from coreml_suite.logger import logger
|
||||
from coreml_suite.models import get_model_config, CoreMLModelWrapperLCM
|
||||
from coreml_suite.nodes import CoreMLSampler
|
||||
|
||||
|
||||
class CoreMLSamplerLCM(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",)},
|
||||
}
|
||||
|
||||
CATEGORY = "Core ML Suite"
|
||||
|
||||
def __init__(self):
|
||||
self.scheduler = LCMScheduler.from_pretrained(
|
||||
os.path.join(os.path.dirname(__file__), "scheduler_config.json")
|
||||
)
|
||||
|
||||
def sample(
|
||||
self,
|
||||
coreml_model,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
positive,
|
||||
latent_image=None,
|
||||
denoise=1.0,
|
||||
**kwargs,
|
||||
):
|
||||
model_config = get_model_config()
|
||||
wrapped_model = CoreMLModelWrapperLCM(model_config, coreml_model)
|
||||
patched_model = ModelPatcher(wrapped_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())}
|
||||
|
||||
positive = positive[0][0]
|
||||
|
||||
callback = latent_preview.prepare_callback(patched_model, steps, None)
|
||||
torch.manual_seed(seed)
|
||||
|
||||
return self._sample(
|
||||
patched_model, steps, cfg, positive, latent_image, denoise, callback
|
||||
)
|
||||
|
||||
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_pred = model.model(
|
||||
latents,
|
||||
ts,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
timestep_cond=w_embedding,
|
||||
)[0]
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents, denoised = self.scheduler.step(
|
||||
model_pred, i, t, latents, return_dict=False
|
||||
)
|
||||
|
||||
if callback:
|
||||
callback(i, denoised, latents, steps)
|
||||
|
||||
denoised = denoised.to(get_torch_device())
|
||||
|
||||
return ({"samples": denoised / 0.1825},)
|
||||
|
||||
def prepare_prompt_embeds(self, batch_size, positive):
|
||||
bs_embed, seq_len, _ = positive.shape
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
prompt_embeds = positive.repeat(1, batch_size, 1)
|
||||
prompt_embeds = prompt_embeds.view(bs_embed * batch_size, seq_len, -1)
|
||||
return prompt_embeds
|
||||
|
||||
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
|
||||
timesteps = self.scheduler.timesteps
|
||||
return timesteps
|
||||
|
||||
def prepare_latents(self, latent_image, device):
|
||||
latent = latent_image["samples"].to(device) * 0.1825
|
||||
latent = latent.to(torch.float16)
|
||||
|
||||
if not torch.any(latent):
|
||||
latents = torch.randn(latent.shape, dtype=torch.float16).to(device)
|
||||
latents *= self.scheduler.init_noise_sigma
|
||||
return latents
|
||||
|
||||
batch_size = latent.shape[0]
|
||||
|
||||
burned = randn_tensor(latent.shape, device=device, dtype=torch.float16)
|
||||
noise = randn_tensor(latent.shape, device=device, dtype=torch.float16)
|
||||
|
||||
latent_timestep = self.scheduler.timesteps[:1].repeat(batch_size)
|
||||
latents = self.scheduler.add_noise(latent, noise, latent_timestep)
|
||||
|
||||
return latents
|
||||
|
||||
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
|
||||
@@ -0,0 +1,524 @@
|
||||
# Copyright 2023 Stanford University Team and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
# DISCLAIMER: This code is strongly influenced by https://github.com/pesser/pytorch_diffusion
|
||||
# and https://github.com/hojonathanho/diffusion
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from diffusers import ConfigMixin, SchedulerMixin
|
||||
from diffusers.configuration_utils import register_to_config
|
||||
from diffusers.utils import BaseOutput
|
||||
|
||||
|
||||
@dataclass
|
||||
# Copied from diffusers.schedulers.scheduling_ddpm.DDPMSchedulerOutput with DDPM->DDIM
|
||||
class LCMSchedulerOutput(BaseOutput):
|
||||
"""
|
||||
Output class for the scheduler's `step` function output.
|
||||
|
||||
Args:
|
||||
prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
|
||||
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the
|
||||
denoising loop.
|
||||
pred_original_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
|
||||
The predicted denoised sample `(x_{0})` based on the model output from the current timestep.
|
||||
`pred_original_sample` can be used to preview progress or for guidance.
|
||||
"""
|
||||
|
||||
prev_sample: torch.FloatTensor
|
||||
denoised: Optional[torch.FloatTensor] = None
|
||||
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_ddpm.betas_for_alpha_bar
|
||||
def betas_for_alpha_bar(
|
||||
num_diffusion_timesteps,
|
||||
max_beta=0.999,
|
||||
alpha_transform_type="cosine",
|
||||
):
|
||||
"""
|
||||
Create a beta schedule that discretizes the given alpha_t_bar function, which defines the cumulative product of
|
||||
(1-beta) over time from t = [0,1].
|
||||
|
||||
Contains a function alpha_bar that takes an argument t and transforms it to the cumulative product of (1-beta) up
|
||||
to that part of the diffusion process.
|
||||
|
||||
|
||||
Args:
|
||||
num_diffusion_timesteps (`int`): the number of betas to produce.
|
||||
max_beta (`float`): the maximum beta to use; use values lower than 1 to
|
||||
prevent singularities.
|
||||
alpha_transform_type (`str`, *optional*, default to `cosine`): the type of noise schedule for alpha_bar.
|
||||
Choose from `cosine` or `exp`
|
||||
|
||||
Returns:
|
||||
betas (`np.ndarray`): the betas used by the scheduler to step the model outputs
|
||||
"""
|
||||
if alpha_transform_type == "cosine":
|
||||
|
||||
def alpha_bar_fn(t):
|
||||
return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
|
||||
|
||||
elif alpha_transform_type == "exp":
|
||||
|
||||
def alpha_bar_fn(t):
|
||||
return math.exp(t * -12.0)
|
||||
|
||||
else:
|
||||
raise ValueError(f"Unsupported alpha_tranform_type: {alpha_transform_type}")
|
||||
|
||||
betas = []
|
||||
for i in range(num_diffusion_timesteps):
|
||||
t1 = i / num_diffusion_timesteps
|
||||
t2 = (i + 1) / num_diffusion_timesteps
|
||||
betas.append(min(1 - alpha_bar_fn(t2) / alpha_bar_fn(t1), max_beta))
|
||||
return torch.tensor(betas, dtype=torch.float32)
|
||||
|
||||
|
||||
def rescale_zero_terminal_snr(betas):
|
||||
"""
|
||||
Rescales betas to have zero terminal SNR Based on https://arxiv.org/pdf/2305.08891.pdf (Algorithm 1)
|
||||
|
||||
|
||||
Args:
|
||||
betas (`torch.FloatTensor`):
|
||||
the betas that the scheduler is being initialized with.
|
||||
|
||||
Returns:
|
||||
`torch.FloatTensor`: rescaled betas with zero terminal SNR
|
||||
"""
|
||||
# Convert betas to alphas_bar_sqrt
|
||||
alphas = 1.0 - betas
|
||||
alphas_cumprod = torch.cumprod(alphas, dim=0)
|
||||
alphas_bar_sqrt = alphas_cumprod.sqrt()
|
||||
|
||||
# Store old values.
|
||||
alphas_bar_sqrt_0 = alphas_bar_sqrt[0].clone()
|
||||
alphas_bar_sqrt_T = alphas_bar_sqrt[-1].clone()
|
||||
|
||||
# Shift so the last timestep is zero.
|
||||
alphas_bar_sqrt -= alphas_bar_sqrt_T
|
||||
|
||||
# Scale so the first timestep is back to the old value.
|
||||
alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T)
|
||||
|
||||
# Convert alphas_bar_sqrt to betas
|
||||
alphas_bar = alphas_bar_sqrt**2 # Revert sqrt
|
||||
alphas = alphas_bar[1:] / alphas_bar[:-1] # Revert cumprod
|
||||
alphas = torch.cat([alphas_bar[0:1], alphas])
|
||||
betas = 1 - alphas
|
||||
|
||||
return betas
|
||||
|
||||
|
||||
class LCMScheduler(SchedulerMixin, ConfigMixin):
|
||||
"""
|
||||
`LCMScheduler` extends the denoising procedure introduced in denoising diffusion probabilistic models (DDPMs) with
|
||||
non-Markovian guidance.
|
||||
|
||||
This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
|
||||
methods the library implements for all schedulers such as loading and saving.
|
||||
|
||||
Args:
|
||||
num_train_timesteps (`int`, defaults to 1000):
|
||||
The number of diffusion steps to train the model.
|
||||
beta_start (`float`, defaults to 0.0001):
|
||||
The starting `beta` value of inference.
|
||||
beta_end (`float`, defaults to 0.02):
|
||||
The final `beta` value.
|
||||
beta_schedule (`str`, defaults to `"linear"`):
|
||||
The beta schedule, a mapping from a beta range to a sequence of betas for stepping the model. Choose from
|
||||
`linear`, `scaled_linear`, or `squaredcos_cap_v2`.
|
||||
trained_betas (`np.ndarray`, *optional*):
|
||||
Pass an array of betas directly to the constructor to bypass `beta_start` and `beta_end`.
|
||||
clip_sample (`bool`, defaults to `True`):
|
||||
Clip the predicted sample for numerical stability.
|
||||
clip_sample_range (`float`, defaults to 1.0):
|
||||
The maximum magnitude for sample clipping. Valid only when `clip_sample=True`.
|
||||
set_alpha_to_one (`bool`, defaults to `True`):
|
||||
Each diffusion step uses the alphas product value at that step and at the previous one. For the final step
|
||||
there is no previous alpha. When this option is `True` the previous alpha product is fixed to `1`,
|
||||
otherwise it uses the alpha value at step 0.
|
||||
steps_offset (`int`, defaults to 0):
|
||||
An offset added to the inference steps. You can use a combination of `offset=1` and
|
||||
`set_alpha_to_one=False` to make the last step use step 0 for the previous alpha product like in Stable
|
||||
Diffusion.
|
||||
prediction_type (`str`, defaults to `epsilon`, *optional*):
|
||||
Prediction type of the scheduler function; can be `epsilon` (predicts the noise of the diffusion process),
|
||||
`sample` (directly predicts the noisy sample`) or `v_prediction` (see section 2.4 of [Imagen
|
||||
Video](https://imagen.research.google/video/paper.pdf) paper).
|
||||
thresholding (`bool`, defaults to `False`):
|
||||
Whether to use the "dynamic thresholding" method. This is unsuitable for latent-space diffusion models such
|
||||
as Stable Diffusion.
|
||||
dynamic_thresholding_ratio (`float`, defaults to 0.995):
|
||||
The ratio for the dynamic thresholding method. Valid only when `thresholding=True`.
|
||||
sample_max_value (`float`, defaults to 1.0):
|
||||
The threshold value for dynamic thresholding. Valid only when `thresholding=True`.
|
||||
timestep_spacing (`str`, defaults to `"leading"`):
|
||||
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
|
||||
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
|
||||
rescale_betas_zero_snr (`bool`, defaults to `False`):
|
||||
Whether to rescale the betas to have zero terminal SNR. This enables the model to generate very bright and
|
||||
dark samples instead of limiting it to samples with medium brightness. Loosely related to
|
||||
[`--offset_noise`](https://github.com/huggingface/diffusers/blob/74fd735eb073eb1d774b1ab4154a0876eb82f055/examples/dreambooth/train_dreambooth.py#L506).
|
||||
"""
|
||||
|
||||
# _compatibles = [e.name for e in KarrasDiffusionSchedulers]
|
||||
order = 1
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
num_train_timesteps: int = 1000,
|
||||
beta_start: float = 0.0001,
|
||||
beta_end: float = 0.02,
|
||||
beta_schedule: str = "linear",
|
||||
trained_betas: Optional[Union[np.ndarray, List[float]]] = None,
|
||||
clip_sample: bool = True,
|
||||
set_alpha_to_one: bool = True,
|
||||
steps_offset: int = 0,
|
||||
prediction_type: str = "epsilon",
|
||||
thresholding: bool = False,
|
||||
dynamic_thresholding_ratio: float = 0.995,
|
||||
clip_sample_range: float = 1.0,
|
||||
sample_max_value: float = 1.0,
|
||||
timestep_spacing: str = "leading",
|
||||
rescale_betas_zero_snr: bool = False,
|
||||
):
|
||||
if trained_betas is not None:
|
||||
self.betas = torch.tensor(trained_betas, dtype=torch.float32)
|
||||
elif beta_schedule == "linear":
|
||||
self.betas = torch.linspace(
|
||||
beta_start, beta_end, num_train_timesteps, dtype=torch.float32
|
||||
)
|
||||
elif beta_schedule == "scaled_linear":
|
||||
# this schedule is very specific to the latent diffusion model.
|
||||
self.betas = (
|
||||
torch.linspace(
|
||||
beta_start**0.5,
|
||||
beta_end**0.5,
|
||||
num_train_timesteps,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
** 2
|
||||
)
|
||||
elif beta_schedule == "squaredcos_cap_v2":
|
||||
# Glide cosine schedule
|
||||
self.betas = betas_for_alpha_bar(num_train_timesteps)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f"{beta_schedule} does is not implemented for {self.__class__}"
|
||||
)
|
||||
|
||||
# Rescale for zero SNR
|
||||
if rescale_betas_zero_snr:
|
||||
self.betas = rescale_zero_terminal_snr(self.betas)
|
||||
|
||||
self.alphas = 1.0 - self.betas
|
||||
self.alphas_cumprod = torch.cumprod(self.alphas, dim=0)
|
||||
|
||||
# At every step in ddim, we are looking into the previous alphas_cumprod
|
||||
# For the final step, there is no previous alphas_cumprod because we are already at 0
|
||||
# `set_alpha_to_one` decides whether we set this parameter simply to one or
|
||||
# whether we use the final alpha of the "non-previous" one.
|
||||
self.final_alpha_cumprod = (
|
||||
torch.tensor(1.0) if set_alpha_to_one else self.alphas_cumprod[0]
|
||||
)
|
||||
|
||||
# standard deviation of the initial noise distribution
|
||||
self.init_noise_sigma = 1.0
|
||||
|
||||
# setable values
|
||||
self.num_inference_steps = None
|
||||
self.timesteps = torch.from_numpy(
|
||||
np.arange(0, num_train_timesteps)[::-1].copy().astype(np.int64)
|
||||
)
|
||||
|
||||
def scale_model_input(
|
||||
self, sample: torch.FloatTensor, timestep: Optional[int] = None
|
||||
) -> torch.FloatTensor:
|
||||
"""
|
||||
Ensures interchangeability with schedulers that need to scale the denoising model input depending on the
|
||||
current timestep.
|
||||
|
||||
Args:
|
||||
sample (`torch.FloatTensor`):
|
||||
The input sample.
|
||||
timestep (`int`, *optional*):
|
||||
The current timestep in the diffusion chain.
|
||||
|
||||
Returns:
|
||||
`torch.FloatTensor`:
|
||||
A scaled input sample.
|
||||
"""
|
||||
return sample
|
||||
|
||||
def _get_variance(self, timestep, prev_timestep):
|
||||
alpha_prod_t = self.alphas_cumprod[timestep]
|
||||
alpha_prod_t_prev = (
|
||||
self.alphas_cumprod[prev_timestep]
|
||||
if prev_timestep >= 0
|
||||
else self.final_alpha_cumprod
|
||||
)
|
||||
beta_prod_t = 1 - alpha_prod_t
|
||||
beta_prod_t_prev = 1 - alpha_prod_t_prev
|
||||
|
||||
variance = (beta_prod_t_prev / beta_prod_t) * (
|
||||
1 - alpha_prod_t / alpha_prod_t_prev
|
||||
)
|
||||
|
||||
return variance
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_ddpm.DDPMScheduler._threshold_sample
|
||||
def _threshold_sample(self, sample: torch.FloatTensor) -> torch.FloatTensor:
|
||||
"""
|
||||
"Dynamic thresholding: At each sampling step we set s to a certain percentile absolute pixel value in xt0 (the
|
||||
prediction of x_0 at timestep t), and if s > 1, then we threshold xt0 to the range [-s, s] and then divide by
|
||||
s. Dynamic thresholding pushes saturated pixels (those near -1 and 1) inwards, thereby actively preventing
|
||||
pixels from saturation at each step. We find that dynamic thresholding results in significantly better
|
||||
photorealism as well as better image-text alignment, especially when using very large guidance weights."
|
||||
|
||||
https://arxiv.org/abs/2205.11487
|
||||
"""
|
||||
dtype = sample.dtype
|
||||
batch_size, channels, height, width = sample.shape
|
||||
|
||||
if dtype not in (torch.float32, torch.float64):
|
||||
# upcast for quantile calculation, and clamp not implemented for cpu half
|
||||
sample = sample.float()
|
||||
|
||||
# Flatten sample for doing quantile calculation along each image
|
||||
sample = sample.reshape(batch_size, channels * height * width)
|
||||
|
||||
abs_sample = sample.abs() # "a certain percentile absolute pixel value"
|
||||
|
||||
s = torch.quantile(abs_sample, self.config.dynamic_thresholding_ratio, dim=1)
|
||||
s = torch.clamp(
|
||||
s, min=1, max=self.config.sample_max_value
|
||||
) # When clamped to min=1, equivalent to standard clipping to [-1, 1]
|
||||
|
||||
# (batch_size, 1) because clamp will broadcast along dim=0
|
||||
s = s.unsqueeze(1)
|
||||
# "we threshold xt0 to the range [-s, s] and then divide by s"
|
||||
sample = torch.clamp(sample, -s, s) / s
|
||||
|
||||
sample = sample.reshape(batch_size, channels, height, width)
|
||||
sample = sample.to(dtype)
|
||||
|
||||
return sample
|
||||
|
||||
def set_timesteps(
|
||||
self,
|
||||
num_inference_steps: int,
|
||||
lcm_origin_steps: int,
|
||||
device: Union[str, torch.device] = None,
|
||||
):
|
||||
"""
|
||||
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
|
||||
|
||||
Args:
|
||||
num_inference_steps (`int`):
|
||||
The number of diffusion steps used when generating samples with a pre-trained model.
|
||||
"""
|
||||
|
||||
if num_inference_steps > self.config.num_train_timesteps:
|
||||
raise ValueError(
|
||||
f"`num_inference_steps`: {num_inference_steps} cannot be larger than `self.config.train_timesteps`:"
|
||||
f" {self.config.num_train_timesteps} as the unet model trained with this scheduler can only handle"
|
||||
f" maximal {self.config.num_train_timesteps} timesteps."
|
||||
)
|
||||
|
||||
self.num_inference_steps = num_inference_steps
|
||||
|
||||
# LCM Timesteps Setting: # Linear Spacing
|
||||
c = self.config.num_train_timesteps // lcm_origin_steps
|
||||
lcm_origin_timesteps = (
|
||||
np.asarray(list(range(1, lcm_origin_steps + 1))) * c - 1
|
||||
) # LCM Training Steps Schedule
|
||||
skipping_step = len(lcm_origin_timesteps) // num_inference_steps
|
||||
# LCM Inference Steps Schedule
|
||||
timesteps = lcm_origin_timesteps[::-skipping_step][:num_inference_steps]
|
||||
|
||||
self.timesteps = torch.from_numpy(timesteps.copy()).to(device)
|
||||
|
||||
def get_scalings_for_boundary_condition_discrete(self, t):
|
||||
self.sigma_data = 0.5 # Default: 0.5
|
||||
|
||||
# By dividing 0.1: This is almost a delta function at t=0.
|
||||
c_skip = self.sigma_data**2 / ((t / 0.1) ** 2 + self.sigma_data**2)
|
||||
c_out = (t / 0.1) / ((t / 0.1) ** 2 + self.sigma_data**2) ** 0.5
|
||||
return c_skip, c_out
|
||||
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.FloatTensor,
|
||||
timeindex: int,
|
||||
timestep: int,
|
||||
sample: torch.FloatTensor,
|
||||
eta: float = 0.0,
|
||||
use_clipped_model_output: bool = False,
|
||||
generator=None,
|
||||
variance_noise: Optional[torch.FloatTensor] = None,
|
||||
return_dict: bool = True,
|
||||
) -> Union[LCMSchedulerOutput, Tuple]:
|
||||
"""
|
||||
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
|
||||
process from the learned model outputs (most often the predicted noise).
|
||||
|
||||
Args:
|
||||
model_output (`torch.FloatTensor`):
|
||||
The direct output from learned diffusion model.
|
||||
timestep (`float`):
|
||||
The current discrete timestep in the diffusion chain.
|
||||
sample (`torch.FloatTensor`):
|
||||
A current instance of a sample created by the diffusion process.
|
||||
eta (`float`):
|
||||
The weight of noise for added noise in diffusion step.
|
||||
use_clipped_model_output (`bool`, defaults to `False`):
|
||||
If `True`, computes "corrected" `model_output` from the clipped predicted original sample. Necessary
|
||||
because predicted original sample is clipped to [-1, 1] when `self.config.clip_sample` is `True`. If no
|
||||
clipping has happened, "corrected" `model_output` would coincide with the one provided as input and
|
||||
`use_clipped_model_output` has no effect.
|
||||
generator (`torch.Generator`, *optional*):
|
||||
A random number generator.
|
||||
variance_noise (`torch.FloatTensor`):
|
||||
Alternative to generating noise with `generator` by directly providing the noise for the variance
|
||||
itself. Useful for methods such as [`CycleDiffusion`].
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~schedulers.scheduling_lcm.LCMSchedulerOutput`] or `tuple`.
|
||||
|
||||
Returns:
|
||||
[`~schedulers.scheduling_utils.LCMSchedulerOutput`] or `tuple`:
|
||||
If return_dict is `True`, [`~schedulers.scheduling_lcm.LCMSchedulerOutput`] is returned, otherwise a
|
||||
tuple is returned where the first element is the sample tensor.
|
||||
|
||||
"""
|
||||
if self.num_inference_steps is None:
|
||||
raise ValueError(
|
||||
"Number of inference steps is 'None', you need to run 'set_timesteps' after creating the scheduler"
|
||||
)
|
||||
|
||||
# 1. get previous step value
|
||||
prev_timeindex = timeindex + 1
|
||||
if prev_timeindex < len(self.timesteps):
|
||||
prev_timestep = self.timesteps[prev_timeindex]
|
||||
else:
|
||||
prev_timestep = timestep
|
||||
|
||||
# 2. compute alphas, betas
|
||||
alpha_prod_t = self.alphas_cumprod[timestep]
|
||||
alpha_prod_t_prev = (
|
||||
self.alphas_cumprod[prev_timestep]
|
||||
if prev_timestep >= 0
|
||||
else self.final_alpha_cumprod
|
||||
)
|
||||
|
||||
beta_prod_t = 1 - alpha_prod_t
|
||||
beta_prod_t_prev = 1 - alpha_prod_t_prev
|
||||
|
||||
# 3. Get scalings for boundary conditions
|
||||
c_skip, c_out = self.get_scalings_for_boundary_condition_discrete(timestep)
|
||||
|
||||
# 4. Different Parameterization:
|
||||
parameterization = self.config.prediction_type
|
||||
|
||||
if parameterization == "epsilon": # noise-prediction
|
||||
pred_x0 = (sample - beta_prod_t.sqrt() * model_output) / alpha_prod_t.sqrt()
|
||||
|
||||
elif parameterization == "sample": # x-prediction
|
||||
pred_x0 = model_output
|
||||
|
||||
elif parameterization == "v_prediction": # v-prediction
|
||||
pred_x0 = alpha_prod_t.sqrt() * sample - beta_prod_t.sqrt() * model_output
|
||||
|
||||
# 4. Denoise model output using boundary conditions
|
||||
denoised = c_out * pred_x0 + c_skip * sample
|
||||
|
||||
# 5. Sample z ~ N(0, I), For MultiStep Inference
|
||||
# Noise is not used for one-step sampling.
|
||||
if len(self.timesteps) > 1:
|
||||
noise = torch.randn(model_output.shape).to(model_output.device)
|
||||
prev_sample = (
|
||||
alpha_prod_t_prev.sqrt() * denoised + beta_prod_t_prev.sqrt() * noise
|
||||
)
|
||||
else:
|
||||
prev_sample = denoised
|
||||
|
||||
if not return_dict:
|
||||
return (prev_sample, denoised)
|
||||
|
||||
return LCMSchedulerOutput(prev_sample=prev_sample, denoised=denoised)
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_ddpm.DDPMScheduler.add_noise
|
||||
|
||||
def add_noise(
|
||||
self,
|
||||
original_samples: torch.FloatTensor,
|
||||
noise: torch.FloatTensor,
|
||||
timesteps: torch.IntTensor,
|
||||
) -> torch.FloatTensor:
|
||||
# Make sure alphas_cumprod and timestep have same device and dtype as original_samples
|
||||
alphas_cumprod = self.alphas_cumprod.to(
|
||||
device=original_samples.device, dtype=original_samples.dtype
|
||||
)
|
||||
timesteps = timesteps.to(original_samples.device)
|
||||
|
||||
sqrt_alpha_prod = alphas_cumprod[timesteps] ** 0.5
|
||||
sqrt_alpha_prod = sqrt_alpha_prod.flatten()
|
||||
while len(sqrt_alpha_prod.shape) < len(original_samples.shape):
|
||||
sqrt_alpha_prod = sqrt_alpha_prod.unsqueeze(-1)
|
||||
|
||||
sqrt_one_minus_alpha_prod = (1 - alphas_cumprod[timesteps]) ** 0.5
|
||||
sqrt_one_minus_alpha_prod = sqrt_one_minus_alpha_prod.flatten()
|
||||
while len(sqrt_one_minus_alpha_prod.shape) < len(original_samples.shape):
|
||||
sqrt_one_minus_alpha_prod = sqrt_one_minus_alpha_prod.unsqueeze(-1)
|
||||
|
||||
noisy_samples = (
|
||||
sqrt_alpha_prod * original_samples + sqrt_one_minus_alpha_prod * noise
|
||||
)
|
||||
return noisy_samples
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_ddpm.DDPMScheduler.get_velocity
|
||||
def get_velocity(
|
||||
self,
|
||||
sample: torch.FloatTensor,
|
||||
noise: torch.FloatTensor,
|
||||
timesteps: torch.IntTensor,
|
||||
) -> torch.FloatTensor:
|
||||
# Make sure alphas_cumprod and timestep have same device and dtype as sample
|
||||
alphas_cumprod = self.alphas_cumprod.to(
|
||||
device=sample.device, dtype=sample.dtype
|
||||
)
|
||||
timesteps = timesteps.to(sample.device)
|
||||
|
||||
sqrt_alpha_prod = alphas_cumprod[timesteps] ** 0.5
|
||||
sqrt_alpha_prod = sqrt_alpha_prod.flatten()
|
||||
while len(sqrt_alpha_prod.shape) < len(sample.shape):
|
||||
sqrt_alpha_prod = sqrt_alpha_prod.unsqueeze(-1)
|
||||
|
||||
sqrt_one_minus_alpha_prod = (1 - alphas_cumprod[timesteps]) ** 0.5
|
||||
sqrt_one_minus_alpha_prod = sqrt_one_minus_alpha_prod.flatten()
|
||||
while len(sqrt_one_minus_alpha_prod.shape) < len(sample.shape):
|
||||
sqrt_one_minus_alpha_prod = sqrt_one_minus_alpha_prod.unsqueeze(-1)
|
||||
|
||||
velocity = sqrt_alpha_prod * noise - sqrt_one_minus_alpha_prod * sample
|
||||
return velocity
|
||||
|
||||
def __len__(self):
|
||||
return self.config.num_train_timesteps
|
||||
@@ -0,0 +1,71 @@
|
||||
import os
|
||||
|
||||
from coremltools import ComputeUnit
|
||||
from python_coreml_stable_diffusion.coreml_model import CoreMLModel
|
||||
|
||||
from coreml_suite.lcm import lcm_converter
|
||||
|
||||
|
||||
class CoreMLConverterLCM:
|
||||
"""Converts a LCM model to Core ML."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"height": ("INT", {"default": 512, "min": 512, "max": 768, "step": 8}),
|
||||
"width": ("INT", {"default": 512, "min": 512, "max": 768, "step": 8}),
|
||||
"batch_size": ("INT", {"default": 1, "min": 1, "max": 64}),
|
||||
"compute_unit": (
|
||||
[
|
||||
ComputeUnit.CPU_AND_NE.name,
|
||||
ComputeUnit.CPU_AND_GPU.name,
|
||||
ComputeUnit.ALL.name,
|
||||
ComputeUnit.CPU_ONLY.name,
|
||||
],
|
||||
),
|
||||
# "controlnet_support": ("BOOLEAN", {"default": False}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("COREML_UNET",)
|
||||
RETURN_NAMES = ("coreml_model",)
|
||||
FUNCTION = "convert"
|
||||
|
||||
def convert(
|
||||
self, height, width, batch_size, compute_unit, controlnet_support=False
|
||||
):
|
||||
"""Converts a LCM model to Core ML.
|
||||
|
||||
Args:
|
||||
height (int): Height of the target image.
|
||||
width (int): Width of the target image.
|
||||
batch_size (int): Batch size.
|
||||
compute_unit (str): Compute unit to use when loading the model.
|
||||
|
||||
Returns:
|
||||
coreml_model: The converted Core ML model.
|
||||
|
||||
The converted model is also saved to "models/unet" directory and
|
||||
can be loaded with the "LCMCoreMLLoaderUNet" node.
|
||||
"""
|
||||
h = height
|
||||
w = width
|
||||
sample_size = (h // 8, w // 8)
|
||||
batch_size = batch_size
|
||||
cn_support_str = "_cn" if controlnet_support else ""
|
||||
|
||||
out_name = f"{lcm_converter.MODEL_NAME}_{batch_size}x{w}x{h}{cn_support_str}"
|
||||
|
||||
out_path = lcm_converter.get_out_path("unet", f"{out_name}")
|
||||
|
||||
if not os.path.exists(out_path):
|
||||
lcm_converter.convert(
|
||||
out_path=out_path,
|
||||
sample_size=sample_size,
|
||||
batch_size=batch_size,
|
||||
controlnet_support=controlnet_support,
|
||||
)
|
||||
target_path = lcm_converter.compile_model(out_path=out_path, out_name=out_name)
|
||||
|
||||
return (CoreMLModel(target_path, compute_unit, "compiled"),)
|
||||
@@ -0,0 +1,19 @@
|
||||
{
|
||||
"_class_name": "LCMScheduler",
|
||||
"_diffusers_version": "0.22.0.dev0",
|
||||
"beta_end": 0.012,
|
||||
"beta_schedule": "scaled_linear",
|
||||
"beta_start": 0.00085,
|
||||
"clip_sample": true,
|
||||
"clip_sample_range": 1.0,
|
||||
"dynamic_thresholding_ratio": 0.995,
|
||||
"num_train_timesteps": 1000,
|
||||
"prediction_type": "epsilon",
|
||||
"rescale_betas_zero_snr": false,
|
||||
"sample_max_value": 1.0,
|
||||
"set_alpha_to_one": true,
|
||||
"steps_offset": 0,
|
||||
"thresholding": false,
|
||||
"timestep_spacing": "leading",
|
||||
"trained_betas": null
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
from overrides import overrides
|
||||
from python_coreml_stable_diffusion.unet import UNet2DConditionModel, TimestepEmbedding
|
||||
|
||||
|
||||
class UNet2DConditionModelLCM(UNet2DConditionModel):
|
||||
def __init__(
|
||||
self,
|
||||
time_cond_proj_dim=None,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
timestep_input_dim = self.config.block_out_channels[0]
|
||||
time_embed_dim = self.config.block_out_channels[0] * 4
|
||||
|
||||
time_embedding = TimestepEmbedding(
|
||||
timestep_input_dim, time_embed_dim, cond_proj_dim=time_cond_proj_dim
|
||||
)
|
||||
self.time_embedding = time_embedding
|
||||
|
||||
@overrides(check_signature=False)
|
||||
def forward(
|
||||
self,
|
||||
sample,
|
||||
timestep,
|
||||
encoder_hidden_states,
|
||||
timestep_cond,
|
||||
*additional_residuals,
|
||||
):
|
||||
# 0. Project (or look-up) time embeddings
|
||||
t_emb = self.time_proj(timestep)
|
||||
emb = self.time_embedding(t_emb, timestep_cond)
|
||||
|
||||
# 1. center input if necessary
|
||||
if self.config.center_input_sample:
|
||||
sample = 2 * sample - 1.0
|
||||
|
||||
# 2. pre-process
|
||||
sample = self.conv_in(sample)
|
||||
|
||||
# 3. down
|
||||
down_block_res_samples = (sample,)
|
||||
for downsample_block in self.down_blocks:
|
||||
if (
|
||||
hasattr(downsample_block, "attentions")
|
||||
and downsample_block.attentions is not None
|
||||
):
|
||||
sample, res_samples = downsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
)
|
||||
else:
|
||||
sample, res_samples = downsample_block(hidden_states=sample, temb=emb)
|
||||
|
||||
down_block_res_samples += res_samples
|
||||
|
||||
if additional_residuals:
|
||||
new_down_block_res_samples = ()
|
||||
for i, down_block_res_sample in enumerate(down_block_res_samples):
|
||||
down_block_res_sample = down_block_res_sample + additional_residuals[i]
|
||||
new_down_block_res_samples += (down_block_res_sample,)
|
||||
down_block_res_samples = new_down_block_res_samples
|
||||
|
||||
# 4. mid
|
||||
sample = self.mid_block(
|
||||
sample, emb, encoder_hidden_states=encoder_hidden_states
|
||||
)
|
||||
|
||||
if additional_residuals:
|
||||
sample = sample + additional_residuals[-1]
|
||||
|
||||
# 5. up
|
||||
for upsample_block in self.up_blocks:
|
||||
res_samples = down_block_res_samples[-len(upsample_block.resnets) :]
|
||||
down_block_res_samples = down_block_res_samples[
|
||||
: -len(upsample_block.resnets)
|
||||
]
|
||||
|
||||
if (
|
||||
hasattr(upsample_block, "attentions")
|
||||
and upsample_block.attentions is not None
|
||||
):
|
||||
sample = upsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
res_hidden_states_tuple=res_samples,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
)
|
||||
else:
|
||||
sample = upsample_block(
|
||||
hidden_states=sample, temb=emb, res_hidden_states_tuple=res_samples
|
||||
)
|
||||
|
||||
# 6. post-process
|
||||
sample = self.conv_norm_out(sample)
|
||||
sample = self.conv_act(sample)
|
||||
sample = self.conv_out(sample)
|
||||
|
||||
return (sample,)
|
||||
+42
-21
@@ -38,32 +38,29 @@ class CoreMLModelWrapper(BaseModel):
|
||||
c_adm=None,
|
||||
control=None,
|
||||
transformer_options={},
|
||||
**kwargs,
|
||||
):
|
||||
chunked_in = self.chunk_inputs(x, t, c_crossattn, control)
|
||||
chunked_in = self.chunk_inputs(
|
||||
x, t, c_crossattn, control, kwargs.get("timestep_cond")
|
||||
)
|
||||
chunked_out = [
|
||||
self._apply_model(
|
||||
x, t, c_concat, c_crossattn, c_adm, control, transformer_options
|
||||
)
|
||||
for x, t, c_crossattn, control in zip(*chunked_in)
|
||||
self._apply_model(x, t, c_crossattn, control, ts_cond)
|
||||
for x, t, c_crossattn, control, ts_cond in zip(*chunked_in)
|
||||
]
|
||||
|
||||
merged_out = merge_chunks(chunked_out, x.shape)
|
||||
return merged_out
|
||||
|
||||
def _apply_model(self, x, t, c_crossattn, control=None, ts_cond=None):
|
||||
model_input_kwargs = self.prepare_inputs(x, t, c_crossattn, control, ts_cond)
|
||||
|
||||
np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"]
|
||||
return torch.from_numpy(np_out).to(x.device)
|
||||
|
||||
def get_dtype(self):
|
||||
# Hardcoding torch-compatible dtype (used for memory allocation)
|
||||
return torch.float16
|
||||
|
||||
def _apply_model(
|
||||
self,
|
||||
x,
|
||||
t,
|
||||
c_concat=None,
|
||||
c_crossattn=None,
|
||||
c_adm=None,
|
||||
control=None,
|
||||
transformer_options={},
|
||||
):
|
||||
def prepare_inputs(self, x, t, c_crossattn, control, ts_cond=None):
|
||||
sample = x.cpu().numpy().astype(np.float16)
|
||||
|
||||
context = c_crossattn.cpu().numpy().astype(np.float16)
|
||||
@@ -78,12 +75,15 @@ class CoreMLModelWrapper(BaseModel):
|
||||
}
|
||||
residual_kwargs = extract_residual_kwargs(self.diffusion_model, control)
|
||||
model_input_kwargs |= residual_kwargs
|
||||
# model_input_kwargs = expand_inputs(model_input_kwargs)
|
||||
|
||||
np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"]
|
||||
return torch.from_numpy(np_out).to(x.device)
|
||||
if ts_cond is not None:
|
||||
model_input_kwargs["timestep_cond"] = (
|
||||
ts_cond.cpu().numpy().astype(np.float16)
|
||||
)
|
||||
|
||||
def chunk_inputs(self, x, t, c_crossattn, control):
|
||||
return model_input_kwargs
|
||||
|
||||
def chunk_inputs(self, x, t, c_crossattn, control, ts_cond=None):
|
||||
sample_shape = self.expected_inputs["sample"]["shape"]
|
||||
timestep_shape = self.expected_inputs["timestep"]["shape"]
|
||||
hidden_shape = self.expected_inputs["encoder_hidden_states"]["shape"]
|
||||
@@ -97,8 +97,29 @@ class CoreMLModelWrapper(BaseModel):
|
||||
if control is not None:
|
||||
chunked_control = chunk_control(control, sample_shape[0])
|
||||
|
||||
return chunked_x, ts, chunked_context, chunked_control
|
||||
chunked_ts_cond = [None] * len(chunked_x)
|
||||
if ts_cond is not None:
|
||||
ts_cond_shape = self.expected_inputs["timestep_cond"]["shape"]
|
||||
chunked_ts_cond = chunk_batch(ts_cond, ts_cond_shape)
|
||||
|
||||
return chunked_x, ts, chunked_context, chunked_control, chunked_ts_cond
|
||||
|
||||
@property
|
||||
def expected_inputs(self):
|
||||
return self.diffusion_model.expected_inputs
|
||||
|
||||
def __call__(self, latents, ts, encoder_hidden_states, **kwargs):
|
||||
return (
|
||||
self.apply_model(latents, ts, c_crossattn=encoder_hidden_states, **kwargs),
|
||||
)
|
||||
|
||||
|
||||
class CoreMLModelWrapperLCM(CoreMLModelWrapper):
|
||||
def __init__(self, model_config, coreml_model):
|
||||
super().__init__(model_config, coreml_model)
|
||||
self.config = None
|
||||
|
||||
def __call__(self, latents, ts, encoder_hidden_states, **kwargs):
|
||||
return (
|
||||
self.apply_model(latents, ts, c_crossattn=encoder_hidden_states, **kwargs),
|
||||
)
|
||||
|
||||
+30
-16
@@ -7,7 +7,11 @@ import torch
|
||||
from comfy.model_management import get_torch_device
|
||||
from coreml_suite.latents import chunk_batch, merge_chunks
|
||||
from coreml_suite.controlnet import chunk_control
|
||||
from coreml_suite.models import CoreMLModelWrapper, get_model_config
|
||||
from coreml_suite.models import (
|
||||
CoreMLModelWrapper,
|
||||
get_model_config,
|
||||
CoreMLModelWrapperLCM,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -16,6 +20,7 @@ def coreml_model():
|
||||
model.expected_inputs = {
|
||||
"sample": {"shape": (2, 4, 64, 64)},
|
||||
"timestep": {"shape": (2,)},
|
||||
"timestep_cond": {"shape": (2, 256)},
|
||||
"encoder_hidden_states": {"shape": (2, 768, 1, 77)},
|
||||
"additional_residual_0": {"shape": (2, 320, 64, 64)},
|
||||
"additional_residual_1": {"shape": (2, 640, 32, 32)},
|
||||
@@ -54,6 +59,22 @@ def test_merge_chunks(batch_size):
|
||||
assert torch.equal(input_tensor, merged)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def inputs():
|
||||
x = torch.randn(1, 4, 64, 64).to(get_torch_device())
|
||||
t = torch.randn([1]).to(get_torch_device())
|
||||
c_crossattn = torch.randn(1, 77, 768).to(get_torch_device())
|
||||
control = {
|
||||
"output": [
|
||||
torch.randn(1, 320, 64, 64).to(get_torch_device()),
|
||||
torch.randn(1, 640, 32, 32).to(get_torch_device()),
|
||||
],
|
||||
}
|
||||
timestep_cond = torch.randn(1, 256).to(get_torch_device())
|
||||
|
||||
return x, t, c_crossattn, control, timestep_cond
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"b, target_size, num_chunks",
|
||||
[
|
||||
@@ -95,29 +116,22 @@ def test_chunking_no_control():
|
||||
assert chunked == [None, None]
|
||||
|
||||
|
||||
def test_chunking_inputs(coreml_model, model_config):
|
||||
def test_chunking_inputs(coreml_model, model_config, inputs):
|
||||
model = CoreMLModelWrapper(model_config, coreml_model)
|
||||
x = torch.randn(1, 4, 64, 64).to(get_torch_device())
|
||||
t = torch.randn([1]).to(get_torch_device())
|
||||
c_crossattn = torch.randn(1, 77, 768).to(get_torch_device())
|
||||
control = {
|
||||
"output": [
|
||||
torch.randn(1, 320, 64, 64).to(get_torch_device()),
|
||||
torch.randn(1, 640, 32, 32).to(get_torch_device()),
|
||||
],
|
||||
}
|
||||
|
||||
chunked_x, ts, chunked_context, chunked_control = model.chunk_inputs(
|
||||
x, t, c_crossattn, control
|
||||
chunked_x, ts, chunked_context, chunked_cn, chunked_ts_cond = model.chunk_inputs(
|
||||
*inputs
|
||||
)
|
||||
|
||||
assert len(chunked_x) == 1
|
||||
assert len(ts) == 1
|
||||
assert len(chunked_context) == 1
|
||||
assert len(chunked_control) == 1
|
||||
assert len(chunked_cn) == 1
|
||||
assert len(chunked_ts_cond) == 1
|
||||
|
||||
assert chunked_x[0].shape == (2, 4, 64, 64)
|
||||
assert ts[0].shape == (2,)
|
||||
assert chunked_context[0].shape == (2, 77, 768)
|
||||
assert chunked_control[0]["output"][0].shape == (2, 320, 64, 64)
|
||||
assert chunked_control[0]["output"][1].shape == (2, 640, 32, 32)
|
||||
assert chunked_cn[0]["output"][0].shape == (2, 320, 64, 64)
|
||||
assert chunked_cn[0]["output"][1].shape == (2, 640, 32, 32)
|
||||
assert chunked_ts_cond[0].shape == (2, 256)
|
||||
|
||||
Reference in New Issue
Block a user