first commit
This commit is contained in:
@@ -0,0 +1 @@
|
||||
__pycache__
|
||||
@@ -1,2 +1,46 @@
|
||||
# DiffMorpher-ComfyUI
|
||||
a custom node for [DiffMorpher](https://github.com/Kevin-thu/DiffMorpher.git)
|
||||
a custom node for [DiffMorpher](https://github.com/Kevin-thu/DiffMorpher.git),you can find base workflow in [`doc`](./doc/)
|
||||
|
||||
## Example
|
||||
image_0 | image_1 | output
|
||||
----- | ---- | ----
|
||||
 |  | 
|
||||
|
||||
## How to use
|
||||
```
|
||||
# in ComfyUI/custom_nodes
|
||||
git clone https://github.com/AIFSH/DiffMorpher-ComfyUI.git
|
||||
cd DiffMorpher-ComfyUI
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
weights will be downloaded from huggingface
|
||||
|
||||
## Tutorial
|
||||
|
||||
|
||||
DiffMorpherNode
|
||||
|
||||
required
|
||||
- `image__0`: the first image (default: "")
|
||||
- `prompt_0`: Prompt of the first image (default: "")
|
||||
- `image_1`: the second image (default: "")
|
||||
- `prompt_1`: Prompt of the second image (default: "")
|
||||
- `use_adain`: Use AdaIN (default: False)
|
||||
- `use_reschedule`: Use reschedule sampling (default: False)
|
||||
- `lamb`: Hyperparameter $\lambda \in [0,1]$ for self-attention replacement, where a larger $\lambda$ indicates more replacements (default: 0.6)
|
||||
- `save_inter`: Save intermediate results (default: False) if True, frame saved in `ComfyUI/output/diffmorpher`
|
||||
- `num_frames`: Number of frames to generate (default: 50)
|
||||
- `duration`: Duration of each frame (default: 50)
|
||||
|
||||
optional
|
||||
- `model_path`: Pretrained model path (default: "stabilityai/stable-diffusion-2-1-base")
|
||||
- `lora_0`: Path of the lora directory of the first image (default: "")
|
||||
- `lora_1`: Path of the lora directory of the second image (default: "")
|
||||
|
||||
|
||||
## ask for answer as soon as you want
|
||||
wechat: aifsh_98
|
||||
need donate if you mand it,
|
||||
but please feel free to new issue for answering
|
||||
|
||||
Windows环境配置太难?可以添加微信:aifsh_98,赞赏获取Windows一键包,当然你也可以提issue等待大佬为你答疑解惑。
|
||||
+181
@@ -0,0 +1,181 @@
|
||||
import os
|
||||
import folder_paths
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
import math
|
||||
import torch
|
||||
import time
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from .diffmorpher.model import DiffMorpherPipeline
|
||||
|
||||
output_dir = folder_paths.get_output_directory()
|
||||
|
||||
def get_64x_num(num):
|
||||
return math.ceil(num / 64) * 64
|
||||
|
||||
def crop_and_resize(image, height, width):
|
||||
image = np.array(image)
|
||||
image_height, image_width, _ = image.shape
|
||||
if image_height / image_width < height / width:
|
||||
croped_width = int(image_height / height * width)
|
||||
left = (image_width - croped_width) // 2
|
||||
image = image[:, left: left+croped_width]
|
||||
image = Image.fromarray(image).resize((width, height))
|
||||
else:
|
||||
croped_height = int(image_width / width * height)
|
||||
left = (image_height - croped_height) // 2
|
||||
image = image[left: left+croped_height, :]
|
||||
image = Image.fromarray(image).resize((width, height))
|
||||
return image
|
||||
|
||||
class DiffMorpherNode:
|
||||
def __init__(self) -> None:
|
||||
self.model_path = None
|
||||
self.pipeline = None
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required":{
|
||||
"image_0":("IMAGE",),
|
||||
"prompt_0":("TEXT",),
|
||||
"image_1":("IMAGE",),
|
||||
"prompt_1":("TEXT",),
|
||||
"num_frames":("INT",{
|
||||
"default":16
|
||||
}),
|
||||
"duration":("INT",{
|
||||
"default":100
|
||||
}),
|
||||
"use_adain":("BOOLEAN",{
|
||||
"default":True
|
||||
}),
|
||||
"use_reschedule":("BOOLEAN",{
|
||||
"default":True
|
||||
}),
|
||||
"lamb":("FLOAT",{
|
||||
"min": 0.0,
|
||||
"max":1.0,
|
||||
"default":0.6,
|
||||
"display": "slider"
|
||||
}),
|
||||
"save_inter":("BOOLEAN",{
|
||||
"default":True
|
||||
}),
|
||||
},
|
||||
"optional":{
|
||||
"diffusers_model":(folder_paths.get_filename_list("diffusers"),),
|
||||
"lora_0":(folder_paths.get_filename_list("loras"),),
|
||||
"lora_1":(folder_paths.get_filename_list("loras"),)
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("GIF",)
|
||||
#RETURN_NAMES = ("image_output_name",)
|
||||
|
||||
FUNCTION = "generate"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "AIFSH_DiffMorpher"
|
||||
|
||||
def comfy2image(self,image):
|
||||
image = image.numpy()[0] * 255
|
||||
image = image.astype(np.uint8)
|
||||
image_pil = Image.fromarray(image).convert("RGB")
|
||||
org_w, org_h = image_pil.size
|
||||
height, width = (768,get_64x_num(768*org_w/org_h)) if org_h > org_w else (get_64x_num(768*org_h/org_w),768)
|
||||
print(f"crop and resize from ({org_w},{org_h}) to ({width},{height})")
|
||||
return crop_and_resize(image_pil,height,width)
|
||||
|
||||
def generate(self,image_0,prompt_0,image_1,prompt_1,num_frames,duration,
|
||||
use_adain,use_reschedule,lamb,save_inter,diffusers_model=None,lora_0=None,lora_1=None):
|
||||
|
||||
if diffusers_model is None:
|
||||
# stabilityai/stable-diffusion-2-1-base
|
||||
model_path = os.path.join(folder_paths.models_dir, "diffusers","stable-diffusion-2-1-base")
|
||||
snapshot_download(repo_id="stabilityai/stable-diffusion-2-1-base",
|
||||
allow_patterns=["*.json","*.txt","*.fp16.safetensors"],
|
||||
local_dir=model_path)
|
||||
|
||||
else:
|
||||
model_path = folder_paths.get_full_path("diffusers",diffusers_model)
|
||||
|
||||
if self.model_path != model_path:
|
||||
self.model_path = model_path
|
||||
self.pipeline = DiffMorpherPipeline.from_pretrained(self.model_path, torch_dtype=torch.float16,variant="fp16",use_safetensors=True)
|
||||
|
||||
self.pipeline.to("cuda")
|
||||
out_dir = os.path.join(output_dir,"diffmorpher")
|
||||
os.makedirs(out_dir,exist_ok=True)
|
||||
images = self.pipeline(img_0=self.comfy2image(image_0),
|
||||
img_1=self.comfy2image(image_1),
|
||||
prompt_0=prompt_0,prompt_1=prompt_1,
|
||||
save_lora_dir=os.path.join(folder_paths.models_dir, "loras"),
|
||||
load_lora_path_0=lora_0,load_lora_path_1=lora_1,
|
||||
use_adain=use_adain,
|
||||
use_reschedule=use_reschedule,
|
||||
lamd=lamb,
|
||||
output_path=out_dir,
|
||||
num_frames=num_frames,
|
||||
save_intermediates=save_inter,
|
||||
use_lora= lora_0 and lora_1,
|
||||
)
|
||||
output_path = os.path.join(output_dir,f"diffmorpher_{time.time_ns()}.gif")
|
||||
images[0].save(output_path, save_all=True,
|
||||
append_images=images[1:], duration=duration, loop=0)
|
||||
|
||||
return (output_path,)
|
||||
|
||||
|
||||
|
||||
|
||||
class PreViewGIF:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":{
|
||||
"gif":("GIF",),
|
||||
}}
|
||||
|
||||
CATEGORY = "AIFSH_DiffMorpher"
|
||||
DESCRIPTION = "hello world!"
|
||||
|
||||
RETURN_TYPES = ()
|
||||
|
||||
OUTPUT_NODE = True
|
||||
|
||||
FUNCTION = "load_gif"
|
||||
|
||||
def load_gif(self, gif):
|
||||
video_name = os.path.basename(gif)
|
||||
video_path_name = os.path.basename(os.path.dirname(gif))
|
||||
return {"ui":{"gif":[video_name,video_path_name]}}
|
||||
|
||||
|
||||
class TextNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"text": ("STRING", {"multiline": True, "dynamicPrompts": True})
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("TEXT",)
|
||||
#RETURN_NAMES = ("image_output_name",)
|
||||
|
||||
FUNCTION = "text"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "AIFSH_DiffMorpher"
|
||||
|
||||
def text(self,text):
|
||||
return (text,)
|
||||
|
||||
WEB_DIRECTORY = "./web"
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"TextNode":TextNode,
|
||||
"PreViewGIF":PreViewGIF,
|
||||
"DiffMorpherNode": DiffMorpherNode
|
||||
}
|
||||
@@ -0,0 +1,643 @@
|
||||
import os
|
||||
from diffusers.models import AutoencoderKL, UNet2DConditionModel
|
||||
from diffusers.models.attention_processor import AttnProcessor
|
||||
from diffusers.pipelines.stable_diffusion.safety_checker import StableDiffusionSafetyChecker
|
||||
from diffusers.schedulers import KarrasDiffusionSchedulers
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import tqdm
|
||||
import numpy as np
|
||||
import safetensors
|
||||
from PIL import Image
|
||||
from torchvision import transforms
|
||||
from transformers import CLIPImageProcessor, CLIPTextModel, CLIPTokenizer
|
||||
from diffusers import StableDiffusionPipeline
|
||||
from argparse import ArgumentParser
|
||||
import inspect
|
||||
|
||||
from .utils.model_utils import get_img, slerp, do_replace_attn
|
||||
from .utils.lora_utils import train_lora, load_lora
|
||||
from .utils.alpha_scheduler import AlphaScheduler
|
||||
|
||||
|
||||
class StoreProcessor():
|
||||
def __init__(self, original_processor, value_dict, name):
|
||||
self.original_processor = original_processor
|
||||
self.value_dict = value_dict
|
||||
self.name = name
|
||||
self.value_dict[self.name] = dict()
|
||||
self.id = 0
|
||||
|
||||
def __call__(self, attn, hidden_states, *args, encoder_hidden_states=None, attention_mask=None, **kwargs):
|
||||
# Is self attention
|
||||
if encoder_hidden_states is None:
|
||||
self.value_dict[self.name][self.id] = hidden_states.detach()
|
||||
self.id += 1
|
||||
res = self.original_processor(attn, hidden_states, *args,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
**kwargs)
|
||||
|
||||
return res
|
||||
|
||||
|
||||
class LoadProcessor():
|
||||
def __init__(self, original_processor, name, img0_dict, img1_dict, alpha, beta=0, lamd=0.6):
|
||||
super().__init__()
|
||||
self.original_processor = original_processor
|
||||
self.name = name
|
||||
self.img0_dict = img0_dict
|
||||
self.img1_dict = img1_dict
|
||||
self.alpha = alpha
|
||||
self.beta = beta
|
||||
self.lamd = lamd
|
||||
self.id = 0
|
||||
|
||||
def __call__(self, attn, hidden_states, *args, encoder_hidden_states=None, attention_mask=None, **kwargs):
|
||||
# Is self attention
|
||||
if encoder_hidden_states is None:
|
||||
if self.id < 50 * self.lamd:
|
||||
map0 = self.img0_dict[self.name][self.id]
|
||||
map1 = self.img1_dict[self.name][self.id]
|
||||
cross_map = self.beta * hidden_states + \
|
||||
(1 - self.beta) * ((1 - self.alpha) * map0 + self.alpha * map1)
|
||||
# cross_map = self.beta * hidden_states + \
|
||||
# (1 - self.beta) * slerp(map0, map1, self.alpha)
|
||||
# cross_map = slerp(slerp(map0, map1, self.alpha),
|
||||
# hidden_states, self.beta)
|
||||
# cross_map = hidden_states
|
||||
# cross_map = torch.cat(
|
||||
# ((1 - self.alpha) * map0, self.alpha * map1), dim=1)
|
||||
|
||||
res = self.original_processor(attn, hidden_states, *args,
|
||||
encoder_hidden_states=cross_map,
|
||||
attention_mask=attention_mask,
|
||||
**kwargs)
|
||||
else:
|
||||
res = self.original_processor(attn, hidden_states, *args,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
**kwargs)
|
||||
|
||||
self.id += 1
|
||||
# if self.id == len(self.img0_dict[self.name]):
|
||||
if self.id == len(self.img0_dict[self.name]):
|
||||
self.id = 0
|
||||
else:
|
||||
res = self.original_processor(attn, hidden_states, *args,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
**kwargs)
|
||||
|
||||
return res
|
||||
|
||||
|
||||
class DiffMorpherPipeline(StableDiffusionPipeline):
|
||||
|
||||
def __init__(self,
|
||||
vae: AutoencoderKL,
|
||||
text_encoder: CLIPTextModel,
|
||||
tokenizer: CLIPTokenizer,
|
||||
unet: UNet2DConditionModel,
|
||||
scheduler: KarrasDiffusionSchedulers,
|
||||
safety_checker: StableDiffusionSafetyChecker,
|
||||
feature_extractor: CLIPImageProcessor,
|
||||
image_encoder=None,
|
||||
requires_safety_checker: bool = True,
|
||||
):
|
||||
sig = inspect.signature(super().__init__)
|
||||
params = sig.parameters
|
||||
if 'image_encoder' in params:
|
||||
super().__init__(vae, text_encoder, tokenizer, unet, scheduler,
|
||||
safety_checker, feature_extractor, image_encoder, requires_safety_checker)
|
||||
else:
|
||||
super().__init__(vae, text_encoder, tokenizer, unet, scheduler,
|
||||
safety_checker, feature_extractor, requires_safety_checker)
|
||||
self.img0_dict = dict()
|
||||
self.img1_dict = dict()
|
||||
|
||||
def inv_step(
|
||||
self,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: int,
|
||||
x: torch.FloatTensor,
|
||||
eta=0.,
|
||||
verbose=False
|
||||
):
|
||||
"""
|
||||
Inverse sampling for DDIM Inversion
|
||||
"""
|
||||
if verbose:
|
||||
print("timestep: ", timestep)
|
||||
next_step = timestep
|
||||
timestep = min(timestep - self.scheduler.config.num_train_timesteps //
|
||||
self.scheduler.num_inference_steps, 999)
|
||||
alpha_prod_t = self.scheduler.alphas_cumprod[
|
||||
timestep] if timestep >= 0 else self.scheduler.final_alpha_cumprod
|
||||
alpha_prod_t_next = self.scheduler.alphas_cumprod[next_step]
|
||||
beta_prod_t = 1 - alpha_prod_t
|
||||
pred_x0 = (x - beta_prod_t**0.5 * model_output) / alpha_prod_t**0.5
|
||||
pred_dir = (1 - alpha_prod_t_next)**0.5 * model_output
|
||||
x_next = alpha_prod_t_next**0.5 * pred_x0 + pred_dir
|
||||
return x_next, pred_x0
|
||||
|
||||
@torch.no_grad()
|
||||
def invert(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
prompt,
|
||||
num_inference_steps=50,
|
||||
num_actual_inference_steps=None,
|
||||
guidance_scale=1.,
|
||||
eta=0.0,
|
||||
**kwds):
|
||||
"""
|
||||
invert a real image into noise map with determinisc DDIM inversion
|
||||
"""
|
||||
DEVICE = torch.device(
|
||||
"cuda") if torch.cuda.is_available() else torch.device("cpu")
|
||||
batch_size = image.shape[0]
|
||||
if isinstance(prompt, list):
|
||||
if batch_size == 1:
|
||||
image = image.expand(len(prompt), -1, -1, -1)
|
||||
elif isinstance(prompt, str):
|
||||
if batch_size > 1:
|
||||
prompt = [prompt] * batch_size
|
||||
|
||||
# text embeddings
|
||||
text_input = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=77,
|
||||
return_tensors="pt"
|
||||
)
|
||||
text_embeddings = self.text_encoder(text_input.input_ids.to(DEVICE))[0]
|
||||
print("input text embeddings :", text_embeddings.shape)
|
||||
# define initial latents
|
||||
latents = self.image2latent(image)
|
||||
|
||||
# unconditional embedding for classifier free guidance
|
||||
if guidance_scale > 1.:
|
||||
max_length = text_input.input_ids.shape[-1]
|
||||
unconditional_input = self.tokenizer(
|
||||
[""] * batch_size,
|
||||
padding="max_length",
|
||||
max_length=77,
|
||||
return_tensors="pt"
|
||||
)
|
||||
unconditional_embeddings = self.text_encoder(
|
||||
unconditional_input.input_ids.to(DEVICE))[0]
|
||||
text_embeddings = torch.cat(
|
||||
[unconditional_embeddings, text_embeddings], dim=0)
|
||||
|
||||
print("latents shape: ", latents.shape)
|
||||
# interative sampling
|
||||
self.scheduler.set_timesteps(num_inference_steps)
|
||||
print("Valid timesteps: ", reversed(self.scheduler.timesteps))
|
||||
# print("attributes: ", self.scheduler.__dict__)
|
||||
latents_list = [latents]
|
||||
pred_x0_list = [latents]
|
||||
for i, t in enumerate(tqdm.tqdm(reversed(self.scheduler.timesteps), desc="DDIM Inversion")):
|
||||
if num_actual_inference_steps is not None and i >= num_actual_inference_steps:
|
||||
continue
|
||||
|
||||
if guidance_scale > 1.:
|
||||
model_inputs = torch.cat([latents] * 2)
|
||||
else:
|
||||
model_inputs = latents
|
||||
|
||||
# predict the noise
|
||||
noise_pred = self.unet(
|
||||
model_inputs, t, encoder_hidden_states=text_embeddings).sample
|
||||
if guidance_scale > 1.:
|
||||
noise_pred_uncon, noise_pred_con = noise_pred.chunk(2, dim=0)
|
||||
noise_pred = noise_pred_uncon + guidance_scale * \
|
||||
(noise_pred_con - noise_pred_uncon)
|
||||
# compute the previous noise sample x_t-1 -> x_t
|
||||
latents, pred_x0 = self.inv_step(noise_pred, t, latents)
|
||||
latents_list.append(latents)
|
||||
pred_x0_list.append(pred_x0)
|
||||
|
||||
return latents
|
||||
|
||||
@torch.no_grad()
|
||||
def ddim_inversion(self, latent, cond):
|
||||
timesteps = reversed(self.scheduler.timesteps)
|
||||
with torch.autocast(device_type='cuda', dtype=torch.float32):
|
||||
for i, t in enumerate(tqdm.tqdm(timesteps, desc="DDIM inversion")):
|
||||
cond_batch = cond.repeat(latent.shape[0], 1, 1)
|
||||
|
||||
alpha_prod_t = self.scheduler.alphas_cumprod[t]
|
||||
alpha_prod_t_prev = (
|
||||
self.scheduler.alphas_cumprod[timesteps[i - 1]]
|
||||
if i > 0 else self.scheduler.final_alpha_cumprod
|
||||
)
|
||||
|
||||
mu = alpha_prod_t ** 0.5
|
||||
mu_prev = alpha_prod_t_prev ** 0.5
|
||||
sigma = (1 - alpha_prod_t) ** 0.5
|
||||
sigma_prev = (1 - alpha_prod_t_prev) ** 0.5
|
||||
|
||||
eps = self.unet(
|
||||
latent, t, encoder_hidden_states=cond_batch).sample
|
||||
|
||||
pred_x0 = (latent - sigma_prev * eps) / mu_prev
|
||||
latent = mu * pred_x0 + sigma * eps
|
||||
# if save_latents:
|
||||
# torch.save(latent, os.path.join(save_path, f'noisy_latents_{t}.pt'))
|
||||
# torch.save(latent, os.path.join(save_path, f'noisy_latents_{t}.pt'))
|
||||
return latent
|
||||
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: int,
|
||||
x: torch.FloatTensor,
|
||||
):
|
||||
"""
|
||||
predict the sample of the next step in the denoise process.
|
||||
"""
|
||||
prev_timestep = timestep - \
|
||||
self.scheduler.config.num_train_timesteps // self.scheduler.num_inference_steps
|
||||
alpha_prod_t = self.scheduler.alphas_cumprod[timestep]
|
||||
alpha_prod_t_prev = self.scheduler.alphas_cumprod[
|
||||
prev_timestep] if prev_timestep > 0 else self.scheduler.final_alpha_cumprod
|
||||
beta_prod_t = 1 - alpha_prod_t
|
||||
pred_x0 = (x - beta_prod_t**0.5 * model_output) / alpha_prod_t**0.5
|
||||
pred_dir = (1 - alpha_prod_t_prev)**0.5 * model_output
|
||||
x_prev = alpha_prod_t_prev**0.5 * pred_x0 + pred_dir
|
||||
return x_prev, pred_x0
|
||||
|
||||
@torch.no_grad()
|
||||
def image2latent(self, image):
|
||||
DEVICE = torch.device(
|
||||
"cuda") if torch.cuda.is_available() else torch.device("cpu")
|
||||
if type(image) is Image:
|
||||
image = np.array(image)
|
||||
image = torch.from_numpy(image).float() / 127.5 - 1
|
||||
image = image.permute(2, 0, 1).unsqueeze(0)
|
||||
|
||||
image = image.half()
|
||||
# input image density range [-1, 1]
|
||||
latents = self.vae.encode(image.to(DEVICE))['latent_dist'].mean
|
||||
latents = latents * 0.18215
|
||||
return latents
|
||||
|
||||
@torch.no_grad()
|
||||
def latent2image(self, latents, return_type='np'):
|
||||
latents = 1 / 0.18215 * latents.detach()
|
||||
image = self.vae.decode(latents)['sample']
|
||||
if return_type == 'np':
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
image = image.cpu().permute(0, 2, 3, 1).numpy()[0]
|
||||
image = (image * 255).astype(np.uint8)
|
||||
elif return_type == "pt":
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
|
||||
return image
|
||||
|
||||
def latent2image_grad(self, latents):
|
||||
latents = 1 / 0.18215 * latents
|
||||
image = self.vae.decode(latents)['sample']
|
||||
|
||||
return image # range [-1, 1]
|
||||
|
||||
@torch.no_grad()
|
||||
def cal_latent(self, num_inference_steps, guidance_scale, unconditioning, img_noise_0, img_noise_1, text_embeddings_0, text_embeddings_1, lora_0, lora_1, alpha, use_lora, fix_lora=None):
|
||||
# latents = torch.cos(alpha * torch.pi / 2) * img_noise_0 + \
|
||||
# torch.sin(alpha * torch.pi / 2) * img_noise_1
|
||||
# latents = (1 - alpha) * img_noise_0 + alpha * img_noise_1
|
||||
# latents = latents / ((1 - alpha) ** 2 + alpha ** 2)
|
||||
latents = slerp(img_noise_0, img_noise_1, alpha, self.use_adain)
|
||||
latents = latents.half()
|
||||
text_embeddings = (1 - alpha) * text_embeddings_0 + \
|
||||
alpha * text_embeddings_1
|
||||
text_embeddings = text_embeddings.half()
|
||||
|
||||
self.scheduler.set_timesteps(num_inference_steps)
|
||||
if use_lora:
|
||||
if fix_lora is not None:
|
||||
self.unet = load_lora(self.unet, lora_0, lora_1, fix_lora)
|
||||
else:
|
||||
self.unet = load_lora(self.unet, lora_0, lora_1, alpha)
|
||||
|
||||
for i, t in enumerate(tqdm.tqdm(self.scheduler.timesteps, desc=f"DDIM Sampler, alpha={alpha}")):
|
||||
|
||||
if guidance_scale > 1.:
|
||||
model_inputs = torch.cat([latents] * 2)
|
||||
else:
|
||||
model_inputs = latents
|
||||
if unconditioning is not None and isinstance(unconditioning, list):
|
||||
_, text_embeddings = text_embeddings.chunk(2)
|
||||
text_embeddings = torch.cat(
|
||||
[unconditioning[i].expand(*text_embeddings.shape), text_embeddings])
|
||||
# predict the noise
|
||||
noise_pred = self.unet(
|
||||
model_inputs, t, encoder_hidden_states=text_embeddings).sample
|
||||
if guidance_scale > 1.0:
|
||||
noise_pred_uncon, noise_pred_con = noise_pred.chunk(
|
||||
2, dim=0)
|
||||
noise_pred = noise_pred_uncon + guidance_scale * \
|
||||
(noise_pred_con - noise_pred_uncon)
|
||||
# compute the previous noise sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(
|
||||
noise_pred, t, latents, return_dict=False)[0]
|
||||
return latents
|
||||
|
||||
@torch.no_grad()
|
||||
def get_text_embeddings(self, prompt, guidance_scale, neg_prompt, batch_size):
|
||||
DEVICE = torch.device(
|
||||
"cuda") if torch.cuda.is_available() else torch.device("cpu")
|
||||
# text embeddings
|
||||
text_input = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=77,
|
||||
return_tensors="pt"
|
||||
)
|
||||
text_embeddings = self.text_encoder(text_input.input_ids.cuda())[0]
|
||||
|
||||
if guidance_scale > 1.:
|
||||
if neg_prompt:
|
||||
uc_text = neg_prompt
|
||||
else:
|
||||
uc_text = ""
|
||||
unconditional_input = self.tokenizer(
|
||||
[uc_text] * batch_size,
|
||||
padding="max_length",
|
||||
max_length=77,
|
||||
return_tensors="pt"
|
||||
)
|
||||
unconditional_embeddings = self.text_encoder(
|
||||
unconditional_input.input_ids.to(DEVICE))[0]
|
||||
text_embeddings = torch.cat(
|
||||
[unconditional_embeddings, text_embeddings], dim=0)
|
||||
|
||||
return text_embeddings
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
img_0=None,
|
||||
img_1=None,
|
||||
img_path_0=None,
|
||||
img_path_1=None,
|
||||
prompt_0="",
|
||||
prompt_1="",
|
||||
save_lora_dir="./lora",
|
||||
load_lora_path_0=None,
|
||||
load_lora_path_1=None,
|
||||
lora_steps=200,
|
||||
lora_lr=2e-4,
|
||||
lora_rank=16,
|
||||
batch_size=1,
|
||||
height=512,
|
||||
width=512,
|
||||
num_inference_steps=50,
|
||||
num_actual_inference_steps=None,
|
||||
guidance_scale=1,
|
||||
attn_beta=0,
|
||||
lamd=0.6,
|
||||
use_lora=True,
|
||||
use_adain=True,
|
||||
use_reschedule=True,
|
||||
output_path="./results",
|
||||
num_frames=50,
|
||||
fix_lora=None,
|
||||
progress=tqdm,
|
||||
unconditioning=None,
|
||||
neg_prompt=None,
|
||||
save_intermediates=False,
|
||||
**kwds):
|
||||
|
||||
# if isinstance(prompt, list):
|
||||
# batch_size = len(prompt)
|
||||
# elif isinstance(prompt, str):
|
||||
# if batch_size > 1:
|
||||
# prompt = [prompt] * batch_size
|
||||
self.scheduler.set_timesteps(num_inference_steps)
|
||||
self.use_lora = use_lora
|
||||
self.use_adain = use_adain
|
||||
self.use_reschedule = use_reschedule
|
||||
self.output_path = output_path
|
||||
|
||||
if img_0 is None:
|
||||
img_0 = Image.open(img_path_0).convert("RGB")
|
||||
# else:
|
||||
# img_0 = Image.fromarray(img_0).convert("RGB")
|
||||
|
||||
if img_1 is None:
|
||||
img_1 = Image.open(img_path_1).convert("RGB")
|
||||
# else:
|
||||
# img_1 = Image.fromarray(img_1).convert("RGB")
|
||||
|
||||
if self.use_lora:
|
||||
print("Loading lora...")
|
||||
if not load_lora_path_0:
|
||||
|
||||
weight_name = f"{output_path.split('/')[-1]}_lora_0.ckpt"
|
||||
load_lora_path_0 = save_lora_dir + "/" + weight_name
|
||||
if not os.path.exists(load_lora_path_0):
|
||||
train_lora(img_0, prompt_0, save_lora_dir, None, self.tokenizer, self.text_encoder,
|
||||
self.vae, self.unet, self.scheduler, lora_steps, lora_lr, lora_rank, weight_name=weight_name)
|
||||
print(f"Load from {load_lora_path_0}.")
|
||||
if load_lora_path_0.endswith(".safetensors"):
|
||||
lora_0 = safetensors.torch.load_file(
|
||||
load_lora_path_0, device="cpu")
|
||||
else:
|
||||
lora_0 = torch.load(load_lora_path_0, map_location="cpu")
|
||||
|
||||
if not load_lora_path_1:
|
||||
weight_name = f"{output_path.split('/')[-1]}_lora_1.ckpt"
|
||||
load_lora_path_1 = save_lora_dir + "/" + weight_name
|
||||
if not os.path.exists(load_lora_path_1):
|
||||
train_lora(img_1, prompt_1, save_lora_dir, None, self.tokenizer, self.text_encoder,
|
||||
self.vae, self.unet, self.scheduler, lora_steps, lora_lr, lora_rank, weight_name=weight_name)
|
||||
print(f"Load from {load_lora_path_1}.")
|
||||
if load_lora_path_1.endswith(".safetensors"):
|
||||
lora_1 = safetensors.torch.load_file(
|
||||
load_lora_path_1, device="cpu")
|
||||
else:
|
||||
lora_1 = torch.load(load_lora_path_1, map_location="cpu")
|
||||
else:
|
||||
lora_0 = lora_1 = None
|
||||
|
||||
text_embeddings_0 = self.get_text_embeddings(
|
||||
prompt_0, guidance_scale, neg_prompt, batch_size)
|
||||
text_embeddings_1 = self.get_text_embeddings(
|
||||
prompt_1, guidance_scale, neg_prompt, batch_size)
|
||||
img_0 = get_img(img_0)
|
||||
img_1 = get_img(img_1)
|
||||
if self.use_lora:
|
||||
self.unet = load_lora(self.unet, lora_0, lora_1, 0)
|
||||
img_noise_0 = self.ddim_inversion(
|
||||
self.image2latent(img_0), text_embeddings_0)
|
||||
if self.use_lora:
|
||||
self.unet = load_lora(self.unet, lora_0, lora_1, 1)
|
||||
img_noise_1 = self.ddim_inversion(
|
||||
self.image2latent(img_1), text_embeddings_1)
|
||||
|
||||
print("latents shape: ", img_noise_0.shape)
|
||||
|
||||
original_processor = list(self.unet.attn_processors.values())[0]
|
||||
|
||||
def morph(alpha_list, progress, desc):
|
||||
images = []
|
||||
if attn_beta is not None:
|
||||
if self.use_lora:
|
||||
self.unet = load_lora(
|
||||
self.unet, lora_0, lora_1, 0 if fix_lora is None else fix_lora)
|
||||
|
||||
attn_processor_dict = {}
|
||||
for k in self.unet.attn_processors.keys():
|
||||
if do_replace_attn(k):
|
||||
if self.use_lora:
|
||||
attn_processor_dict[k] = StoreProcessor(self.unet.attn_processors[k],
|
||||
self.img0_dict, k)
|
||||
else:
|
||||
attn_processor_dict[k] = StoreProcessor(original_processor,
|
||||
self.img0_dict, k)
|
||||
else:
|
||||
attn_processor_dict[k] = self.unet.attn_processors[k]
|
||||
self.unet.set_attn_processor(attn_processor_dict)
|
||||
|
||||
latents = self.cal_latent(
|
||||
num_inference_steps,
|
||||
guidance_scale,
|
||||
unconditioning,
|
||||
img_noise_0,
|
||||
img_noise_1,
|
||||
text_embeddings_0,
|
||||
text_embeddings_1,
|
||||
lora_0,
|
||||
lora_1,
|
||||
alpha_list[0],
|
||||
False,
|
||||
fix_lora
|
||||
)
|
||||
first_image = self.latent2image(latents)
|
||||
first_image = Image.fromarray(first_image)
|
||||
if save_intermediates:
|
||||
first_image.save(f"{self.output_path}/{0:02d}.png")
|
||||
|
||||
if self.use_lora:
|
||||
self.unet = load_lora(
|
||||
self.unet, lora_0, lora_1, 1 if fix_lora is None else fix_lora)
|
||||
attn_processor_dict = {}
|
||||
for k in self.unet.attn_processors.keys():
|
||||
if do_replace_attn(k):
|
||||
if self.use_lora:
|
||||
attn_processor_dict[k] = StoreProcessor(self.unet.attn_processors[k],
|
||||
self.img1_dict, k)
|
||||
else:
|
||||
attn_processor_dict[k] = StoreProcessor(original_processor,
|
||||
self.img1_dict, k)
|
||||
else:
|
||||
attn_processor_dict[k] = self.unet.attn_processors[k]
|
||||
|
||||
self.unet.set_attn_processor(attn_processor_dict)
|
||||
|
||||
latents = self.cal_latent(
|
||||
num_inference_steps,
|
||||
guidance_scale,
|
||||
unconditioning,
|
||||
img_noise_0,
|
||||
img_noise_1,
|
||||
text_embeddings_0,
|
||||
text_embeddings_1,
|
||||
lora_0,
|
||||
lora_1,
|
||||
alpha_list[-1],
|
||||
False,
|
||||
fix_lora
|
||||
)
|
||||
last_image = self.latent2image(latents)
|
||||
last_image = Image.fromarray(last_image)
|
||||
if save_intermediates:
|
||||
last_image.save(
|
||||
f"{self.output_path}/{num_frames - 1:02d}.png")
|
||||
|
||||
for i in progress.tqdm(range(1, num_frames - 1), desc=desc):
|
||||
alpha = alpha_list[i]
|
||||
if self.use_lora:
|
||||
self.unet = load_lora(
|
||||
self.unet, lora_0, lora_1, alpha if fix_lora is None else fix_lora)
|
||||
|
||||
attn_processor_dict = {}
|
||||
for k in self.unet.attn_processors.keys():
|
||||
if do_replace_attn(k):
|
||||
if self.use_lora:
|
||||
attn_processor_dict[k] = LoadProcessor(
|
||||
self.unet.attn_processors[k], k, self.img0_dict, self.img1_dict, alpha, attn_beta, lamd)
|
||||
else:
|
||||
attn_processor_dict[k] = LoadProcessor(
|
||||
original_processor, k, self.img0_dict, self.img1_dict, alpha, attn_beta, lamd)
|
||||
else:
|
||||
attn_processor_dict[k] = self.unet.attn_processors[k]
|
||||
|
||||
self.unet.set_attn_processor(attn_processor_dict)
|
||||
|
||||
latents = self.cal_latent(
|
||||
num_inference_steps,
|
||||
guidance_scale,
|
||||
unconditioning,
|
||||
img_noise_0,
|
||||
img_noise_1,
|
||||
text_embeddings_0,
|
||||
text_embeddings_1,
|
||||
lora_0,
|
||||
lora_1,
|
||||
alpha_list[i],
|
||||
False,
|
||||
fix_lora
|
||||
)
|
||||
image = self.latent2image(latents)
|
||||
image = Image.fromarray(image)
|
||||
if save_intermediates:
|
||||
image.save(f"{self.output_path}/{i:02d}.png")
|
||||
images.append(image)
|
||||
|
||||
images = [first_image] + images + [last_image]
|
||||
|
||||
else:
|
||||
for k, alpha in enumerate(alpha_list):
|
||||
|
||||
latents = self.cal_latent(
|
||||
num_inference_steps,
|
||||
guidance_scale,
|
||||
unconditioning,
|
||||
img_noise_0,
|
||||
img_noise_1,
|
||||
text_embeddings_0,
|
||||
text_embeddings_1,
|
||||
lora_0,
|
||||
lora_1,
|
||||
alpha_list[k],
|
||||
self.use_lora,
|
||||
fix_lora
|
||||
)
|
||||
image = self.latent2image(latents)
|
||||
image = Image.fromarray(image)
|
||||
if save_intermediates:
|
||||
image.save(f"{self.output_path}/{k:02d}.png")
|
||||
images.append(image)
|
||||
|
||||
return images
|
||||
|
||||
with torch.no_grad():
|
||||
if self.use_reschedule:
|
||||
alpha_scheduler = AlphaScheduler()
|
||||
alpha_list = list(torch.linspace(0, 1, num_frames))
|
||||
images_pt = morph(alpha_list, progress, "Sampling...")
|
||||
images_pt = [transforms.ToTensor()(img).unsqueeze(0)
|
||||
for img in images_pt]
|
||||
alpha_scheduler.from_imgs(images_pt)
|
||||
alpha_list = alpha_scheduler.get_list()
|
||||
print(alpha_list)
|
||||
images = morph(alpha_list, progress, "Reschedule..."
|
||||
)
|
||||
else:
|
||||
alpha_list = list(torch.linspace(0, 1, num_frames))
|
||||
print(alpha_list)
|
||||
images = morph(alpha_list, progress, "Sampling...")
|
||||
|
||||
return images
|
||||
@@ -0,0 +1,54 @@
|
||||
import bisect
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import lpips
|
||||
|
||||
perceptual_loss = lpips.LPIPS()
|
||||
|
||||
|
||||
def distance(img_a, img_b):
|
||||
return perceptual_loss(img_a, img_b).item()
|
||||
# return F.mse_loss(img_a, img_b).item()
|
||||
|
||||
|
||||
class AlphaScheduler:
|
||||
def __init__(self):
|
||||
...
|
||||
|
||||
def from_imgs(self, imgs):
|
||||
self.__num_values = len(imgs)
|
||||
self.__values = [0]
|
||||
for i in range(self.__num_values - 1):
|
||||
dis = distance(imgs[i], imgs[i + 1])
|
||||
self.__values.append(dis)
|
||||
self.__values[i + 1] += self.__values[i]
|
||||
for i in range(self.__num_values):
|
||||
self.__values[i] /= self.__values[-1]
|
||||
|
||||
def save(self, filename):
|
||||
torch.save(torch.tensor(self.__values), filename)
|
||||
|
||||
def load(self, filename):
|
||||
self.__values = torch.load(filename).tolist()
|
||||
self.__num_values = len(self.__values)
|
||||
|
||||
def get_x(self, y):
|
||||
assert y >= 0 and y <= 1
|
||||
id = bisect.bisect_left(self.__values, y)
|
||||
id -= 1
|
||||
if id < 0:
|
||||
id = 0
|
||||
yl = self.__values[id]
|
||||
yr = self.__values[id + 1]
|
||||
xl = id * (1 / (self.__num_values - 1))
|
||||
xr = (id + 1) * (1 / (self.__num_values - 1))
|
||||
x = (y - yl) / (yr - yl) * (xr - xl) + xl
|
||||
return x
|
||||
|
||||
def get_list(self, len=None):
|
||||
if len is None:
|
||||
len = self.__num_values
|
||||
|
||||
ys = torch.linspace(0, 1, len)
|
||||
res = [self.get_x(y) for y in ys]
|
||||
return res
|
||||
@@ -0,0 +1,283 @@
|
||||
from timeit import default_timer as timer
|
||||
from datetime import timedelta
|
||||
from PIL import Image
|
||||
import os
|
||||
import numpy as np
|
||||
from einops import rearrange
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torchvision import transforms
|
||||
import transformers
|
||||
from accelerate import Accelerator
|
||||
from accelerate.utils import set_seed
|
||||
from packaging import version
|
||||
from PIL import Image
|
||||
import tqdm
|
||||
|
||||
from transformers import AutoTokenizer, PretrainedConfig
|
||||
|
||||
import diffusers
|
||||
from diffusers import (
|
||||
AutoencoderKL,
|
||||
DDPMScheduler,
|
||||
DiffusionPipeline,
|
||||
DPMSolverMultistepScheduler,
|
||||
StableDiffusionPipeline,
|
||||
UNet2DConditionModel,
|
||||
)
|
||||
from diffusers.loaders import AttnProcsLayers, LoraLoaderMixin
|
||||
from diffusers.models.attention_processor import (
|
||||
AttnAddedKVProcessor,
|
||||
AttnAddedKVProcessor2_0,
|
||||
LoRAAttnAddedKVProcessor,
|
||||
LoRAAttnProcessor,
|
||||
LoRAAttnProcessor2_0,
|
||||
SlicedAttnAddedKVProcessor,
|
||||
)
|
||||
from diffusers.optimization import get_scheduler
|
||||
from diffusers.utils import check_min_version
|
||||
from diffusers.utils.import_utils import is_xformers_available
|
||||
|
||||
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
|
||||
check_min_version("0.17.0")
|
||||
|
||||
|
||||
def import_model_class_from_model_name_or_path(pretrained_model_name_or_path: str, revision: str):
|
||||
text_encoder_config = PretrainedConfig.from_pretrained(
|
||||
pretrained_model_name_or_path,
|
||||
subfolder="text_encoder",
|
||||
revision=revision,
|
||||
)
|
||||
model_class = text_encoder_config.architectures[0]
|
||||
|
||||
if model_class == "CLIPTextModel":
|
||||
from transformers import CLIPTextModel
|
||||
|
||||
return CLIPTextModel
|
||||
elif model_class == "RobertaSeriesModelWithTransformation":
|
||||
from diffusers.pipelines.alt_diffusion.modeling_roberta_series import RobertaSeriesModelWithTransformation
|
||||
|
||||
return RobertaSeriesModelWithTransformation
|
||||
elif model_class == "T5EncoderModel":
|
||||
from transformers import T5EncoderModel
|
||||
|
||||
return T5EncoderModel
|
||||
else:
|
||||
raise ValueError(f"{model_class} is not supported.")
|
||||
|
||||
def tokenize_prompt(tokenizer, prompt, tokenizer_max_length=None):
|
||||
if tokenizer_max_length is not None:
|
||||
max_length = tokenizer_max_length
|
||||
else:
|
||||
max_length = tokenizer.model_max_length
|
||||
|
||||
text_inputs = tokenizer(
|
||||
prompt,
|
||||
truncation=True,
|
||||
padding="max_length",
|
||||
max_length=max_length,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
return text_inputs
|
||||
|
||||
def encode_prompt(text_encoder, input_ids, attention_mask, text_encoder_use_attention_mask=False):
|
||||
text_input_ids = input_ids.to(text_encoder.device)
|
||||
|
||||
if text_encoder_use_attention_mask:
|
||||
attention_mask = attention_mask.to(text_encoder.device)
|
||||
else:
|
||||
attention_mask = None
|
||||
|
||||
prompt_embeds = text_encoder(
|
||||
text_input_ids,
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
prompt_embeds = prompt_embeds[0]
|
||||
|
||||
return prompt_embeds
|
||||
|
||||
# model_path: path of the model
|
||||
# image: input image, have not been pre-processed
|
||||
# save_lora_dir: the path to save the lora
|
||||
# prompt: the user input prompt
|
||||
# lora_steps: number of lora training step
|
||||
# lora_lr: learning rate of lora training
|
||||
# lora_rank: the rank of lora
|
||||
def train_lora(image, prompt, save_lora_dir, model_path=None, tokenizer=None, text_encoder=None, vae=None, unet=None, noise_scheduler=None, lora_steps=200, lora_lr=2e-4, lora_rank=16, weight_name=None, safe_serialization=False, progress=tqdm):
|
||||
# initialize accelerator
|
||||
accelerator = Accelerator(
|
||||
gradient_accumulation_steps=1,
|
||||
# mixed_precision='fp16'
|
||||
)
|
||||
set_seed(0)
|
||||
|
||||
# Load the tokenizer
|
||||
if tokenizer is None:
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
model_path,
|
||||
subfolder="tokenizer",
|
||||
revision=None,
|
||||
use_fast=False,
|
||||
)
|
||||
# initialize the model
|
||||
if noise_scheduler is None:
|
||||
noise_scheduler = DDPMScheduler.from_pretrained(model_path, subfolder="scheduler")
|
||||
if text_encoder is None:
|
||||
text_encoder_cls = import_model_class_from_model_name_or_path(model_path, revision=None)
|
||||
text_encoder = text_encoder_cls.from_pretrained(
|
||||
model_path, subfolder="text_encoder", revision=None
|
||||
)
|
||||
if vae is None:
|
||||
vae = AutoencoderKL.from_pretrained(
|
||||
model_path, subfolder="vae", revision=None
|
||||
)
|
||||
if unet is None:
|
||||
unet = UNet2DConditionModel.from_pretrained(
|
||||
model_path, subfolder="unet", revision=None
|
||||
)
|
||||
|
||||
# set device and dtype
|
||||
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
|
||||
|
||||
vae.requires_grad_(False)
|
||||
text_encoder.requires_grad_(False)
|
||||
unet.requires_grad_(False)
|
||||
|
||||
unet.to(device)
|
||||
vae.to(device)
|
||||
text_encoder.to(device)
|
||||
|
||||
# initialize UNet LoRA
|
||||
unet_lora_attn_procs = {}
|
||||
for name, attn_processor in unet.attn_processors.items():
|
||||
cross_attention_dim = None if name.endswith("attn1.processor") else unet.config.cross_attention_dim
|
||||
if name.startswith("mid_block"):
|
||||
hidden_size = unet.config.block_out_channels[-1]
|
||||
elif name.startswith("up_blocks"):
|
||||
block_id = int(name[len("up_blocks.")])
|
||||
hidden_size = list(reversed(unet.config.block_out_channels))[block_id]
|
||||
elif name.startswith("down_blocks"):
|
||||
block_id = int(name[len("down_blocks.")])
|
||||
hidden_size = unet.config.block_out_channels[block_id]
|
||||
else:
|
||||
raise NotImplementedError("name must start with up_blocks, mid_blocks, or down_blocks")
|
||||
|
||||
if isinstance(attn_processor, (AttnAddedKVProcessor, SlicedAttnAddedKVProcessor, AttnAddedKVProcessor2_0)):
|
||||
lora_attn_processor_class = LoRAAttnAddedKVProcessor
|
||||
else:
|
||||
lora_attn_processor_class = (
|
||||
LoRAAttnProcessor2_0 if hasattr(F, "scaled_dot_product_attention") else LoRAAttnProcessor
|
||||
)
|
||||
unet_lora_attn_procs[name] = lora_attn_processor_class(
|
||||
hidden_size=hidden_size, cross_attention_dim=cross_attention_dim, rank=lora_rank
|
||||
)
|
||||
unet.set_attn_processor(unet_lora_attn_procs)
|
||||
unet_lora_layers = AttnProcsLayers(unet.attn_processors)
|
||||
|
||||
# Optimizer creation
|
||||
params_to_optimize = (unet_lora_layers.parameters())
|
||||
optimizer = torch.optim.AdamW(
|
||||
params_to_optimize,
|
||||
lr=lora_lr,
|
||||
betas=(0.9, 0.999),
|
||||
weight_decay=1e-2,
|
||||
eps=1e-08,
|
||||
)
|
||||
|
||||
lr_scheduler = get_scheduler(
|
||||
"constant",
|
||||
optimizer=optimizer,
|
||||
num_warmup_steps=0,
|
||||
num_training_steps=lora_steps,
|
||||
num_cycles=1,
|
||||
power=1.0,
|
||||
)
|
||||
|
||||
# prepare accelerator
|
||||
unet_lora_layers = accelerator.prepare_model(unet_lora_layers)
|
||||
optimizer = accelerator.prepare_optimizer(optimizer)
|
||||
lr_scheduler = accelerator.prepare_scheduler(lr_scheduler)
|
||||
|
||||
# initialize text embeddings
|
||||
with torch.no_grad():
|
||||
text_inputs = tokenize_prompt(tokenizer, prompt, tokenizer_max_length=None)
|
||||
text_embedding = encode_prompt(
|
||||
text_encoder,
|
||||
text_inputs.input_ids,
|
||||
text_inputs.attention_mask,
|
||||
text_encoder_use_attention_mask=False
|
||||
)
|
||||
|
||||
if type(image) == np.ndarray:
|
||||
image = Image.fromarray(image)
|
||||
|
||||
# initialize latent distribution
|
||||
image_transforms = transforms.Compose(
|
||||
[
|
||||
transforms.Resize(512, interpolation=transforms.InterpolationMode.BILINEAR),
|
||||
# transforms.RandomCrop(512),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.5], [0.5]),
|
||||
]
|
||||
)
|
||||
|
||||
image = image_transforms(image).to(device)
|
||||
image = image.unsqueeze(dim=0)
|
||||
|
||||
latents_dist = vae.encode(image).latent_dist
|
||||
for _ in progress.tqdm(range(lora_steps), desc="Training LoRA..."):
|
||||
unet.train()
|
||||
model_input = latents_dist.sample() * vae.config.scaling_factor
|
||||
# Sample noise that we'll add to the latents
|
||||
noise = torch.randn_like(model_input)
|
||||
bsz, channels, height, width = model_input.shape
|
||||
# Sample a random timestep for each image
|
||||
timesteps = torch.randint(
|
||||
0, noise_scheduler.config.num_train_timesteps, (bsz,), device=model_input.device
|
||||
)
|
||||
timesteps = timesteps.long()
|
||||
|
||||
# Add noise to the model input according to the noise magnitude at each timestep
|
||||
# (this is the forward diffusion process)
|
||||
noisy_model_input = noise_scheduler.add_noise(model_input, noise, timesteps)
|
||||
|
||||
# Predict the noise residual
|
||||
model_pred = unet(noisy_model_input, timesteps, text_embedding).sample
|
||||
|
||||
# Get the target for loss depending on the prediction type
|
||||
if noise_scheduler.config.prediction_type == "epsilon":
|
||||
target = noise
|
||||
elif noise_scheduler.config.prediction_type == "v_prediction":
|
||||
target = noise_scheduler.get_velocity(model_input, noise, timesteps)
|
||||
else:
|
||||
raise ValueError(f"Unknown prediction type {noise_scheduler.config.prediction_type}")
|
||||
|
||||
loss = F.mse_loss(model_pred.float(), target.float(), reduction="mean")
|
||||
accelerator.backward(loss)
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
optimizer.zero_grad()
|
||||
|
||||
# save the trained lora
|
||||
# unet = unet.to(torch.float32)
|
||||
# vae = vae.to(torch.float32)
|
||||
# text_encoder = text_encoder.to(torch.float32)
|
||||
|
||||
# unwrap_model is used to remove all special modules added when doing distributed training
|
||||
# so here, there is no need to call unwrap_model
|
||||
# unet_lora_layers = accelerator.unwrap_model(unet_lora_layers)
|
||||
LoraLoaderMixin.save_lora_weights(
|
||||
save_directory=save_lora_dir,
|
||||
unet_lora_layers=unet_lora_layers,
|
||||
text_encoder_lora_layers=None,
|
||||
weight_name=weight_name,
|
||||
safe_serialization=safe_serialization
|
||||
)
|
||||
|
||||
def load_lora(unet, lora_0, lora_1, alpha):
|
||||
lora = {}
|
||||
for key in lora_0:
|
||||
lora[key] = (1 - alpha) * lora_0[key] + alpha * lora_1[key]
|
||||
unet.load_attn_procs(lora)
|
||||
return unet
|
||||
@@ -0,0 +1,86 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torchvision import transforms
|
||||
|
||||
def calc_mean_std(feat, eps=1e-5):
|
||||
# eps is a small value added to the variance to avoid divide-by-zero.
|
||||
size = feat.size()
|
||||
|
||||
N, C = size[:2]
|
||||
feat_var = feat.view(N, C, -1).var(dim=2) + eps
|
||||
if len(size) == 3:
|
||||
feat_std = feat_var.sqrt().view(N, C, 1)
|
||||
feat_mean = feat.view(N, C, -1).mean(dim=2).view(N, C, 1)
|
||||
else:
|
||||
feat_std = feat_var.sqrt().view(N, C, 1, 1)
|
||||
feat_mean = feat.view(N, C, -1).mean(dim=2).view(N, C, 1, 1)
|
||||
return feat_mean, feat_std
|
||||
|
||||
|
||||
def get_img(img, resolution=512):
|
||||
norm_mean = [0.5, 0.5, 0.5]
|
||||
norm_std = [0.5, 0.5, 0.5]
|
||||
transform = transforms.Compose([
|
||||
transforms.Resize((resolution, resolution)),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize(norm_mean, norm_std)
|
||||
])
|
||||
img = transform(img)
|
||||
return img.unsqueeze(0)
|
||||
|
||||
@torch.no_grad()
|
||||
def slerp(p0, p1, fract_mixing: float, adain=True):
|
||||
r""" Copied from lunarring/latentblending
|
||||
Helper function to correctly mix two random variables using spherical interpolation.
|
||||
The function will always cast up to float64 for sake of extra 4.
|
||||
Args:
|
||||
p0:
|
||||
First tensor for interpolation
|
||||
p1:
|
||||
Second tensor for interpolation
|
||||
fract_mixing: float
|
||||
Mixing coefficient of interval [0, 1].
|
||||
0 will return in p0
|
||||
1 will return in p1
|
||||
0.x will return a mix between both preserving angular velocity.
|
||||
"""
|
||||
if p0.dtype == torch.float16:
|
||||
recast_to = 'fp16'
|
||||
else:
|
||||
recast_to = 'fp32'
|
||||
|
||||
p0 = p0.double()
|
||||
p1 = p1.double()
|
||||
|
||||
if adain:
|
||||
mean1, std1 = calc_mean_std(p0)
|
||||
mean2, std2 = calc_mean_std(p1)
|
||||
mean = mean1 * (1 - fract_mixing) + mean2 * fract_mixing
|
||||
std = std1 * (1 - fract_mixing) + std2 * fract_mixing
|
||||
|
||||
norm = torch.linalg.norm(p0) * torch.linalg.norm(p1)
|
||||
epsilon = 1e-7
|
||||
dot = torch.sum(p0 * p1) / norm
|
||||
dot = dot.clamp(-1+epsilon, 1-epsilon)
|
||||
|
||||
theta_0 = torch.arccos(dot)
|
||||
sin_theta_0 = torch.sin(theta_0)
|
||||
theta_t = theta_0 * fract_mixing
|
||||
s0 = torch.sin(theta_0 - theta_t) / sin_theta_0
|
||||
s1 = torch.sin(theta_t) / sin_theta_0
|
||||
interp = p0*s0 + p1*s1
|
||||
|
||||
if adain:
|
||||
interp = F.instance_norm(interp) * std + mean
|
||||
|
||||
if recast_to == 'fp16':
|
||||
interp = interp.half()
|
||||
elif recast_to == 'fp32':
|
||||
interp = interp.float()
|
||||
|
||||
return interp
|
||||
|
||||
|
||||
def do_replace_attn(key: str):
|
||||
# return key.startswith('up_blocks.2') or key.startswith('up_blocks.3')
|
||||
return key.startswith('up')
|
||||
BIN
Binary file not shown.
|
After Width: | Height: | Size: 293 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 67 KiB |
@@ -0,0 +1,301 @@
|
||||
{
|
||||
"last_node_id": 11,
|
||||
"last_link_id": 20,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 3,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
66,
|
||||
28
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 314
|
||||
},
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
16
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
"Biden.jpg",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 4,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
81,
|
||||
401
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 314.0000305175781
|
||||
},
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
17
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
"Trump.jpg",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 5,
|
||||
"type": "TextNode",
|
||||
"pos": [
|
||||
412,
|
||||
35
|
||||
],
|
||||
"size": {
|
||||
"0": 400,
|
||||
"1": 200
|
||||
},
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "TEXT",
|
||||
"type": "TEXT",
|
||||
"links": [
|
||||
18
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "TextNode"
|
||||
},
|
||||
"widgets_values": [
|
||||
"A photo of an American man"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 8,
|
||||
"type": "TextNode",
|
||||
"pos": [
|
||||
415,
|
||||
529
|
||||
],
|
||||
"size": {
|
||||
"0": 400,
|
||||
"1": 200
|
||||
},
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "TEXT",
|
||||
"type": "TEXT",
|
||||
"links": [
|
||||
19
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "TextNode"
|
||||
},
|
||||
"widgets_values": [
|
||||
"A photo of an American man"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 11,
|
||||
"type": "DiffMorpherNode",
|
||||
"pos": [
|
||||
861,
|
||||
383
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 310
|
||||
},
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image_0",
|
||||
"type": "IMAGE",
|
||||
"link": 16
|
||||
},
|
||||
{
|
||||
"name": "prompt_0",
|
||||
"type": "TEXT",
|
||||
"link": 18
|
||||
},
|
||||
{
|
||||
"name": "image_1",
|
||||
"type": "IMAGE",
|
||||
"link": 17
|
||||
},
|
||||
{
|
||||
"name": "prompt_1",
|
||||
"type": "TEXT",
|
||||
"link": 19
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "GIF",
|
||||
"type": "GIF",
|
||||
"links": [
|
||||
20
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "DiffMorpherNode"
|
||||
},
|
||||
"widgets_values": [
|
||||
16,
|
||||
100,
|
||||
false,
|
||||
false,
|
||||
0.6,
|
||||
false,
|
||||
null,
|
||||
null,
|
||||
null
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 10,
|
||||
"type": "PreViewGIF",
|
||||
"pos": [
|
||||
998,
|
||||
61
|
||||
],
|
||||
"size": {
|
||||
"0": 210,
|
||||
"1": 26
|
||||
},
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "gif",
|
||||
"type": "GIF",
|
||||
"link": 20
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "PreViewGIF"
|
||||
},
|
||||
"widgets_values": [
|
||||
{
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {}
|
||||
},
|
||||
{
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {}
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
16,
|
||||
3,
|
||||
0,
|
||||
11,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
17,
|
||||
4,
|
||||
0,
|
||||
11,
|
||||
2,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
18,
|
||||
5,
|
||||
0,
|
||||
11,
|
||||
1,
|
||||
"TEXT"
|
||||
],
|
||||
[
|
||||
19,
|
||||
8,
|
||||
0,
|
||||
11,
|
||||
3,
|
||||
"TEXT"
|
||||
],
|
||||
[
|
||||
20,
|
||||
11,
|
||||
0,
|
||||
10,
|
||||
0,
|
||||
"GIF"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 1,
|
||||
"offset": [
|
||||
0,
|
||||
0
|
||||
]
|
||||
}
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 2.2 MiB |
@@ -0,0 +1,10 @@
|
||||
accelerate
|
||||
diffusers
|
||||
einops
|
||||
numpy
|
||||
opencv_python
|
||||
Pillow
|
||||
safetensors
|
||||
tqdm
|
||||
transformers
|
||||
lpips
|
||||
@@ -0,0 +1,154 @@
|
||||
import { app } from "../../../scripts/app.js";
|
||||
import { api } from '../../../scripts/api.js'
|
||||
|
||||
function fitHeight(node) {
|
||||
node.setSize([node.size[0], node.computeSize([node.size[0], node.size[1]])[1]])
|
||||
node?.graph?.setDirtyCanvas(true);
|
||||
}
|
||||
function chainCallback(object, property, callback) {
|
||||
if (object == undefined) {
|
||||
//This should not happen.
|
||||
console.error("Tried to add callback to non-existant object")
|
||||
return;
|
||||
}
|
||||
if (property in object) {
|
||||
const callback_orig = object[property]
|
||||
object[property] = function () {
|
||||
const r = callback_orig.apply(this, arguments);
|
||||
callback.apply(this, arguments);
|
||||
return r
|
||||
};
|
||||
} else {
|
||||
object[property] = callback;
|
||||
}
|
||||
}
|
||||
|
||||
function addPreviewOptions(nodeType) {
|
||||
chainCallback(nodeType.prototype, "getExtraMenuOptions", function(_, options) {
|
||||
// The intended way of appending options is returning a list of extra options,
|
||||
// but this isn't used in widgetInputs.js and would require
|
||||
// less generalization of chainCallback
|
||||
let optNew = []
|
||||
try {
|
||||
const previewWidget = this.widgets.find((w) => w.name === "videopreview");
|
||||
|
||||
let url = null
|
||||
if (previewWidget.videoEl?.hidden == false && previewWidget.videoEl.src) {
|
||||
//Use full quality video
|
||||
//url = api.apiURL('/view?' + new URLSearchParams(previewWidget.value.params));
|
||||
url = previewWidget.videoEl.src
|
||||
}
|
||||
if (url) {
|
||||
optNew.push(
|
||||
{
|
||||
content: "Open preview",
|
||||
callback: () => {
|
||||
window.open(url, "_blank")
|
||||
},
|
||||
},
|
||||
{
|
||||
content: "Save preview",
|
||||
callback: () => {
|
||||
const a = document.createElement("a");
|
||||
a.href = url;
|
||||
a.setAttribute("download", new URLSearchParams(previewWidget.value.params).get("filename"));
|
||||
document.body.append(a);
|
||||
a.click();
|
||||
requestAnimationFrame(() => a.remove());
|
||||
},
|
||||
}
|
||||
);
|
||||
}
|
||||
if(options.length > 0 && options[0] != null && optNew.length > 0) {
|
||||
optNew.push(null);
|
||||
}
|
||||
options.unshift(...optNew);
|
||||
|
||||
} catch (error) {
|
||||
console.log(error);
|
||||
}
|
||||
|
||||
});
|
||||
}
|
||||
function previewVideo(node,file,type){
|
||||
var element = document.createElement("div");
|
||||
const previewNode = node;
|
||||
var previewWidget = node.addDOMWidget("videopreview", "preview", element, {
|
||||
serialize: false,
|
||||
hideOnZoom: false,
|
||||
getValue() {
|
||||
return element.value;
|
||||
},
|
||||
setValue(v) {
|
||||
element.value = v;
|
||||
},
|
||||
});
|
||||
previewWidget.computeSize = function(width) {
|
||||
if (this.aspectRatio && !this.parentEl.hidden) {
|
||||
let height = (previewNode.size[0]-20)/ this.aspectRatio + 10;
|
||||
if (!(height > 0)) {
|
||||
height = 0;
|
||||
}
|
||||
this.computedHeight = height + 10;
|
||||
return [width, height];
|
||||
}
|
||||
return [width, -4];//no loaded src, widget should not display
|
||||
}
|
||||
// element.style['pointer-events'] = "none"
|
||||
previewWidget.value = {hidden: false, paused: false, params: {}}
|
||||
previewWidget.parentEl = document.createElement("div");
|
||||
previewWidget.parentEl.className = "video_preview";
|
||||
previewWidget.parentEl.style['width'] = "100%"
|
||||
element.appendChild(previewWidget.parentEl);
|
||||
previewWidget.videoEl = document.createElement("img");
|
||||
previewWidget.videoEl.controls = true;
|
||||
previewWidget.videoEl.loop = false;
|
||||
previewWidget.videoEl.muted = false;
|
||||
previewWidget.videoEl.style['width'] = "100%"
|
||||
previewWidget.videoEl.addEventListener("loadedmetadata", () => {
|
||||
|
||||
previewWidget.aspectRatio = previewWidget.videoEl.videoWidth / previewWidget.videoEl.videoHeight;
|
||||
fitHeight(this);
|
||||
});
|
||||
previewWidget.videoEl.addEventListener("error", () => {
|
||||
//TODO: consider a way to properly notify the user why a preview isn't shown.
|
||||
previewWidget.parentEl.hidden = true;
|
||||
fitHeight(this);
|
||||
});
|
||||
|
||||
let params = {
|
||||
"filename": file,
|
||||
"type": type,
|
||||
}
|
||||
|
||||
previewWidget.parentEl.hidden = previewWidget.value.hidden;
|
||||
previewWidget.videoEl.autoplay = !previewWidget.value.paused && !previewWidget.value.hidden;
|
||||
let target_width = 256
|
||||
if (element.style?.width) {
|
||||
//overscale to allow scrolling. Endpoint won't return higher than native
|
||||
target_width = element.style.width.slice(0,-2)*2;
|
||||
}
|
||||
if (!params.force_size || params.force_size.includes("?") || params.force_size == "Disabled") {
|
||||
params.force_size = target_width+"x?"
|
||||
} else {
|
||||
let size = params.force_size.split("x")
|
||||
let ar = parseInt(size[0])/parseInt(size[1])
|
||||
params.force_size = target_width+"x"+(target_width/ar)
|
||||
}
|
||||
|
||||
previewWidget.videoEl.src = api.apiURL('/view?' + new URLSearchParams(params));
|
||||
|
||||
previewWidget.videoEl.hidden = false;
|
||||
previewWidget.parentEl.appendChild(previewWidget.videoEl)
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: "DiffMorpher.VideoPreviewer",
|
||||
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||
if (nodeData?.name == "PreViewGIF") {
|
||||
nodeType.prototype.onExecuted = function (data) {
|
||||
previewVideo(this, data.gif[0], data.gif[1]);
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
Reference in New Issue
Block a user