img2img works
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
@@ -7,7 +7,11 @@ import torch
|
||||
from comfy.model_management import get_torch_device
|
||||
from coreml_suite.latents import chunk_batch, merge_chunks
|
||||
from coreml_suite.controlnet import chunk_control
|
||||
from coreml_suite.models import CoreMLModelWrapper, get_model_config
|
||||
from coreml_suite.models import (
|
||||
CoreMLModelWrapper,
|
||||
get_model_config,
|
||||
CoreMLModelWrapperLCM,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -16,6 +20,7 @@ def coreml_model():
|
||||
model.expected_inputs = {
|
||||
"sample": {"shape": (2, 4, 64, 64)},
|
||||
"timestep": {"shape": (2,)},
|
||||
"timestep_cond": {"shape": (2, 256)},
|
||||
"encoder_hidden_states": {"shape": (2, 768, 1, 77)},
|
||||
"additional_residual_0": {"shape": (2, 320, 64, 64)},
|
||||
"additional_residual_1": {"shape": (2, 640, 32, 32)},
|
||||
@@ -54,6 +59,22 @@ def test_merge_chunks(batch_size):
|
||||
assert torch.equal(input_tensor, merged)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def inputs():
|
||||
x = torch.randn(1, 4, 64, 64).to(get_torch_device())
|
||||
t = torch.randn([1]).to(get_torch_device())
|
||||
c_crossattn = torch.randn(1, 77, 768).to(get_torch_device())
|
||||
control = {
|
||||
"output": [
|
||||
torch.randn(1, 320, 64, 64).to(get_torch_device()),
|
||||
torch.randn(1, 640, 32, 32).to(get_torch_device()),
|
||||
],
|
||||
}
|
||||
timestep_cond = torch.randn(1, 256).to(get_torch_device())
|
||||
|
||||
return x, t, c_crossattn, control, timestep_cond
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"b, target_size, num_chunks",
|
||||
[
|
||||
@@ -95,29 +116,22 @@ def test_chunking_no_control():
|
||||
assert chunked == [None, None]
|
||||
|
||||
|
||||
def test_chunking_inputs(coreml_model, model_config):
|
||||
def test_chunking_inputs(coreml_model, model_config, inputs):
|
||||
model = CoreMLModelWrapper(model_config, coreml_model)
|
||||
x = torch.randn(1, 4, 64, 64).to(get_torch_device())
|
||||
t = torch.randn([1]).to(get_torch_device())
|
||||
c_crossattn = torch.randn(1, 77, 768).to(get_torch_device())
|
||||
control = {
|
||||
"output": [
|
||||
torch.randn(1, 320, 64, 64).to(get_torch_device()),
|
||||
torch.randn(1, 640, 32, 32).to(get_torch_device()),
|
||||
],
|
||||
}
|
||||
|
||||
chunked_x, ts, chunked_context, chunked_control = model.chunk_inputs(
|
||||
x, t, c_crossattn, control
|
||||
chunked_x, ts, chunked_context, chunked_cn, chunked_ts_cond = model.chunk_inputs(
|
||||
*inputs
|
||||
)
|
||||
|
||||
assert len(chunked_x) == 1
|
||||
assert len(ts) == 1
|
||||
assert len(chunked_context) == 1
|
||||
assert len(chunked_control) == 1
|
||||
assert len(chunked_cn) == 1
|
||||
assert len(chunked_ts_cond) == 1
|
||||
|
||||
assert chunked_x[0].shape == (2, 4, 64, 64)
|
||||
assert ts[0].shape == (2,)
|
||||
assert chunked_context[0].shape == (2, 77, 768)
|
||||
assert chunked_control[0]["output"][0].shape == (2, 320, 64, 64)
|
||||
assert chunked_control[0]["output"][1].shape == (2, 640, 32, 32)
|
||||
assert chunked_cn[0]["output"][0].shape == (2, 320, 64, 64)
|
||||
assert chunked_cn[0]["output"][1].shape == (2, 640, 32, 32)
|
||||
assert chunked_ts_cond[0].shape == (2, 256)
|
||||
|
||||
Reference in New Issue
Block a user