img2img works

This commit is contained in:
aszc-dev
2023-11-03 01:29:32 +01:00
parent e22d8187cd
commit 44a380ffdf
8 changed files with 249 additions and 84 deletions
+2 -4
View File
@@ -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)
+5 -7
View File
@@ -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
+2
View File
@@ -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(
+69 -37
View File
@@ -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
+4 -2
View File
@@ -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:
+100
View File
@@ -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,)
+37 -18
View File
@@ -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),
)
+30 -16
View File
@@ -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)