From 44a380ffdf084e9c1fd788ddf6717387a6df3ef3 Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Fri, 3 Nov 2023 00:27:15 +0100 Subject: [PATCH] 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)