Compatible with the latest version of ComfyUI

This commit is contained in:
sylym
2023-04-05 12:48:38 +08:00
parent 14869385cc
commit 45c25ec9fb
5 changed files with 81 additions and 100 deletions
+5 -5
View File
@@ -1,4 +1,4 @@
# Vid2vid Node Suite for ComfyUI # Vid2vid Node Suite for [ComfyUI](https://github.com/comfyanonymous/ComfyUI)
A node suite for ComfyUI that allows you to load image sequence and generate new image sequence with different styles or content. A node suite for ComfyUI that allows you to load image sequence and generate new image sequence with different styles or content.
@@ -54,7 +54,7 @@ Load image sequence from a folder.
- n_sample_frames - n_sample_frames
- The number of images in the sequence. The number of images in `image_sequence_folder` must be greater than or equal to `sample_start_idx - 1 + n_sample_frames * sample_frame_rate`. - The number of images in the sequence. The number of images in `image_sequence_folder` must be greater than or equal to `sample_start_idx - 1 + n_sample_frames * sample_frame_rate`.
- If you want to use the node `CheckpointLoaderSimpleSequence` to generate a sequence of pictures, set `n_sample_frames` >= 3 can improve the time consistency of the output image sequence. - If you want to use the node `CheckpointLoaderSimpleSequence` to generate a sequence of pictures, set `n_sample_frames` >= 3.
### LoadImageMaskSequence ### LoadImageMaskSequence
@@ -192,9 +192,9 @@ Same function as KSampler node, but added support for noise vector and image mas
## Limits ## Limits
- UNet3DCoditionModel has high demand for GPU memory. If you encounter out of memory error, try to reduce `n_sample_frames`. If your computer has less than 6GB of GPU memory, `n_sample_frames` should not exceed 3, but this may reduce the quality of the generated results - UNet3DCoditionModel has high demand for GPU memory. If you encounter out of memory error, try to reduce `n_sample_frames`. However, `n_sample_frames` must be greater than or equal to 3.
- Some custom nodes do not support processing image sequences. The nodes listed below have been tested and are working properly: - Some custom nodes do not support processing image sequences. The nodes listed below have been tested and are working properly:
- Official node - [Official node](https://github.com/comfyanonymous/ComfyUI)
- comfy_controlnet_preprocessors - [comfy_controlnet_preprocessors](https://github.com/Fannovel16/comfy_controlnet_preprocessors)
+1 -1
View File
@@ -254,7 +254,7 @@ class CheckpointLoaderSimpleSequence:
def INPUT_TYPES(s): def INPUT_TYPES(s):
return {"required": { "ckpt_name": (folder_paths.get_filename_list("checkpoints"), ), return {"required": { "ckpt_name": (folder_paths.get_filename_list("checkpoints"), ),
}} }}
RETURN_TYPES = ("ORIGINAL_MODEL", "CLIP", "VAE") RETURN_TYPES = ("MODEL", "CLIP", "VAE")
FUNCTION = "load_checkpoint" FUNCTION = "load_checkpoint"
CATEGORY = "vid2vid" CATEGORY = "vid2vid"
-1
View File
@@ -1,4 +1,3 @@
from diffusers.models import UNet2DConditionModel
from .tuneavideo.models.unet import UNet3DConditionModel from .tuneavideo.models.unet import UNet3DConditionModel
from diffusers import DDIMScheduler from diffusers import DDIMScheduler
+46 -9
View File
@@ -1,15 +1,18 @@
import torch import torch
from comfy import model_management from comfy import model_management
from comfy.sd import load_torch_file, load_model_weights, ModelPatcher, VAE, CLIP from comfy.sd import load_model_weights, ModelPatcher, VAE, CLIP
from comfy import utils
from comfy import clip_vision
from comfy.ldm.util import instantiate_from_config from comfy.ldm.util import instantiate_from_config
from .convert_from_ckpt import convert_unet_checkpoint from .convert_from_ckpt import convert_unet_checkpoint
from omegaconf import OmegaConf from omegaconf import OmegaConf
def load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, embedding_directory=None): def load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, output_clipvision=False, embedding_directory=None):
sd = load_torch_file(ckpt_path) sd = utils.load_torch_file(ckpt_path)
sd_keys = sd.keys() sd_keys = sd.keys()
clip = None clip = None
clipvision = None
vae = None vae = None
fp16 = model_management.should_use_fp16() fp16 = model_management.should_use_fp16()
@@ -34,6 +37,29 @@ def load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, e
w.cond_stage_model = clip.cond_stage_model w.cond_stage_model = clip.cond_stage_model
load_state_dict_to = [w] load_state_dict_to = [w]
clipvision_key = "embedder.model.visual.transformer.resblocks.0.attn.in_proj_weight"
noise_aug_config = None
if clipvision_key in sd_keys:
size = sd[clipvision_key].shape[1]
if output_clipvision:
clipvision = clip_vision.load_clipvision_from_sd(sd)
noise_aug_key = "noise_augmentor.betas"
if noise_aug_key in sd_keys:
noise_aug_config = {}
params = {}
noise_schedule_config = {}
noise_schedule_config["timesteps"] = sd[noise_aug_key].shape[0]
noise_schedule_config["beta_schedule"] = "squaredcos_cap_v2"
params["noise_schedule_config"] = noise_schedule_config
noise_aug_config['target'] = "ldm.modules.encoders.noise_aug_modules.CLIPEmbeddingNoiseAugmentation"
if size == 1280: #h
params["timestep_dim"] = 1024
elif size == 1024: #l
params["timestep_dim"] = 768
noise_aug_config['params'] = params
sd_config = { sd_config = {
"linear_start": 0.00085, "linear_start": 0.00085,
"linear_end": 0.012, "linear_end": 0.012,
@@ -82,7 +108,13 @@ def load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, e
sd_config["unet_config"] = {"target": "ldm.modules.diffusionmodules.openaimodel.UNetModel", "params": unet_config} sd_config["unet_config"] = {"target": "ldm.modules.diffusionmodules.openaimodel.UNetModel", "params": unet_config}
model_config = {"target": "ldm.models.diffusion.ddpm.LatentDiffusion", "params": sd_config} model_config = {"target": "ldm.models.diffusion.ddpm.LatentDiffusion", "params": sd_config}
if unet_config["in_channels"] > 4: #inpainting model if noise_aug_config is not None: #SD2.x unclip model
sd_config["noise_aug_config"] = noise_aug_config
sd_config["image_size"] = 96
sd_config["embedding_dropout"] = 0.25
sd_config["conditioning_key"] = 'crossattn-adm'
model_config["target"] = "ldm.models.diffusion.ddpm.ImageEmbeddingConditionedLatentDiffusion"
elif unet_config["in_channels"] > 4: #inpainting model
sd_config["conditioning_key"] = "hybrid" sd_config["conditioning_key"] = "hybrid"
sd_config["finetune_keys"] = None sd_config["finetune_keys"] = None
model_config["target"] = "ldm.models.diffusion.ddpm.LatentInpaintDiffusion" model_config["target"] = "ldm.models.diffusion.ddpm.LatentInpaintDiffusion"
@@ -94,6 +126,11 @@ def load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, e
else: else:
unet_config["num_heads"] = 8 #SD1.x unet_config["num_heads"] = 8 #SD1.x
unclip = 'model.diffusion_model.label_emb.0.0.weight'
if unclip in sd_keys:
unet_config["num_classes"] = "sequential"
unet_config["adm_in_channels"] = sd[unclip].shape[1]
if unet_config["context_dim"] == 1024 and unet_config["in_channels"] == 4: #only SD2.x non inpainting models are v prediction if unet_config["context_dim"] == 1024 and unet_config["in_channels"] == 4: #only SD2.x non inpainting models are v prediction
k = "model.diffusion_model.output_blocks.11.1.transformer_blocks.0.norm1.bias" k = "model.diffusion_model.output_blocks.11.1.transformer_blocks.0.norm1.bias"
out = sd[k] out = sd[k]
@@ -102,13 +139,13 @@ def load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, e
model = instantiate_from_config(model_config) model = instantiate_from_config(model_config)
model = load_model_weights(model, sd, verbose=False, load_state_dict_to=load_state_dict_to) model = load_model_weights(model, sd, verbose=False, load_state_dict_to=load_state_dict_to)
with torch.inference_mode(mode=False):
model.model.diffusion_model = convert_unet_checkpoint(sd, OmegaConf.create({"model": model_config}))
#with torch.inference_mode(mode=False):
model.model.diffusion_model = convert_unet_checkpoint(sd, OmegaConf.create({"model": model_config}))
if model_management.xformers_enabled(): if model_management.xformers_enabled():
model.model.diffusion_model.enable_xformers_memory_efficient_attention() model.model.diffusion_model.enable_xformers_memory_efficient_attention()
#if fp16: if fp16:
# model = model.half() model = model.half()
return (ModelPatcher(model), clip, vae) return (ModelPatcher(model), clip, vae, clipvision)
+29 -84
View File
@@ -3,8 +3,6 @@
from dataclasses import dataclass from dataclasses import dataclass
from typing import List, Optional, Tuple, Union from typing import List, Optional, Tuple, Union
import os
import json
from einops import rearrange from einops import rearrange
import torch import torch
@@ -281,34 +279,25 @@ class UNet3DConditionModel(ModelMixin, ConfigMixin):
x: torch.FloatTensor, x: torch.FloatTensor,
timesteps: Union[torch.Tensor, float, int], timesteps: Union[torch.Tensor, float, int],
context: torch.Tensor, context: torch.Tensor,
y=None,
control: Optional[torch.Tensor] = None, control: Optional[torch.Tensor] = None,
class_labels: Optional[torch.Tensor] = None, transformer_options={}
attention_mask: Optional[torch.Tensor] = None,
) -> Union[UNet3DConditionOutput, Tuple]: ) -> Union[UNet3DConditionOutput, Tuple]:
# prepare timestep and sample for comfyui # prepare timestep and sample for comfyui
if not torch.all(torch.eq(timesteps, timesteps[0])): if not torch.all(torch.eq(timesteps, timesteps[0])):
raise ValueError("All timesteps must be equal.") raise ValueError("All timesteps must be equal.")
timestep = timesteps[0] timesteps = timesteps[0]
if not torch.all(torch.eq(context, context[0])): if not torch.all(torch.eq(context, context[0])):
raise ValueError("All contexts must be equal.") raise ValueError("All contexts must be equal.")
context = context[0].unsqueeze(0) context = context[0].unsqueeze(0)
sample = rearrange(x.unsqueeze(0), "b f c h w -> b c f h w") sample = rearrange(x.unsqueeze(0), "b f c h w -> b c f h w")
sample = sample.type(self.dtype) sample = sample.to(dtype=self.dtype)
context = context.type(self.dtype) context = context.to(dtype=self.dtype)
down_block_additional_residuals = None del x
mid_block_additional_residual = None
if control is not None:
down_block_additional_residuals = []
for output in control["output"]:
down_block_additional_residuals.append(rearrange(output.unsqueeze(0), "a b c d e -> a c b d e"))
mid_block_additional_residual = rearrange(control["middle"][0].unsqueeze(0), "a b c d e -> a c b d e")
del x, timesteps, control
torch.cuda.empty_cache() torch.cuda.empty_cache()
# By default samples have to be AT least a multiple of the overall upsampling factor. # By default samples have to be AT least a multiple of the overall upsampling factor.
@@ -325,28 +314,10 @@ class UNet3DConditionModel(ModelMixin, ConfigMixin):
logger.info("Forward upsample size to force interpolation output size.") logger.info("Forward upsample size to force interpolation output size.")
forward_upsample_size = True forward_upsample_size = True
# prepare attention_mask
if attention_mask is not None:
attention_mask = (1 - attention_mask.to(sample.dtype)) * -10000.0
attention_mask = attention_mask.unsqueeze(1)
# center input if necessary # center input if necessary
if self.config.center_input_sample: if self.config.center_input_sample:
sample = 2 * sample - 1.0 sample = 2 * sample - 1.0
# time
timesteps = timestep
if not torch.is_tensor(timesteps):
# This would be a good case for the `match` statement (Python 3.10+)
is_mps = sample.device.type == "mps"
if isinstance(timestep, float):
dtype = torch.float32 if is_mps else torch.float64
else:
dtype = torch.int32 if is_mps else torch.int64
timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device)
elif len(timesteps.shape) == 0:
timesteps = timesteps[None].to(sample.device)
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML # broadcast to batch dimension in a way that's compatible with ONNX/Core ML
timesteps = timesteps.expand(sample.shape[0]) timesteps = timesteps.expand(sample.shape[0])
@@ -359,13 +330,13 @@ class UNet3DConditionModel(ModelMixin, ConfigMixin):
emb = self.time_embedding(t_emb) emb = self.time_embedding(t_emb)
if self.class_embedding is not None: if self.class_embedding is not None:
if class_labels is None: if y is None:
raise ValueError("class_labels should be provided when num_class_embeds > 0") raise ValueError("y should be provided when num_class_embeds > 0")
if self.config.class_embed_type == "timestep": if self.config.class_embed_type == "timestep":
class_labels = self.time_proj(class_labels) y = self.time_proj(y)
class_emb = self.class_embedding(class_labels).to(dtype=self.dtype) class_emb = self.class_embedding(y).to(dtype=self.dtype)
emb = emb + class_emb emb = emb + class_emb
# pre-process # pre-process
@@ -379,31 +350,43 @@ class UNet3DConditionModel(ModelMixin, ConfigMixin):
hidden_states=sample, hidden_states=sample,
temb=emb, temb=emb,
encoder_hidden_states=context, encoder_hidden_states=context,
attention_mask=attention_mask,
) )
else: else:
sample, res_samples = downsample_block(hidden_states=sample, temb=emb) sample, res_samples = downsample_block(hidden_states=sample, temb=emb)
down_block_res_samples += res_samples down_block_res_samples += res_samples
if down_block_additional_residuals is not None: if control is not None and "input" in control and len(control['input']) > 0:
new_down_block_res_samples = () new_down_block_res_samples = ()
for down_block_res_sample, down_block_additional_residual in zip( for down_block_res_sample, down_block_additional_residual in zip(
down_block_res_samples, down_block_additional_residuals down_block_res_samples, reversed(control['input'])
): ):
down_block_res_sample += down_block_additional_residual if down_block_additional_residual is not None:
down_block_res_sample += rearrange(down_block_additional_residual.unsqueeze(0), "a b c d e -> a c b d e")
new_down_block_res_samples += (down_block_res_sample,) new_down_block_res_samples += (down_block_res_sample,)
down_block_res_samples = new_down_block_res_samples down_block_res_samples = new_down_block_res_samples
# mid # mid
sample = self.mid_block( sample = self.mid_block(
sample, emb, encoder_hidden_states=context, attention_mask=attention_mask sample, emb, encoder_hidden_states=context
) )
if mid_block_additional_residual is not None: if control is not None and "middle" in control and len(control['middle']) > 0:
sample += mid_block_additional_residual sample += rearrange(control["middle"][0].unsqueeze(0), "a b c d e -> a c b d e")
if control is not None and "output" in control and len(control['output']) > 0:
new_down_block_res_samples = ()
for down_block_res_sample, down_block_additional_residual in zip(
down_block_res_samples, control["output"]
):
if down_block_additional_residual is not None:
down_block_res_sample += rearrange(down_block_additional_residual.unsqueeze(0), "a b c d e -> a c b d e")
new_down_block_res_samples += (down_block_res_sample,)
down_block_res_samples = new_down_block_res_samples
# up # up
for i, upsample_block in enumerate(self.up_blocks): for i, upsample_block in enumerate(self.up_blocks):
@@ -424,7 +407,6 @@ class UNet3DConditionModel(ModelMixin, ConfigMixin):
res_hidden_states_tuple=res_samples, res_hidden_states_tuple=res_samples,
encoder_hidden_states=context, encoder_hidden_states=context,
upsample_size=upsample_size, upsample_size=upsample_size,
attention_mask=attention_mask,
) )
else: else:
sample = upsample_block( sample = upsample_block(
@@ -439,40 +421,3 @@ class UNet3DConditionModel(ModelMixin, ConfigMixin):
sample = rearrange(sample.squeeze(0), "c f h w -> f c h w") sample = rearrange(sample.squeeze(0), "c f h w -> f c h w")
return sample return sample
@classmethod
def from_pretrained_2d(cls, pretrained_model_path, subfolder=None):
if subfolder is not None:
pretrained_model_path = os.path.join(pretrained_model_path, subfolder)
config_file = os.path.join(pretrained_model_path, 'config.json')
if not os.path.isfile(config_file):
raise RuntimeError(f"{config_file} does not exist")
with open(config_file, "r") as f:
config = json.load(f)
config["_class_name"] = cls.__name__
config["down_block_types"] = [
"CrossAttnDownBlock3D",
"CrossAttnDownBlock3D",
"CrossAttnDownBlock3D",
"DownBlock3D"
]
config["up_block_types"] = [
"UpBlock3D",
"CrossAttnUpBlock3D",
"CrossAttnUpBlock3D",
"CrossAttnUpBlock3D"
]
from diffusers.utils import WEIGHTS_NAME
model = cls.from_config(config)
model_file = os.path.join(pretrained_model_path, WEIGHTS_NAME)
if not os.path.isfile(model_file):
raise RuntimeError(f"{model_file} does not exist")
state_dict = torch.load(model_file, map_location="cpu")
for k, v in model.state_dict().items():
if '_temp.' in k:
state_dict.update({k: v})
model.load_state_dict(state_dict)
return model