first commit

This commit is contained in:
AIFSH
2024-07-16 17:50:08 +08:00
parent 78ecdbf9be
commit 18fb8868bb
14 changed files with 1758 additions and 1 deletions
+1
View File
@@ -0,0 +1 @@
__pycache__
+45 -1
View File
@@ -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
----- | ---- | ----
![](./doc/Biden.jpg) | ![](./doc/15.png) | ![](./doc/diffmorpher_1721118982208625000.gif)
## 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
View File
@@ -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
}
+643
View File
@@ -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
View File
+54
View File
@@ -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
+283
View File
@@ -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
+86
View File
@@ -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
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 293 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 67 KiB

+301
View File
@@ -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

+10
View File
@@ -0,0 +1,10 @@
accelerate
diffusers
einops
numpy
opencv_python
Pillow
safetensors
tqdm
transformers
lpips
+154
View File
@@ -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]);
}
}
}
});