From 0092ad5e758acbc8d893690eb3f01caad9cb03a0 Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Sat, 28 Oct 2023 12:24:48 +0200 Subject: [PATCH 01/12] Prepare LCM Model Wrapper --- .gitignore | 2 +- coreml_suite/models.py | 37 ++++++++++++++++++++++++------------- 2 files changed, 25 insertions(+), 14 deletions(-) diff --git a/.gitignore b/.gitignore index 6ff69c2..f1358c3 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,3 @@ playground/ experiments/ -__pycache__/ +__pycache__/ \ No newline at end of file diff --git a/coreml_suite/models.py b/coreml_suite/models.py index b456f0d..ebb9d45 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -50,20 +50,21 @@ class CoreMLModelWrapper(BaseModel): merged_out = merge_chunks(chunked_out, x.shape) return merged_out + def _apply_model(self, x, t, c_concat=None, c_crossattn=None, c_adm=None, + control=None, transformer_options={}): + model_input_kwargs = self.prepare_inputs(x, t, c_crossattn, control) + residual_kwargs = extract_residual_kwargs(self.diffusion_model, + control) + model_input_kwargs |= residual_kwargs + + 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): sample = x.cpu().numpy().astype(np.float16) context = c_crossattn.cpu().numpy().astype(np.float16) @@ -78,10 +79,8 @@ 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) + return model_input_kwargs def chunk_inputs(self, x, t, c_crossattn, control): sample_shape = self.expected_inputs["sample"]["shape"] @@ -102,3 +101,15 @@ class CoreMLModelWrapper(BaseModel): @property def expected_inputs(self): return self.diffusion_model.expected_inputs + +class CoreMLModelWrapperLCM(CoreMLModelWrapper): + def __init__(self, model_config, coreml_model): + super().__init__(model_config, coreml_model) + self.config = None + + def _apply_model(self, x, t, c_concat=None, c_crossattn=None, c_adm=None, + control=None, transformer_options={}): + model_input_kwargs = self.prepare_inputs(x, t, c_crossattn, control) + + np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"] + return torch.from_numpy(np_out).to(x.device) From 9d509ad8f4d5460b8a5cf120d15fc542a0f572f2 Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Sun, 29 Oct 2023 13:58:33 +0100 Subject: [PATCH 02/12] WIP: LCM --- __init__.py | 5 + coreml_suite/lcm/__init__.py | 4 + coreml_suite/lcm/lcm_converter.py | 231 +++++++++++ coreml_suite/lcm/lcm_pipeline.py | 290 ++++++++++++++ coreml_suite/lcm/lcm_sampler.py | 93 +++++ coreml_suite/lcm/lcm_scheduler.py | 524 +++++++++++++++++++++++++ coreml_suite/lcm/nodes.py | 52 +++ coreml_suite/lcm/scheduler_config.json | 19 + coreml_suite/models.py | 24 +- 9 files changed, 1230 insertions(+), 12 deletions(-) create mode 100644 coreml_suite/lcm/__init__.py create mode 100644 coreml_suite/lcm/lcm_converter.py create mode 100644 coreml_suite/lcm/lcm_pipeline.py create mode 100644 coreml_suite/lcm/lcm_sampler.py create mode 100644 coreml_suite/lcm/lcm_scheduler.py create mode 100644 coreml_suite/lcm/nodes.py create mode 100644 coreml_suite/lcm/scheduler_config.json diff --git a/__init__.py b/__init__.py index 8a0d5c9..68f4c48 100644 --- a/__init__.py +++ b/__init__.py @@ -4,14 +4,19 @@ import sys sys.path.append(os.path.dirname(__file__)) from coreml_suite.nodes import CoreMLLoaderUNet, CoreMLSampler, CoreMLModelAdapter +from coreml_suite.lcm import CoreMLConverterLCM, CoreMLSamplerLCM NODE_CLASS_MAPPINGS = { "CoreMLUNetLoader": CoreMLLoaderUNet, "CoreMLSampler": CoreMLSampler, "CoreMLModelAdapter": CoreMLModelAdapter, + "CoreMLSamplerLCM": CoreMLSamplerLCM, + "CoreMLConverterLCM": CoreMLConverterLCM, } NODE_DISPLAY_NAME_MAPPINGS = { "CoreMLUNetLoader": "Load Core ML UNet", "CoreMLSampler": "Core ML Sampler", "CoreMLModelAdapter": "Core ML Adapter (Experimental)", + "CoreMLSamplerLCM": "Core ML LCM Sampler", + "CoreMLConverterLCM": "Convert LCM to Core ML", } diff --git a/coreml_suite/lcm/__init__.py b/coreml_suite/lcm/__init__.py new file mode 100644 index 0000000..7ac4c9c --- /dev/null +++ b/coreml_suite/lcm/__init__.py @@ -0,0 +1,4 @@ +from lcm_sampler import CoreMLSamplerLCM +from .nodes import CoreMLConverterLCM + +__all__ = ["CoreMLSamplerLCM", "CoreMLConverterLCM"] diff --git a/coreml_suite/lcm/lcm_converter.py b/coreml_suite/lcm/lcm_converter.py new file mode 100644 index 0000000..54c5446 --- /dev/null +++ b/coreml_suite/lcm/lcm_converter.py @@ -0,0 +1,231 @@ +import os +import shutil +import logging +import time +import gc + +import numpy as np +import torch +from diffusers import UNet2DConditionModel +from python_coreml_stable_diffusion.unet import ( + UNet2DConditionModel as CoreMLUNet2DConditionModel, +) +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_V2 + + +def get_unets(): + ref_unet = UNet2DConditionModel.from_pretrained( + MODEL_VERSION, + subfolder="unet", + device_map=None, + low_cpu_mem_usage=False, + ) + + ref_config = ref_unet.config + + cml_unet = CoreMLUNet2DConditionModel().eval() + cml_unet.load_state_dict(ref_unet.state_dict(), strict=False) + + del ref_unet + gc.collect() + + return cml_unet, ref_config + + +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)), + ] + ) + sample_unet_inputs_spec = { + k: (v.shape, v.dtype) for k, v in sample_unet_inputs.items() + } + return sample_unet_inputs, sample_unet_inputs_spec + + +def convert( + out_path: str, batch_size: int = 1, sample_size: tuple[int, int] = (64, 64) +): + coreml_unet, unet_config = get_unets() + + sample_shape = ( + batch_size, # B + unet_config.in_channels, # C + sample_size[0], # H + sample_size[1], # W + ) + + encoder_hidden_states_shape = get_encoder_hidden_states_shape( + unet_config, batch_size + ) + + scheduler = get_scheduler() + + sample_inputs, sample_inputs_spec = get_sample_input( + batch_size, encoder_hidden_states_shape, sample_shape, scheduler + ) + + logger.info(f"Sample UNet inputs spec: {sample_inputs_spec}") + logger.info("JIT tracing..") + traced_unet = torch.jit.trace(coreml_unet, example_kwarg_inputs=sample_inputs) + 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 + + out_name = f"{MODEL_NAME}_{w}x{h}_batch{batch_size}" + + 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) diff --git a/coreml_suite/lcm/lcm_pipeline.py b/coreml_suite/lcm/lcm_pipeline.py new file mode 100644 index 0000000..13dc1e5 --- /dev/null +++ b/coreml_suite/lcm/lcm_pipeline.py @@ -0,0 +1,290 @@ +import torch +from diffusers import DiffusionPipeline, AutoencoderKL, UNet2DConditionModel +from transformers import CLIPTokenizer, CLIPTextModel, CLIPImageProcessor +from diffusers.pipelines.stable_diffusion import StableDiffusionPipelineOutput +from diffusers.image_processor import VaeImageProcessor +from typing import List, Optional, Union, Dict, Any +from comfy.model_management import get_torch_device + + +# from diffusers import logging +# logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +class LatentConsistencyModelPipeline(DiffusionPipeline): + def __init__( + self, + vae: AutoencoderKL, + text_encoder: CLIPTextModel, + tokenizer: CLIPTokenizer, + unet: UNet2DConditionModel, + scheduler: None, + safety_checker: None, + feature_extractor: CLIPImageProcessor, + ): + super().__init__() + + self.register_modules( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + unet=unet, + scheduler=scheduler, + safety_checker=safety_checker, + feature_extractor=feature_extractor, + ) + self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1) + self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor) + + def _encode_prompt( + self, + prompt, + device, + num_images_per_prompt, + prompt_embeds: None, + ): + r""" + Encodes the prompt into text encoder hidden states. + + Args: + prompt (`str` or `List[str]`, *optional*): + prompt to be encoded + device: (`torch.device`): + torch device + num_images_per_prompt (`int`): + number of images that should be generated per prompt + prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not + provided, text embeddings will be generated from `prompt` input argument. + """ + + if prompt is not None and isinstance(prompt, str): + batch_size = 1 + elif prompt is not None and isinstance(prompt, list): + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + if prompt_embeds is None: + text_inputs = self.tokenizer( + prompt, + padding="max_length", + max_length=self.tokenizer.model_max_length, + truncation=True, + return_tensors="pt", + ) + text_input_ids = text_inputs.input_ids + untruncated_ids = self.tokenizer( + prompt, padding="longest", return_tensors="pt" + ).input_ids + + if untruncated_ids.shape[-1] >= text_input_ids.shape[ + -1 + ] and not torch.equal(text_input_ids, untruncated_ids): + removed_text = self.tokenizer.batch_decode( + untruncated_ids[:, self.tokenizer.model_max_length - 1 : -1] + ) + print( + "The following part of your input was truncated because CLIP can only handle sequences up to" + f" {self.tokenizer.model_max_length} tokens: {removed_text}" + ) + + if ( + hasattr(self.text_encoder.config, "use_attention_mask") + and self.text_encoder.config.use_attention_mask + ): + attention_mask = text_inputs.attention_mask.to(device) + else: + attention_mask = None + + prompt_embeds = self.text_encoder( + text_input_ids.to(device), + attention_mask=attention_mask, + ) + prompt_embeds = prompt_embeds[0] + + if self.text_encoder is not None: + prompt_embeds_dtype = self.text_encoder.dtype + elif self.unet is not None: + prompt_embeds_dtype = self.unet.dtype + else: + prompt_embeds_dtype = prompt_embeds.dtype + + prompt_embeds = prompt_embeds.to(dtype=prompt_embeds_dtype, device=device) + + bs_embed, seq_len, _ = prompt_embeds.shape + # duplicate text embeddings for each generation per prompt, using mps friendly method + prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1) + prompt_embeds = prompt_embeds.view( + bs_embed * num_images_per_prompt, seq_len, -1 + ) + + # Don't need to get uncond prompt embedding because of LCM Guided Distillation + return prompt_embeds + + # ¯\_(ツ)_/¯ + def run_safety_checker(self, image, device, dtype): + return image, None + + def prepare_latents( + self, + batch_size, + num_channels_latents, + height, + width, + dtype, + device, + latents=None, + ): + shape = ( + batch_size, + num_channels_latents, + height // self.vae_scale_factor, + width // self.vae_scale_factor, + ) + if latents is None: + latents = torch.randn(shape, dtype=dtype).to(device) + else: + latents = latents.to(device) + # scale the initial noise by the standard deviation required by the scheduler + latents = latents * self.scheduler.init_noise_sigma + 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 + + @torch.no_grad() + def __call__( + self, + prompt: Union[str, List[str]] = None, + height: Optional[int] = 768, + width: Optional[int] = 768, + guidance_scale: float = 7.5, + num_images_per_prompt: Optional[int] = 1, + latents: Optional[torch.FloatTensor] = None, + num_inference_steps: int = 4, + lcm_origin_steps: int = 50, + prompt_embeds: Optional[torch.FloatTensor] = None, + output_type: Optional[str] = "pil", + return_dict: bool = True, + cross_attention_kwargs: Optional[Dict[str, Any]] = None, + ): + # 0. Default height and width to unet + height = height or self.unet.config.sample_size * self.vae_scale_factor + width = width or self.unet.config.sample_size * self.vae_scale_factor + + # 2. Define call parameters + if prompt is not None and isinstance(prompt, str): + batch_size = 1 + elif prompt is not None and isinstance(prompt, list): + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + device = get_torch_device() + # do_classifier_free_guidance = guidance_scale > 0.0 # In LCM Implementation: cfg_noise = noise_cond + cfg_scale * (noise_cond - noise_uncond) , (cfg_scale > 0.0 using CFG) + + # 3. Encode input prompt + prompt_embeds = self._encode_prompt( + prompt, + device, + num_images_per_prompt, + prompt_embeds=prompt_embeds, + ) + + # 4. Prepare timesteps + self.scheduler.set_timesteps(num_inference_steps, lcm_origin_steps) + timesteps = self.scheduler.timesteps + + # 5. Prepare latent variable + num_channels_latents = self.unet.config.in_channels + latents = self.prepare_latents( + batch_size * num_images_per_prompt, + num_channels_latents, + height, + width, + prompt_embeds.dtype, + device, + latents, + ) + bs = batch_size * num_images_per_prompt + + # 6. Get Guidance Scale Embedding + w = torch.tensor(guidance_scale).repeat(bs) + w_embedding = self.get_w_embedding(w, embedding_dim=256).to( + device=device, dtype=latents.dtype + ) + + # 7. LCM MultiStep Sampling Loop: + with self.progress_bar(total=num_inference_steps) as progress_bar: + for i, t in enumerate(timesteps): + ts = torch.full((bs,), t, device=device, dtype=torch.long) + latents = latents.to(prompt_embeds.dtype) + + # model prediction (v-prediction, eps, x) + model_pred = self.unet( + latents, + ts, + timestep_cond=w_embedding, + encoder_hidden_states=prompt_embeds, + cross_attention_kwargs=cross_attention_kwargs, + return_dict=False, + )[0] + + # compute the previous noisy sample x_t -> x_t-1 + latents, denoised = self.scheduler.step( + model_pred, i, t, latents, return_dict=False + ) + + # # call the callback, if provided + # if i == len(timesteps) - 1: + progress_bar.update() + + denoised = denoised.to(prompt_embeds.dtype) + if not output_type == "latent": + image = self.vae.decode( + denoised / self.vae.config.scaling_factor, return_dict=False + )[0] + image, has_nsfw_concept = self.run_safety_checker( + image, device, prompt_embeds.dtype + ) + else: + image = denoised + has_nsfw_concept = None + + if has_nsfw_concept is None: + do_denormalize = [True] * image.shape[0] + else: + do_denormalize = [not has_nsfw for has_nsfw in has_nsfw_concept] + + image = self.image_processor.postprocess( + image, output_type=output_type, do_denormalize=do_denormalize + ) + + if not return_dict: + return (image, has_nsfw_concept) + + return StableDiffusionPipelineOutput( + images=image, nsfw_content_detected=has_nsfw_concept + ) diff --git a/coreml_suite/lcm/lcm_sampler.py b/coreml_suite/lcm/lcm_sampler.py new file mode 100644 index 0000000..12c9d73 --- /dev/null +++ b/coreml_suite/lcm/lcm_sampler.py @@ -0,0 +1,93 @@ +import os +import time + +import torch + +from comfy.model_management import get_torch_device +from coreml_suite.lcm.lcm_pipeline import LatentConsistencyModelPipeline +from coreml_suite.lcm.lcm_scheduler import LCMScheduler + + +class CoreMLSamplerLCM: + def __init__(self): + self.scheduler = LCMScheduler.from_pretrained( + os.path.join(os.path.dirname(__file__), "scheduler_config.json") + ) + self.pipe = None + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}), + "steps": ("INT", {"default": 4, "min": 1, "max": 10000}), + "cfg": ( + "FLOAT", + { + "default": 8.0, + "min": 0.0, + "max": 100.0, + "step": 0.5, + "round": 0.01, + }, + ), + "height": ("INT", {"default": 512, "min": 512, "max": 768}), + "width": ("INT", {"default": 512, "min": 512, "max": 768}), + "num_images": ("INT", {"default": 1, "min": 1, "max": 64}), + "use_fp16": ("BOOLEAN", {"default": True}), + "positive_prompt": ("STRING", {"multiline": True}), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "sample" + CATEGORY = "sampling" + + def sample( + self, + model, + seed, + steps, + cfg, + positive_prompt, + height, + width, + num_images, + use_fp16, + ): + if self.pipe is None: + self.pipe = LatentConsistencyModelPipeline.from_pretrained( + pretrained_model_name_or_path="SimianLuo/LCM_Dreamshaper_v7", + scheduler=self.scheduler, + safety_checker=None, + ) + + if use_fp16: + self.pipe.to(torch_device=get_torch_device(), torch_dtype=torch.float16) + else: + self.pipe.to(torch_device=get_torch_device(), torch_dtype=torch.float32) + + coreml_unet = model.model + coreml_unet.config = self.pipe.unet.config + + self.pipe.unet = coreml_unet + + torch.manual_seed(seed) + start_time = time.time() + + result = self.pipe( + prompt=positive_prompt, + width=width, + height=height, + guidance_scale=cfg, + num_inference_steps=steps, + num_images_per_prompt=num_images, + lcm_origin_steps=50, + output_type="np", + ).images + + print("LCM inference time: ", time.time() - start_time, "seconds") + images_tensor = torch.from_numpy(result) + + return (images_tensor,) diff --git a/coreml_suite/lcm/lcm_scheduler.py b/coreml_suite/lcm/lcm_scheduler.py new file mode 100644 index 0000000..cd5b209 --- /dev/null +++ b/coreml_suite/lcm/lcm_scheduler.py @@ -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 diff --git a/coreml_suite/lcm/nodes.py b/coreml_suite/lcm/nodes.py new file mode 100644 index 0000000..980e1a0 --- /dev/null +++ b/coreml_suite/lcm/nodes.py @@ -0,0 +1,52 @@ +import os + +from coreml_suite.lcm import lcm_converter + + +class CoreMLConverterLCM: + """Converts a LCM model to Core ML.""" + + RETURN_TYPES = ("COMBO",) + RETURN_NAMES = ("model_name",) + FUNCTION = "convert" + + @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": 4, "min": 1, "max": 64}), + } + } + + def convert(self, height, width, batch_size): + """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. + + Returns: + 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 + + out_name = f"{lcm_converter.MODEL_NAME}_{w}x{h}_batch{batch_size}" + + 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 + ) + target_path = lcm_converter.compile_model(out_path=out_path, out_name=out_name) + + return (target_path.split("/")[-1],) diff --git a/coreml_suite/lcm/scheduler_config.json b/coreml_suite/lcm/scheduler_config.json new file mode 100644 index 0000000..b6dba68 --- /dev/null +++ b/coreml_suite/lcm/scheduler_config.json @@ -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 +} \ No newline at end of file diff --git a/coreml_suite/models.py b/coreml_suite/models.py index ebb9d45..a85fdc2 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -30,14 +30,14 @@ class CoreMLModelWrapper(BaseModel): self.diffusion_model = coreml_model def apply_model( - self, - x, - t, - c_concat=None, - c_crossattn=None, - c_adm=None, - control=None, - transformer_options={}, + self, + x, + t, + c_concat=None, + c_crossattn=None, + c_adm=None, + control=None, + transformer_options={}, ): chunked_in = self.chunk_inputs(x, t, c_crossattn, control) chunked_out = [ @@ -53,9 +53,6 @@ class CoreMLModelWrapper(BaseModel): def _apply_model(self, x, t, c_concat=None, c_crossattn=None, c_adm=None, control=None, transformer_options={}): model_input_kwargs = self.prepare_inputs(x, t, c_crossattn, control) - residual_kwargs = extract_residual_kwargs(self.diffusion_model, - control) - model_input_kwargs |= residual_kwargs np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"] return torch.from_numpy(np_out).to(x.device) @@ -112,4 +109,7 @@ class CoreMLModelWrapperLCM(CoreMLModelWrapper): model_input_kwargs = self.prepare_inputs(x, t, c_crossattn, control) np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"] - return torch.from_numpy(np_out).to(x.device) + return (torch.from_numpy(np_out).to(x.device),) + + def __call__(self, latents, t, encoder_hidden_states, **kwargs): + return self.apply_model(latents, t, c_crossattn=encoder_hidden_states, **kwargs) From 213088241d449ad54dc5d2f6fabfad731986740d Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Mon, 30 Oct 2023 11:01:53 +0100 Subject: [PATCH 03/12] LCM Converter works --- coreml_suite/lcm/__init__.py | 2 +- coreml_suite/lcm/lcm_converter.py | 2 +- coreml_suite/lcm/lcm_sampler.py | 11 ++++++++--- coreml_suite/lcm/nodes.py | 24 +++++++++++++++++------- coreml_suite/models.py | 9 +++++---- 5 files changed, 32 insertions(+), 16 deletions(-) diff --git a/coreml_suite/lcm/__init__.py b/coreml_suite/lcm/__init__.py index 7ac4c9c..8ae122a 100644 --- a/coreml_suite/lcm/__init__.py +++ b/coreml_suite/lcm/__init__.py @@ -1,4 +1,4 @@ -from lcm_sampler import CoreMLSamplerLCM +from .lcm_sampler import CoreMLSamplerLCM from .nodes import CoreMLConverterLCM __all__ = ["CoreMLSamplerLCM", "CoreMLConverterLCM"] diff --git a/coreml_suite/lcm/lcm_converter.py b/coreml_suite/lcm/lcm_converter.py index 54c5446..6356586 100644 --- a/coreml_suite/lcm/lcm_converter.py +++ b/coreml_suite/lcm/lcm_converter.py @@ -25,7 +25,7 @@ MODEL_NAME = MODEL_VERSION.split("/")[-1] + "_4k" import python_coreml_stable_diffusion.unet as unet -unet.ATTENTION_IMPLEMENTATION_IN_EFFECT = unet.AttentionImplementations.SPLIT_EINSUM_V2 +unet.ATTENTION_IMPLEMENTATION_IN_EFFECT = unet.AttentionImplementations.SPLIT_EINSUM def get_unets(): diff --git a/coreml_suite/lcm/lcm_sampler.py b/coreml_suite/lcm/lcm_sampler.py index 12c9d73..1c9c8ca 100644 --- a/coreml_suite/lcm/lcm_sampler.py +++ b/coreml_suite/lcm/lcm_sampler.py @@ -6,6 +6,7 @@ import torch from comfy.model_management import get_torch_device from coreml_suite.lcm.lcm_pipeline import LatentConsistencyModelPipeline from coreml_suite.lcm.lcm_scheduler import LCMScheduler +from coreml_suite.models import get_model_config, CoreMLModelWrapperLCM class CoreMLSamplerLCM: @@ -19,7 +20,7 @@ class CoreMLSamplerLCM: def INPUT_TYPES(s): return { "required": { - "model": ("MODEL",), + "coreml_model": ("COREML_UNET",), "seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}), "steps": ("INT", {"default": 4, "min": 1, "max": 10000}), "cfg": ( @@ -46,7 +47,7 @@ class CoreMLSamplerLCM: def sample( self, - model, + coreml_model, seed, steps, cfg, @@ -56,6 +57,10 @@ class CoreMLSamplerLCM: num_images, use_fp16, ): + + model_config = get_model_config() + wrapped_model = CoreMLModelWrapperLCM(model_config, coreml_model) + if self.pipe is None: self.pipe = LatentConsistencyModelPipeline.from_pretrained( pretrained_model_name_or_path="SimianLuo/LCM_Dreamshaper_v7", @@ -68,7 +73,7 @@ class CoreMLSamplerLCM: else: self.pipe.to(torch_device=get_torch_device(), torch_dtype=torch.float32) - coreml_unet = model.model + coreml_unet = wrapped_model coreml_unet.config = self.pipe.unet.config self.pipe.unet = coreml_unet diff --git a/coreml_suite/lcm/nodes.py b/coreml_suite/lcm/nodes.py index 980e1a0..9028668 100644 --- a/coreml_suite/lcm/nodes.py +++ b/coreml_suite/lcm/nodes.py @@ -1,15 +1,14 @@ 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.""" - RETURN_TYPES = ("COMBO",) - RETURN_NAMES = ("model_name",) - FUNCTION = "convert" - @classmethod def INPUT_TYPES(cls): return { @@ -17,19 +16,30 @@ class CoreMLConverterLCM: "height": ("INT", {"default": 512, "min": 512, "max": 768, "step": 8}), "width": ("INT", {"default": 512, "min": 512, "max": 768, "step": 8}), "batch_size": ("INT", {"default": 4, "min": 1, "max": 64}), + "compute_unit": ([ + ComputeUnit.CPU_AND_NE.name, + ComputeUnit.CPU_AND_GPU.name, + ComputeUnit.ALL.name, + ComputeUnit.CPU_ONLY.name, + ],) } } - def convert(self, height, width, batch_size): + RETURN_TYPES = ("COREML_UNET",) + RETURN_NAMES = ("coreml_model",) + FUNCTION = "convert" + + def convert(self, height, width, batch_size, compute_unit): """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: - MODEL: The converted Core ML model. + 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. @@ -49,4 +59,4 @@ class CoreMLConverterLCM: ) target_path = lcm_converter.compile_model(out_path=out_path, out_name=out_name) - return (target_path.split("/")[-1],) + return (CoreMLModel(target_path, compute_unit, "compiled"),) diff --git a/coreml_suite/models.py b/coreml_suite/models.py index a85fdc2..24f455c 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -51,7 +51,7 @@ class CoreMLModelWrapper(BaseModel): return merged_out def _apply_model(self, x, t, c_concat=None, c_crossattn=None, c_adm=None, - control=None, transformer_options={}): + control=None, transformer_options={}): model_input_kwargs = self.prepare_inputs(x, t, c_crossattn, control) np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"] @@ -99,17 +99,18 @@ class CoreMLModelWrapper(BaseModel): def expected_inputs(self): return self.diffusion_model.expected_inputs + class CoreMLModelWrapperLCM(CoreMLModelWrapper): def __init__(self, model_config, coreml_model): super().__init__(model_config, coreml_model) self.config = None def _apply_model(self, x, t, c_concat=None, c_crossattn=None, c_adm=None, - control=None, transformer_options={}): + control=None, transformer_options={}): model_input_kwargs = self.prepare_inputs(x, t, c_crossattn, control) np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"] return (torch.from_numpy(np_out).to(x.device),) - def __call__(self, latents, t, encoder_hidden_states, **kwargs): - return self.apply_model(latents, t, c_crossattn=encoder_hidden_states, **kwargs) + def __call__(self, latents, ts, encoder_hidden_states, **kwargs): + return self.apply_model(latents, ts, c_crossattn=encoder_hidden_states) From 8a814b7a56b0a9ee34db6b8445ea8474c5bd1acc Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Tue, 31 Oct 2023 12:24:53 +0100 Subject: [PATCH 04/12] Fix LCM Sampler --- coreml_suite/models.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/coreml_suite/models.py b/coreml_suite/models.py index 24f455c..115a4fc 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -110,7 +110,7 @@ class CoreMLModelWrapperLCM(CoreMLModelWrapper): model_input_kwargs = self.prepare_inputs(x, t, c_crossattn, control) np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"] - return (torch.from_numpy(np_out).to(x.device),) + return torch.from_numpy(np_out).to(x.device) def __call__(self, latents, ts, encoder_hidden_states, **kwargs): - return self.apply_model(latents, ts, c_crossattn=encoder_hidden_states) + return (self.apply_model(latents, ts, c_crossattn=encoder_hidden_states),) From 1aa5a19b2aa7c100a8282a9b833590cbedeaf9cb Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Tue, 31 Oct 2023 19:48:46 +0100 Subject: [PATCH 05/12] Simplify LCM Sampler --- __init__.py | 6 +++--- coreml_suite/lcm/__init__.py | 4 ++-- coreml_suite/lcm/lcm_sampler.py | 38 ++++++++++++--------------------- 3 files changed, 19 insertions(+), 29 deletions(-) diff --git a/__init__.py b/__init__.py index 68f4c48..8410c43 100644 --- a/__init__.py +++ b/__init__.py @@ -4,19 +4,19 @@ import sys sys.path.append(os.path.dirname(__file__)) from coreml_suite.nodes import CoreMLLoaderUNet, CoreMLSampler, CoreMLModelAdapter -from coreml_suite.lcm import CoreMLConverterLCM, CoreMLSamplerLCM +from coreml_suite.lcm import CoreMLConverterLCM, CoreMLSamplerLCM_Simple NODE_CLASS_MAPPINGS = { "CoreMLUNetLoader": CoreMLLoaderUNet, "CoreMLSampler": CoreMLSampler, "CoreMLModelAdapter": CoreMLModelAdapter, - "CoreMLSamplerLCM": CoreMLSamplerLCM, + "Core ML LCM Sampler (Simple)": CoreMLSamplerLCM_Simple, "CoreMLConverterLCM": CoreMLConverterLCM, } NODE_DISPLAY_NAME_MAPPINGS = { "CoreMLUNetLoader": "Load Core ML UNet", "CoreMLSampler": "Core ML Sampler", "CoreMLModelAdapter": "Core ML Adapter (Experimental)", - "CoreMLSamplerLCM": "Core ML LCM Sampler", + "Core ML LCM Sampler (Simple)": "Core ML LCM Sampler (Simple)", "CoreMLConverterLCM": "Convert LCM to Core ML", } diff --git a/coreml_suite/lcm/__init__.py b/coreml_suite/lcm/__init__.py index 8ae122a..2df93e0 100644 --- a/coreml_suite/lcm/__init__.py +++ b/coreml_suite/lcm/__init__.py @@ -1,4 +1,4 @@ -from .lcm_sampler import CoreMLSamplerLCM +from .lcm_sampler import CoreMLSamplerLCM_Simple from .nodes import CoreMLConverterLCM -__all__ = ["CoreMLSamplerLCM", "CoreMLConverterLCM"] +__all__ = ["CoreMLSamplerLCM_Simple", "CoreMLConverterLCM"] diff --git a/coreml_suite/lcm/lcm_sampler.py b/coreml_suite/lcm/lcm_sampler.py index 1c9c8ca..08d23a7 100644 --- a/coreml_suite/lcm/lcm_sampler.py +++ b/coreml_suite/lcm/lcm_sampler.py @@ -1,5 +1,4 @@ import os -import time import torch @@ -9,7 +8,7 @@ from coreml_suite.lcm.lcm_scheduler import LCMScheduler from coreml_suite.models import get_model_config, CoreMLModelWrapperLCM -class CoreMLSamplerLCM: +class CoreMLSamplerLCM_Simple: def __init__(self): self.scheduler = LCMScheduler.from_pretrained( os.path.join(os.path.dirname(__file__), "scheduler_config.json") @@ -33,10 +32,7 @@ class CoreMLSamplerLCM: "round": 0.01, }, ), - "height": ("INT", {"default": 512, "min": 512, "max": 768}), - "width": ("INT", {"default": 512, "min": 512, "max": 768}), "num_images": ("INT", {"default": 1, "min": 1, "max": 64}), - "use_fp16": ("BOOLEAN", {"default": True}), "positive_prompt": ("STRING", {"multiline": True}), } } @@ -46,17 +42,16 @@ class CoreMLSamplerLCM: CATEGORY = "sampling" def sample( - self, - coreml_model, - seed, - steps, - cfg, - positive_prompt, - height, - width, - num_images, - use_fp16, + self, + coreml_model, + seed, + steps, + cfg, + positive_prompt, + num_images, ): + height = coreml_model.expected_inputs["sample"]["shape"][2] * 8 + width = coreml_model.expected_inputs["sample"]["shape"][3] * 8 model_config = get_model_config() wrapped_model = CoreMLModelWrapperLCM(model_config, coreml_model) @@ -68,18 +63,14 @@ class CoreMLSamplerLCM: safety_checker=None, ) - if use_fp16: - self.pipe.to(torch_device=get_torch_device(), torch_dtype=torch.float16) - else: - self.pipe.to(torch_device=get_torch_device(), torch_dtype=torch.float32) + self.pipe.to(torch_device=get_torch_device(), torch_dtype=torch.float16) - coreml_unet = wrapped_model - coreml_unet.config = self.pipe.unet.config + coreml_unet = wrapped_model + coreml_unet.config = self.pipe.unet.config - self.pipe.unet = coreml_unet + self.pipe.unet = coreml_unet torch.manual_seed(seed) - start_time = time.time() result = self.pipe( prompt=positive_prompt, @@ -92,7 +83,6 @@ class CoreMLSamplerLCM: output_type="np", ).images - print("LCM inference time: ", time.time() - start_time, "seconds") images_tensor = torch.from_numpy(result) return (images_tensor,) From b90591dfd4d6ebcf4802e2f167e4eefd0a20acc1 Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Wed, 1 Nov 2023 01:19:30 +0100 Subject: [PATCH 06/12] Add support for controlnet to LCM converter --- coreml_suite/lcm/lcm_converter.py | 80 ++++++++++++++++++++++++++----- coreml_suite/lcm/nodes.py | 10 ++-- 2 files changed, 73 insertions(+), 17 deletions(-) diff --git a/coreml_suite/lcm/lcm_converter.py b/coreml_suite/lcm/lcm_converter.py index 6356586..be3b4cc 100644 --- a/coreml_suite/lcm/lcm_converter.py +++ b/coreml_suite/lcm/lcm_converter.py @@ -41,10 +41,7 @@ def get_unets(): cml_unet = CoreMLUNet2DConditionModel().eval() cml_unet.load_state_dict(ref_unet.state_dict(), strict=False) - del ref_unet - gc.collect() - - return cml_unet, ref_config + return cml_unet, ref_unet def get_encoder_hidden_states_shape(unet_config, batch_size): @@ -101,7 +98,7 @@ def load_coreml_model(out_path): def convert_to_coreml( - submodule_name, torchscript_module, sample_inputs, output_names, out_path + 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") @@ -162,37 +159,92 @@ def get_sample_input(batch_size, encoder_hidden_states_shape, sample_shape, sche ("encoder_hidden_states", torch.rand(*encoder_hidden_states_shape)), ] ) + 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, sample_unet_inputs_spec + 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) + out_path: str, batch_size: int = 1, sample_size: tuple[int, int] = (64, 64), + controlnet_support: bool = False ): - coreml_unet, unet_config = get_unets() + coreml_unet, ref_unet = get_unets() sample_shape = ( batch_size, # B - unet_config.in_channels, # C + ref_unet.config.in_channels, # C sample_size[0], # H sample_size[1], # W ) encoder_hidden_states_shape = get_encoder_hidden_states_shape( - unet_config, batch_size + ref_unet.config, batch_size ) scheduler = get_scheduler() - sample_inputs, sample_inputs_spec = get_sample_input( + 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_kwarg_inputs=sample_inputs) + traced_unet = torch.jit.trace(coreml_unet, example_inputs=list(sample_inputs.values())) logger.info("Done.") coreml_sample_inputs = get_coreml_inputs(sample_inputs) @@ -223,7 +275,9 @@ if __name__ == "__main__": sample_size = (h // 8, w // 8) batch_size = 4 - out_name = f"{MODEL_NAME}_{w}x{h}_batch{batch_size}" + 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): diff --git a/coreml_suite/lcm/nodes.py b/coreml_suite/lcm/nodes.py index 9028668..296d8a9 100644 --- a/coreml_suite/lcm/nodes.py +++ b/coreml_suite/lcm/nodes.py @@ -21,7 +21,8 @@ class CoreMLConverterLCM: ComputeUnit.CPU_AND_GPU.name, ComputeUnit.ALL.name, ComputeUnit.CPU_ONLY.name, - ],) + ],), + "controlnet_support": ("BOOLEAN", {"default": False}), } } @@ -29,7 +30,7 @@ class CoreMLConverterLCM: RETURN_NAMES = ("coreml_model",) FUNCTION = "convert" - def convert(self, height, width, batch_size, compute_unit): + def convert(self, height, width, batch_size, compute_unit, controlnet_support): """Converts a LCM model to Core ML. Args: @@ -48,14 +49,15 @@ class CoreMLConverterLCM: 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}_{w}x{h}_batch{batch_size}" + 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 + 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) From 1937f39cca9b0ec899f87f7dbb9eb3d0f39cb9ad Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Wed, 1 Nov 2023 01:25:51 +0100 Subject: [PATCH 07/12] Add support for CN models to LCM --- coreml_suite/lcm/lcm_converter.py | 40 +++++++++++++++++++------------ coreml_suite/lcm/lcm_sampler.py | 14 +++++------ coreml_suite/lcm/nodes.py | 19 +++++++++------ coreml_suite/models.py | 11 +-------- 4 files changed, 45 insertions(+), 39 deletions(-) diff --git a/coreml_suite/lcm/lcm_converter.py b/coreml_suite/lcm/lcm_converter.py index be3b4cc..5259584 100644 --- a/coreml_suite/lcm/lcm_converter.py +++ b/coreml_suite/lcm/lcm_converter.py @@ -98,7 +98,7 @@ def load_coreml_model(out_path): def convert_to_coreml( - submodule_name, torchscript_module, sample_inputs, output_names, out_path + 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") @@ -171,6 +171,7 @@ def get_unet_inputs_spec(sample_unet_inputs): 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] @@ -183,27 +184,32 @@ def add_cnet_support(sample_shape, reference_unet): reference_unet.conv_in, ) additional_residuals_shapes.append( - (batch_size, reference_unet.conv_in.out_channels, out_h, out_w)) + (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 + (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: + 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) + 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)) + ( + 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) + (batch_size, reference_unet.mid_block.resnets[-1].out_channels, out_h, out_w) ) additional_inputs = {} @@ -215,8 +221,10 @@ def add_cnet_support(sample_shape, reference_unet): def convert( - out_path: str, batch_size: int = 1, sample_size: tuple[int, int] = (64, 64), - controlnet_support: bool = False + out_path: str, + batch_size: int = 1, + sample_size: tuple[int, int] = (64, 64), + controlnet_support: bool = False, ): coreml_unet, ref_unet = get_unets() @@ -244,7 +252,9 @@ def convert( 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())) + traced_unet = torch.jit.trace( + coreml_unet, example_inputs=list(sample_inputs.values()) + ) logger.info("Done.") coreml_sample_inputs = get_coreml_inputs(sample_inputs) diff --git a/coreml_suite/lcm/lcm_sampler.py b/coreml_suite/lcm/lcm_sampler.py index 08d23a7..d9b18d7 100644 --- a/coreml_suite/lcm/lcm_sampler.py +++ b/coreml_suite/lcm/lcm_sampler.py @@ -42,13 +42,13 @@ class CoreMLSamplerLCM_Simple: CATEGORY = "sampling" def sample( - self, - coreml_model, - seed, - steps, - cfg, - positive_prompt, - num_images, + self, + coreml_model, + seed, + steps, + cfg, + positive_prompt, + num_images, ): height = coreml_model.expected_inputs["sample"]["shape"][2] * 8 width = coreml_model.expected_inputs["sample"]["shape"][3] * 8 diff --git a/coreml_suite/lcm/nodes.py b/coreml_suite/lcm/nodes.py index 296d8a9..fa9a527 100644 --- a/coreml_suite/lcm/nodes.py +++ b/coreml_suite/lcm/nodes.py @@ -16,12 +16,14 @@ class CoreMLConverterLCM: "height": ("INT", {"default": 512, "min": 512, "max": 768, "step": 8}), "width": ("INT", {"default": 512, "min": 512, "max": 768, "step": 8}), "batch_size": ("INT", {"default": 4, "min": 1, "max": 64}), - "compute_unit": ([ - ComputeUnit.CPU_AND_NE.name, - ComputeUnit.CPU_AND_GPU.name, - ComputeUnit.ALL.name, - ComputeUnit.CPU_ONLY.name, - ],), + "compute_unit": ( + [ + ComputeUnit.CPU_AND_NE.name, + ComputeUnit.CPU_AND_GPU.name, + ComputeUnit.ALL.name, + ComputeUnit.CPU_ONLY.name, + ], + ), "controlnet_support": ("BOOLEAN", {"default": False}), } } @@ -57,7 +59,10 @@ class CoreMLConverterLCM: 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 + 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) diff --git a/coreml_suite/models.py b/coreml_suite/models.py index 115a4fc..a7d4b32 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -41,9 +41,7 @@ class CoreMLModelWrapper(BaseModel): ): chunked_in = self.chunk_inputs(x, t, c_crossattn, control) chunked_out = [ - self._apply_model( - x, t, c_concat, c_crossattn, c_adm, control, transformer_options - ) + self._apply_model(x, t, c_crossattn, control) for x, t, c_crossattn, control in zip(*chunked_in) ] @@ -105,12 +103,5 @@ class CoreMLModelWrapperLCM(CoreMLModelWrapper): super().__init__(model_config, coreml_model) self.config = None - def _apply_model(self, x, t, c_concat=None, c_crossattn=None, c_adm=None, - control=None, transformer_options={}): - model_input_kwargs = self.prepare_inputs(x, t, c_crossattn, control) - - np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"] - return torch.from_numpy(np_out).to(x.device) - def __call__(self, latents, ts, encoder_hidden_states, **kwargs): return (self.apply_model(latents, ts, c_crossattn=encoder_hidden_states),) From e22d8187cd55bd58bfbb45afa20d119fdd494db5 Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Wed, 1 Nov 2023 21:46:24 +0100 Subject: [PATCH 08/12] Add more advanced LCM Sampler --- __init__.py | 8 ++- coreml_suite/lcm/__init__.py | 4 +- coreml_suite/lcm/lcm_sampler.py | 122 ++++++++++++++++++++++++++++++++ coreml_suite/models.py | 3 +- 4 files changed, 132 insertions(+), 5 deletions(-) diff --git a/__init__.py b/__init__.py index 8410c43..6897878 100644 --- a/__init__.py +++ b/__init__.py @@ -4,12 +4,17 @@ import sys sys.path.append(os.path.dirname(__file__)) from coreml_suite.nodes import CoreMLLoaderUNet, CoreMLSampler, CoreMLModelAdapter -from coreml_suite.lcm import CoreMLConverterLCM, CoreMLSamplerLCM_Simple +from coreml_suite.lcm import ( + CoreMLSamplerLCM, + CoreMLConverterLCM, + CoreMLSamplerLCM_Simple, +) NODE_CLASS_MAPPINGS = { "CoreMLUNetLoader": CoreMLLoaderUNet, "CoreMLSampler": CoreMLSampler, "CoreMLModelAdapter": CoreMLModelAdapter, + "Core ML LCM Sampler": CoreMLSamplerLCM, "Core ML LCM Sampler (Simple)": CoreMLSamplerLCM_Simple, "CoreMLConverterLCM": CoreMLConverterLCM, } @@ -17,6 +22,7 @@ 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 Sampler (Simple)": "Core ML LCM Sampler (Simple)", "CoreMLConverterLCM": "Convert LCM to Core ML", } diff --git a/coreml_suite/lcm/__init__.py b/coreml_suite/lcm/__init__.py index 2df93e0..3886985 100644 --- a/coreml_suite/lcm/__init__.py +++ b/coreml_suite/lcm/__init__.py @@ -1,4 +1,4 @@ -from .lcm_sampler import CoreMLSamplerLCM_Simple +from .lcm_sampler import CoreMLSamplerLCM, CoreMLSamplerLCM_Simple from .nodes import CoreMLConverterLCM -__all__ = ["CoreMLSamplerLCM_Simple", "CoreMLConverterLCM"] +__all__ = ["CoreMLSamplerLCM", "CoreMLSamplerLCM_Simple", "CoreMLConverterLCM"] diff --git a/coreml_suite/lcm/lcm_sampler.py b/coreml_suite/lcm/lcm_sampler.py index d9b18d7..d9caece 100644 --- a/coreml_suite/lcm/lcm_sampler.py +++ b/coreml_suite/lcm/lcm_sampler.py @@ -1,11 +1,17 @@ import os +import numpy as np import torch +import comfy.utils +import latent_preview from comfy.model_management import get_torch_device +from comfy.model_patcher import ModelPatcher from coreml_suite.lcm.lcm_pipeline import LatentConsistencyModelPipeline 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_Simple: @@ -86,3 +92,119 @@ class CoreMLSamplerLCM_Simple: images_tensor = torch.from_numpy(result) return (images_tensor,) + + +class CoreMLSamplerLCM(CoreMLSampler): + @classmethod + def INPUT_TYPES(s): + old_required = CoreMLSampler.INPUT_TYPES()["required"].copy() + 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] + + torch.manual_seed(seed) + + return self._sample(patched_model, steps, cfg, positive, latent_image, denoise) + + def _sample(self, model, steps, cfg, positive, latent_image, denoise): + batch_size = latent_image["samples"].shape[0] + + 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) + + device = get_torch_device() + # callback = latent_preview.prepare_callback(model, steps, None) + + # Prepare timesteps + 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 + + # Prepare latent variable + latents = self.prepare_latents(latent_image, device) + + # LCM MultiStep Sampling Loop: + progress_bar = comfy.utils.ProgressBar(total=steps) + for i, t in enumerate(timesteps): + ts = torch.full((batch_size,), t, device=device, dtype=torch.float16) + + # model prediction (v-prediction, eps, x) + model_pred = model.model( + latents, + ts, + encoder_hidden_states=prompt_embeds, + )[0] + + # model_pred *= cfg + + # compute the previous noisy sample x_t -> x_t-1 + latents, denoised = self.scheduler.step( + model_pred, i, t, latents, return_dict=False + ) + + # # call the callback, if provided + # if i == len(timesteps) - 1: + # callback(i, t, latents, steps) + + denoised = denoised.to(get_torch_device()) + + return ({"samples": denoised / 0.1825},) + + def prepare_latents(self, latent_image, device): + latent = latent_image["samples"] + 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] + noise = torch.randn(latent.shape, dtype=torch.float16).cpu() + latent_timestep = self.scheduler.timesteps[:1].repeat(batch_size) + latents = self.scheduler.add_noise(latent, noise, latent_timestep) + + return latents diff --git a/coreml_suite/models.py b/coreml_suite/models.py index a7d4b32..ae4067e 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -48,8 +48,7 @@ class CoreMLModelWrapper(BaseModel): merged_out = merge_chunks(chunked_out, x.shape) return merged_out - def _apply_model(self, x, t, c_concat=None, c_crossattn=None, c_adm=None, - control=None, transformer_options={}): + def _apply_model(self, x, t, c_crossattn, control=None): model_input_kwargs = self.prepare_inputs(x, t, c_crossattn, control) np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"] From 44a380ffdf084e9c1fd788ddf6717387a6df3ef3 Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Fri, 3 Nov 2023 00:27:15 +0100 Subject: [PATCH 09/12] img2img works --- coreml_suite/latents.py | 6 +- coreml_suite/lcm/lcm_converter.py | 12 ++-- coreml_suite/lcm/lcm_pipeline.py | 2 + coreml_suite/lcm/lcm_sampler.py | 106 +++++++++++++++++++----------- coreml_suite/lcm/nodes.py | 6 +- coreml_suite/lcm/unet.py | 100 ++++++++++++++++++++++++++++ coreml_suite/models.py | 55 +++++++++++----- tests/test_chunks.py | 46 ++++++++----- 8 files changed, 249 insertions(+), 84 deletions(-) create mode 100644 coreml_suite/lcm/unet.py diff --git a/coreml_suite/latents.py b/coreml_suite/latents.py index db823dc..7b44b30 100644 --- a/coreml_suite/latents.py +++ b/coreml_suite/latents.py @@ -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) diff --git a/coreml_suite/lcm/lcm_converter.py b/coreml_suite/lcm/lcm_converter.py index 5259584..86f0778 100644 --- a/coreml_suite/lcm/lcm_converter.py +++ b/coreml_suite/lcm/lcm_converter.py @@ -7,9 +7,8 @@ import gc import numpy as np import torch from diffusers import UNet2DConditionModel -from python_coreml_stable_diffusion.unet import ( - UNet2DConditionModel as CoreMLUNet2DConditionModel, -) +from coreml_suite.lcm.unet import UNet2DConditionModelLCM + from transformers import CLIPTextModel import coremltools as ct @@ -36,9 +35,7 @@ def get_unets(): low_cpu_mem_usage=False, ) - ref_config = ref_unet.config - - cml_unet = CoreMLUNet2DConditionModel().eval() + 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 @@ -152,11 +149,12 @@ def get_sample_input(batch_size, encoder_hidden_states_shape, sample_shape, sche ("sample", torch.rand(*sample_shape)), ( "timestep", - torch.tensor([scheduler.timesteps[0].item()] * (batch_size)).to( + 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 diff --git a/coreml_suite/lcm/lcm_pipeline.py b/coreml_suite/lcm/lcm_pipeline.py index 13dc1e5..1d353bd 100644 --- a/coreml_suite/lcm/lcm_pipeline.py +++ b/coreml_suite/lcm/lcm_pipeline.py @@ -243,6 +243,7 @@ class LatentConsistencyModelPipeline(DiffusionPipeline): latents = latents.to(prompt_embeds.dtype) # model prediction (v-prediction, eps, x) + print("latents", latents.shape) model_pred = self.unet( latents, ts, @@ -251,6 +252,7 @@ class LatentConsistencyModelPipeline(DiffusionPipeline): cross_attention_kwargs=cross_attention_kwargs, return_dict=False, )[0] + print("model_pred", model_pred.shape) # compute the previous noisy sample x_t -> x_t-1 latents, denoised = self.scheduler.step( diff --git a/coreml_suite/lcm/lcm_sampler.py b/coreml_suite/lcm/lcm_sampler.py index d9caece..fe25bff 100644 --- a/coreml_suite/lcm/lcm_sampler.py +++ b/coreml_suite/lcm/lcm_sampler.py @@ -2,6 +2,7 @@ import os import numpy as np import torch +from diffusers.utils.torch_utils import randn_tensor import comfy.utils import latent_preview @@ -141,17 +142,49 @@ class CoreMLSamplerLCM(CoreMLSampler): return self._sample(patched_model, steps, cfg, positive, latent_image, denoise) def _sample(self, model, steps, cfg, positive, latent_image, denoise): + device = get_torch_device() batch_size = latent_image["samples"].shape[0] + # callback = latent_preview.prepare_callback(model, steps, None) + 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: + for i, t in enumerate(timesteps): + 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 + ) + + 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 - device = get_torch_device() - # callback = latent_preview.prepare_callback(model, steps, None) - - # Prepare timesteps + 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 @@ -162,49 +195,48 @@ class CoreMLSamplerLCM(CoreMLSampler): timesteps = lcm_origin_timesteps[::-skipping_step][:steps] timesteps = torch.from_numpy(timesteps.copy()).to(device) self.scheduler.timesteps = timesteps - timesteps = self.scheduler.timesteps - - # Prepare latent variable - latents = self.prepare_latents(latent_image, device) - - # LCM MultiStep Sampling Loop: - progress_bar = comfy.utils.ProgressBar(total=steps) - for i, t in enumerate(timesteps): - ts = torch.full((batch_size,), t, device=device, dtype=torch.float16) - - # model prediction (v-prediction, eps, x) - model_pred = model.model( - latents, - ts, - encoder_hidden_states=prompt_embeds, - )[0] - - # model_pred *= cfg - - # compute the previous noisy sample x_t -> x_t-1 - latents, denoised = self.scheduler.step( - model_pred, i, t, latents, return_dict=False - ) - - # # call the callback, if provided - # if i == len(timesteps) - 1: - # callback(i, t, latents, steps) - - denoised = denoised.to(get_torch_device()) - - return ({"samples": denoised / 0.1825},) + return timesteps def prepare_latents(self, latent_image, device): - latent = latent_image["samples"] + 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] - noise = torch.randn(latent.shape, dtype=torch.float16).cpu() + + 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 diff --git a/coreml_suite/lcm/nodes.py b/coreml_suite/lcm/nodes.py index fa9a527..6172911 100644 --- a/coreml_suite/lcm/nodes.py +++ b/coreml_suite/lcm/nodes.py @@ -24,7 +24,7 @@ class CoreMLConverterLCM: ComputeUnit.CPU_ONLY.name, ], ), - "controlnet_support": ("BOOLEAN", {"default": False}), + # "controlnet_support": ("BOOLEAN", {"default": False}), } } @@ -32,7 +32,9 @@ class CoreMLConverterLCM: RETURN_NAMES = ("coreml_model",) FUNCTION = "convert" - def convert(self, height, width, batch_size, compute_unit, controlnet_support): + def convert( + self, height, width, batch_size, compute_unit, controlnet_support=False + ): """Converts a LCM model to Core ML. Args: diff --git a/coreml_suite/lcm/unet.py b/coreml_suite/lcm/unet.py new file mode 100644 index 0000000..b756e12 --- /dev/null +++ b/coreml_suite/lcm/unet.py @@ -0,0 +1,100 @@ +from diffusers.configuration_utils import register_to_config +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,) diff --git a/coreml_suite/models.py b/coreml_suite/models.py index ae4067e..3b456ae 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -30,26 +30,28 @@ class CoreMLModelWrapper(BaseModel): self.diffusion_model = coreml_model def apply_model( - self, - x, - t, - c_concat=None, - c_crossattn=None, - c_adm=None, - control=None, - transformer_options={}, + self, + x, + t, + c_concat=None, + c_crossattn=None, + 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_crossattn, control) - 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): - model_input_kwargs = self.prepare_inputs(x, t, c_crossattn, control) + 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) @@ -58,7 +60,7 @@ class CoreMLModelWrapper(BaseModel): # Hardcoding torch-compatible dtype (used for memory allocation) return torch.float16 - def prepare_inputs(self, x, t, c_crossattn, control): + 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) @@ -74,9 +76,14 @@ class CoreMLModelWrapper(BaseModel): residual_kwargs = extract_residual_kwargs(self.diffusion_model, control) model_input_kwargs |= residual_kwargs + if ts_cond is not None: + model_input_kwargs["timestep_cond"] = ( + ts_cond.cpu().numpy().astype(np.float16) + ) + return model_input_kwargs - def chunk_inputs(self, x, t, c_crossattn, control): + 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"] @@ -90,12 +97,22 @@ 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): @@ -103,4 +120,6 @@ class CoreMLModelWrapperLCM(CoreMLModelWrapper): self.config = None def __call__(self, latents, ts, encoder_hidden_states, **kwargs): - return (self.apply_model(latents, ts, c_crossattn=encoder_hidden_states),) + return ( + self.apply_model(latents, ts, c_crossattn=encoder_hidden_states, **kwargs), + ) diff --git a/tests/test_chunks.py b/tests/test_chunks.py index abf8231..421ba52 100644 --- a/tests/test_chunks.py +++ b/tests/test_chunks.py @@ -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) From c01c60e3c1a57028cc9c3992ac27120299e5f97f Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Fri, 3 Nov 2023 01:05:28 +0100 Subject: [PATCH 10/12] Add progress bar and preview to LCM --- coreml_suite/lcm/lcm_sampler.py | 17 +++++++++++++---- 1 file changed, 13 insertions(+), 4 deletions(-) diff --git a/coreml_suite/lcm/lcm_sampler.py b/coreml_suite/lcm/lcm_sampler.py index fe25bff..5b90321 100644 --- a/coreml_suite/lcm/lcm_sampler.py +++ b/coreml_suite/lcm/lcm_sampler.py @@ -3,6 +3,7 @@ import os import numpy as np import torch from diffusers.utils.torch_utils import randn_tensor +from tqdm import tqdm import comfy.utils import latent_preview @@ -137,14 +138,18 @@ class CoreMLSamplerLCM(CoreMLSampler): 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) + return self._sample( + patched_model, steps, cfg, positive, latent_image, denoise, callback + ) - def _sample(self, model, steps, cfg, positive, latent_image, denoise): + def _sample( + self, model, steps, cfg, positive, latent_image, denoise, callback=None + ): device = get_torch_device() batch_size = latent_image["samples"].shape[0] - # callback = latent_preview.prepare_callback(model, steps, None) prompt_embeds = self.prepare_prompt_embeds(batch_size, positive) @@ -158,7 +163,8 @@ class CoreMLSamplerLCM(CoreMLSampler): ) # LCM MultiStep Sampling Loop: - for i, t in enumerate(timesteps): + 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( @@ -173,6 +179,9 @@ class CoreMLSamplerLCM(CoreMLSampler): 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},) From 1ebd9e72aef28e61f6528c89808a49df762763c8 Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Fri, 3 Nov 2023 01:21:02 +0100 Subject: [PATCH 11/12] Remove Simple LCM Sampler --- __init__.py | 7 +- coreml_suite/lcm/__init__.py | 4 +- coreml_suite/lcm/lcm_pipeline.py | 292 ------------------------------- coreml_suite/lcm/lcm_sampler.py | 82 --------- coreml_suite/lcm/unet.py | 1 - 5 files changed, 4 insertions(+), 382 deletions(-) delete mode 100644 coreml_suite/lcm/lcm_pipeline.py diff --git a/__init__.py b/__init__.py index 6897878..8729cb0 100644 --- a/__init__.py +++ b/__init__.py @@ -7,7 +7,6 @@ from coreml_suite.nodes import CoreMLLoaderUNet, CoreMLSampler, CoreMLModelAdapt from coreml_suite.lcm import ( CoreMLSamplerLCM, CoreMLConverterLCM, - CoreMLSamplerLCM_Simple, ) NODE_CLASS_MAPPINGS = { @@ -15,14 +14,12 @@ NODE_CLASS_MAPPINGS = { "CoreMLSampler": CoreMLSampler, "CoreMLModelAdapter": CoreMLModelAdapter, "Core ML LCM Sampler": CoreMLSamplerLCM, - "Core ML LCM Sampler (Simple)": CoreMLSamplerLCM_Simple, - "CoreMLConverterLCM": CoreMLConverterLCM, + "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 Sampler (Simple)": "Core ML LCM Sampler (Simple)", - "CoreMLConverterLCM": "Convert LCM to Core ML", + "Core ML LCM Converter": "Convert LCM to Core ML", } diff --git a/coreml_suite/lcm/__init__.py b/coreml_suite/lcm/__init__.py index 3886985..8ae122a 100644 --- a/coreml_suite/lcm/__init__.py +++ b/coreml_suite/lcm/__init__.py @@ -1,4 +1,4 @@ -from .lcm_sampler import CoreMLSamplerLCM, CoreMLSamplerLCM_Simple +from .lcm_sampler import CoreMLSamplerLCM from .nodes import CoreMLConverterLCM -__all__ = ["CoreMLSamplerLCM", "CoreMLSamplerLCM_Simple", "CoreMLConverterLCM"] +__all__ = ["CoreMLSamplerLCM", "CoreMLConverterLCM"] diff --git a/coreml_suite/lcm/lcm_pipeline.py b/coreml_suite/lcm/lcm_pipeline.py deleted file mode 100644 index 1d353bd..0000000 --- a/coreml_suite/lcm/lcm_pipeline.py +++ /dev/null @@ -1,292 +0,0 @@ -import torch -from diffusers import DiffusionPipeline, AutoencoderKL, UNet2DConditionModel -from transformers import CLIPTokenizer, CLIPTextModel, CLIPImageProcessor -from diffusers.pipelines.stable_diffusion import StableDiffusionPipelineOutput -from diffusers.image_processor import VaeImageProcessor -from typing import List, Optional, Union, Dict, Any -from comfy.model_management import get_torch_device - - -# from diffusers import logging -# logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class LatentConsistencyModelPipeline(DiffusionPipeline): - def __init__( - self, - vae: AutoencoderKL, - text_encoder: CLIPTextModel, - tokenizer: CLIPTokenizer, - unet: UNet2DConditionModel, - scheduler: None, - safety_checker: None, - feature_extractor: CLIPImageProcessor, - ): - super().__init__() - - self.register_modules( - vae=vae, - text_encoder=text_encoder, - tokenizer=tokenizer, - unet=unet, - scheduler=scheduler, - safety_checker=safety_checker, - feature_extractor=feature_extractor, - ) - self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1) - self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor) - - def _encode_prompt( - self, - prompt, - device, - num_images_per_prompt, - prompt_embeds: None, - ): - r""" - Encodes the prompt into text encoder hidden states. - - Args: - prompt (`str` or `List[str]`, *optional*): - prompt to be encoded - device: (`torch.device`): - torch device - num_images_per_prompt (`int`): - number of images that should be generated per prompt - prompt_embeds (`torch.FloatTensor`, *optional*): - Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not - provided, text embeddings will be generated from `prompt` input argument. - """ - - if prompt is not None and isinstance(prompt, str): - batch_size = 1 - elif prompt is not None and isinstance(prompt, list): - batch_size = len(prompt) - else: - batch_size = prompt_embeds.shape[0] - - if prompt_embeds is None: - text_inputs = self.tokenizer( - prompt, - padding="max_length", - max_length=self.tokenizer.model_max_length, - truncation=True, - return_tensors="pt", - ) - text_input_ids = text_inputs.input_ids - untruncated_ids = self.tokenizer( - prompt, padding="longest", return_tensors="pt" - ).input_ids - - if untruncated_ids.shape[-1] >= text_input_ids.shape[ - -1 - ] and not torch.equal(text_input_ids, untruncated_ids): - removed_text = self.tokenizer.batch_decode( - untruncated_ids[:, self.tokenizer.model_max_length - 1 : -1] - ) - print( - "The following part of your input was truncated because CLIP can only handle sequences up to" - f" {self.tokenizer.model_max_length} tokens: {removed_text}" - ) - - if ( - hasattr(self.text_encoder.config, "use_attention_mask") - and self.text_encoder.config.use_attention_mask - ): - attention_mask = text_inputs.attention_mask.to(device) - else: - attention_mask = None - - prompt_embeds = self.text_encoder( - text_input_ids.to(device), - attention_mask=attention_mask, - ) - prompt_embeds = prompt_embeds[0] - - if self.text_encoder is not None: - prompt_embeds_dtype = self.text_encoder.dtype - elif self.unet is not None: - prompt_embeds_dtype = self.unet.dtype - else: - prompt_embeds_dtype = prompt_embeds.dtype - - prompt_embeds = prompt_embeds.to(dtype=prompt_embeds_dtype, device=device) - - bs_embed, seq_len, _ = prompt_embeds.shape - # duplicate text embeddings for each generation per prompt, using mps friendly method - prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1) - prompt_embeds = prompt_embeds.view( - bs_embed * num_images_per_prompt, seq_len, -1 - ) - - # Don't need to get uncond prompt embedding because of LCM Guided Distillation - return prompt_embeds - - # ¯\_(ツ)_/¯ - def run_safety_checker(self, image, device, dtype): - return image, None - - def prepare_latents( - self, - batch_size, - num_channels_latents, - height, - width, - dtype, - device, - latents=None, - ): - shape = ( - batch_size, - num_channels_latents, - height // self.vae_scale_factor, - width // self.vae_scale_factor, - ) - if latents is None: - latents = torch.randn(shape, dtype=dtype).to(device) - else: - latents = latents.to(device) - # scale the initial noise by the standard deviation required by the scheduler - latents = latents * self.scheduler.init_noise_sigma - 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 - - @torch.no_grad() - def __call__( - self, - prompt: Union[str, List[str]] = None, - height: Optional[int] = 768, - width: Optional[int] = 768, - guidance_scale: float = 7.5, - num_images_per_prompt: Optional[int] = 1, - latents: Optional[torch.FloatTensor] = None, - num_inference_steps: int = 4, - lcm_origin_steps: int = 50, - prompt_embeds: Optional[torch.FloatTensor] = None, - output_type: Optional[str] = "pil", - return_dict: bool = True, - cross_attention_kwargs: Optional[Dict[str, Any]] = None, - ): - # 0. Default height and width to unet - height = height or self.unet.config.sample_size * self.vae_scale_factor - width = width or self.unet.config.sample_size * self.vae_scale_factor - - # 2. Define call parameters - if prompt is not None and isinstance(prompt, str): - batch_size = 1 - elif prompt is not None and isinstance(prompt, list): - batch_size = len(prompt) - else: - batch_size = prompt_embeds.shape[0] - - device = get_torch_device() - # do_classifier_free_guidance = guidance_scale > 0.0 # In LCM Implementation: cfg_noise = noise_cond + cfg_scale * (noise_cond - noise_uncond) , (cfg_scale > 0.0 using CFG) - - # 3. Encode input prompt - prompt_embeds = self._encode_prompt( - prompt, - device, - num_images_per_prompt, - prompt_embeds=prompt_embeds, - ) - - # 4. Prepare timesteps - self.scheduler.set_timesteps(num_inference_steps, lcm_origin_steps) - timesteps = self.scheduler.timesteps - - # 5. Prepare latent variable - num_channels_latents = self.unet.config.in_channels - latents = self.prepare_latents( - batch_size * num_images_per_prompt, - num_channels_latents, - height, - width, - prompt_embeds.dtype, - device, - latents, - ) - bs = batch_size * num_images_per_prompt - - # 6. Get Guidance Scale Embedding - w = torch.tensor(guidance_scale).repeat(bs) - w_embedding = self.get_w_embedding(w, embedding_dim=256).to( - device=device, dtype=latents.dtype - ) - - # 7. LCM MultiStep Sampling Loop: - with self.progress_bar(total=num_inference_steps) as progress_bar: - for i, t in enumerate(timesteps): - ts = torch.full((bs,), t, device=device, dtype=torch.long) - latents = latents.to(prompt_embeds.dtype) - - # model prediction (v-prediction, eps, x) - print("latents", latents.shape) - model_pred = self.unet( - latents, - ts, - timestep_cond=w_embedding, - encoder_hidden_states=prompt_embeds, - cross_attention_kwargs=cross_attention_kwargs, - return_dict=False, - )[0] - print("model_pred", model_pred.shape) - - # compute the previous noisy sample x_t -> x_t-1 - latents, denoised = self.scheduler.step( - model_pred, i, t, latents, return_dict=False - ) - - # # call the callback, if provided - # if i == len(timesteps) - 1: - progress_bar.update() - - denoised = denoised.to(prompt_embeds.dtype) - if not output_type == "latent": - image = self.vae.decode( - denoised / self.vae.config.scaling_factor, return_dict=False - )[0] - image, has_nsfw_concept = self.run_safety_checker( - image, device, prompt_embeds.dtype - ) - else: - image = denoised - has_nsfw_concept = None - - if has_nsfw_concept is None: - do_denormalize = [True] * image.shape[0] - else: - do_denormalize = [not has_nsfw for has_nsfw in has_nsfw_concept] - - image = self.image_processor.postprocess( - image, output_type=output_type, do_denormalize=do_denormalize - ) - - if not return_dict: - return (image, has_nsfw_concept) - - return StableDiffusionPipelineOutput( - images=image, nsfw_content_detected=has_nsfw_concept - ) diff --git a/coreml_suite/lcm/lcm_sampler.py b/coreml_suite/lcm/lcm_sampler.py index 5b90321..984e725 100644 --- a/coreml_suite/lcm/lcm_sampler.py +++ b/coreml_suite/lcm/lcm_sampler.py @@ -5,97 +5,15 @@ import torch from diffusers.utils.torch_utils import randn_tensor from tqdm import tqdm -import comfy.utils import latent_preview from comfy.model_management import get_torch_device from comfy.model_patcher import ModelPatcher -from coreml_suite.lcm.lcm_pipeline import LatentConsistencyModelPipeline 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_Simple: - def __init__(self): - self.scheduler = LCMScheduler.from_pretrained( - os.path.join(os.path.dirname(__file__), "scheduler_config.json") - ) - self.pipe = None - - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "coreml_model": ("COREML_UNET",), - "seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}), - "steps": ("INT", {"default": 4, "min": 1, "max": 10000}), - "cfg": ( - "FLOAT", - { - "default": 8.0, - "min": 0.0, - "max": 100.0, - "step": 0.5, - "round": 0.01, - }, - ), - "num_images": ("INT", {"default": 1, "min": 1, "max": 64}), - "positive_prompt": ("STRING", {"multiline": True}), - } - } - - RETURN_TYPES = ("IMAGE",) - FUNCTION = "sample" - CATEGORY = "sampling" - - def sample( - self, - coreml_model, - seed, - steps, - cfg, - positive_prompt, - num_images, - ): - height = coreml_model.expected_inputs["sample"]["shape"][2] * 8 - width = coreml_model.expected_inputs["sample"]["shape"][3] * 8 - - model_config = get_model_config() - wrapped_model = CoreMLModelWrapperLCM(model_config, coreml_model) - - if self.pipe is None: - self.pipe = LatentConsistencyModelPipeline.from_pretrained( - pretrained_model_name_or_path="SimianLuo/LCM_Dreamshaper_v7", - scheduler=self.scheduler, - safety_checker=None, - ) - - self.pipe.to(torch_device=get_torch_device(), torch_dtype=torch.float16) - - coreml_unet = wrapped_model - coreml_unet.config = self.pipe.unet.config - - self.pipe.unet = coreml_unet - - torch.manual_seed(seed) - - result = self.pipe( - prompt=positive_prompt, - width=width, - height=height, - guidance_scale=cfg, - num_inference_steps=steps, - num_images_per_prompt=num_images, - lcm_origin_steps=50, - output_type="np", - ).images - - images_tensor = torch.from_numpy(result) - - return (images_tensor,) - - class CoreMLSamplerLCM(CoreMLSampler): @classmethod def INPUT_TYPES(s): diff --git a/coreml_suite/lcm/unet.py b/coreml_suite/lcm/unet.py index b756e12..729d41f 100644 --- a/coreml_suite/lcm/unet.py +++ b/coreml_suite/lcm/unet.py @@ -1,4 +1,3 @@ -from diffusers.configuration_utils import register_to_config from overrides import overrides from python_coreml_stable_diffusion.unet import UNet2DConditionModel, TimestepEmbedding From eeae4bd6e306de0a495605d84e6f911f678b3abc Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Fri, 3 Nov 2023 01:27:44 +0100 Subject: [PATCH 12/12] Adjust default values for LCM nodes --- coreml_suite/lcm/lcm_sampler.py | 1 + coreml_suite/lcm/nodes.py | 2 +- 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/coreml_suite/lcm/lcm_sampler.py b/coreml_suite/lcm/lcm_sampler.py index 984e725..3f0796e 100644 --- a/coreml_suite/lcm/lcm_sampler.py +++ b/coreml_suite/lcm/lcm_sampler.py @@ -18,6 +18,7 @@ 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") diff --git a/coreml_suite/lcm/nodes.py b/coreml_suite/lcm/nodes.py index 6172911..7e84339 100644 --- a/coreml_suite/lcm/nodes.py +++ b/coreml_suite/lcm/nodes.py @@ -15,7 +15,7 @@ class CoreMLConverterLCM: "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": 4, "min": 1, "max": 64}), + "batch_size": ("INT", {"default": 1, "min": 1, "max": 64}), "compute_unit": ( [ ComputeUnit.CPU_AND_NE.name,