first commit

This commit is contained in:
Tung Nguyen
2023-07-22 22:11:47 +07:00
commit 900de94c97
10 changed files with 2708 additions and 0 deletions
+4
View File
@@ -0,0 +1,4 @@
__pycache__
/venv
.vscode
*.ckpt
+21
View File
@@ -0,0 +1,21 @@
# AnimateDiff for ComfyUI
[AnimateDiff](https://github.com/guoyww/AnimateDiff/) integration for ComfyUI, adapts from [sd-webui-animatediff](https://github.com/continue-revolution/sd-webui-animatediff). Please read the original repo README for more information.
## How to Use
1. Clone this repo into `custom_nodes` folder.
2. Download motion modules from [Google Drive](https://drive.google.com/drive/folders/1EqLC65eR1-W-sGD0Im7fkED6c8GkiNFI) | [HuggingFace](https://huggingface.co/guoyww/animatediff) | [CivitAI](https://civitai.com/models/108836) | [Baidu NetDisk](https://pan.baidu.com/s/18ZpcSM6poBqxWNHtnyMcxg?pwd=et8y). You only need to download one of `mm_sd_v14.ckpt` | `mm_sd_v15.ckpt`. Put the model weights under `comfyui-animatediff/models/`. DO NOT change model filename.
## Samples
### txt2img
<img width="1254" alt="ComfyUI AnimateDiff Usage" src="https://github.com/ArtVentureX/comfyui-animatediff/assets/133728487/a88e2141-c55f-4bdb-b6ca-9155b6639114">
![AnimateDiff_00001](https://github.com/ArtVentureX/comfyui-animatediff/assets/133728487/e48f148a-886b-4a0d-b589-9fa795b06936)
### img2img
<img width="1121" alt="Screenshot 2023-07-22 at 22 08 00" src="https://github.com/ArtVentureX/comfyui-animatediff/assets/133728487/600f96b0-df21-4437-917f-7eda35ab6363">
![AnimateDiff_00002](https://github.com/ArtVentureX/comfyui-animatediff/assets/133728487/c78d64b9-b308-41ec-9804-bbde654d0b47)
+8
View File
@@ -0,0 +1,8 @@
from .animatediff.logger import logger
from .animatediff.nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
from .animatediff.model_utils import get_available_models
if len(get_available_models()) == 0:
logger.error("No models available. Please download one and put it in models folder")
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
File diff suppressed because it is too large Load Diff
+37
View File
@@ -0,0 +1,37 @@
import sys
import copy
import logging
class ColoredFormatter(logging.Formatter):
COLORS = {
"DEBUG": "\033[0;36m", # CYAN
"INFO": "\033[0;32m", # GREEN
"WARNING": "\033[0;33m", # YELLOW
"ERROR": "\033[0;31m", # RED
"CRITICAL": "\033[0;37;41m", # WHITE ON RED
"RESET": "\033[0m", # RESET COLOR
}
def format(self, record):
colored_record = copy.copy(record)
levelname = colored_record.levelname
seq = self.COLORS.get(levelname, self.COLORS["RESET"])
colored_record.levelname = f"{seq}{levelname}{self.COLORS['RESET']}"
return super().format(colored_record)
# Create a new logger
logger = logging.getLogger("AnimateDiff")
logger.propagate = False
# Add handler if we don't have one.
if not logger.handlers:
handler = logging.StreamHandler(sys.stdout)
handler.setFormatter(
ColoredFormatter("[%(name)s] - %(levelname)s - %(message)s")
)
logger.addHandler(handler)
# Configure logger
logger.setLevel("INFO")
+28
View File
@@ -0,0 +1,28 @@
import os
from huggingface_hub import hf_hub_download
HF_REPO = "guoyww/animatediff"
MODEL_FILES = ["mm_sd_v14.ckpt", "mm_sd_v15.ckpt"]
CURRENT_DIR = os.path.dirname(os.path.realpath(__file__))
MODEL_DIR = os.path.abspath(os.path.join(CURRENT_DIR, "../models"))
def get_available_models():
available_model = [
f for f in MODEL_FILES if os.path.exists(os.path.join(MODEL_DIR, f))
]
return available_model
def download(model_file=MODEL_FILES[-1]):
if not os.path.exists(os.path.join(MODEL_DIR, model_file)):
hf_hub_download(
HF_REPO,
model_file,
cache_dir=MODEL_DIR,
force_download=True,
force_filename=model_file,
)
+345
View File
@@ -0,0 +1,345 @@
import torch
import torch.nn.functional as F
from torch import nn
import math
from einops import rearrange, repeat
from comfy.ldm.modules.attention import FeedForward
from .attention_processor import Attention as CrossAttention
def zero_module(module):
# Zero out the parameters of a module and return it.
for p in module.parameters():
p.detach().zero_()
return module
class MotionWrapper(nn.Module):
def __init__(self):
super().__init__()
self.down_blocks = nn.ModuleList([])
self.up_blocks = nn.ModuleList([])
for i, c in enumerate((320, 640, 1280, 1280)):
self.down_blocks.append(MotionModule(c))
for i, c in enumerate((1280, 1280, 640, 320)):
self.up_blocks.append(MotionModule(c, is_up=True))
class MotionModule(nn.Module):
def __init__(self, in_channels, is_up=False):
super().__init__()
self.motion_modules = nn.ModuleList(
[get_motion_module(in_channels), get_motion_module(in_channels)]
)
if is_up:
self.motion_modules.append(get_motion_module(in_channels))
def get_motion_module(in_channels):
return VanillaTemporalModule(in_channels=in_channels)
class VanillaTemporalModule(nn.Module):
def __init__(
self,
in_channels,
num_attention_heads=8,
num_transformer_block=1,
attention_block_types=("Temporal_Self", "Temporal_Self"),
cross_frame_attention_mode=None,
temporal_position_encoding=True,
temporal_position_encoding_max_len=24,
temporal_attention_dim_div=1,
zero_initialize=True,
):
super().__init__()
self.temporal_transformer = TemporalTransformer3DModel(
in_channels=in_channels,
num_attention_heads=num_attention_heads,
attention_head_dim=in_channels
// num_attention_heads
// temporal_attention_dim_div,
num_layers=num_transformer_block,
attention_block_types=attention_block_types,
cross_frame_attention_mode=cross_frame_attention_mode,
temporal_position_encoding=temporal_position_encoding,
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
)
if zero_initialize:
self.temporal_transformer.proj_out = zero_module(
self.temporal_transformer.proj_out
)
def forward(self, input_tensor, encoder_hidden_states, attention_mask=None):
input_cond, input_uncond = input_tensor.chunk(2)
hidden_states = torch.stack([input_cond, input_uncond], dim=0)
hidden_states = rearrange(hidden_states, "b f c h w -> b c f h w")
hidden_states = self.temporal_transformer(
hidden_states, encoder_hidden_states, attention_mask
)
hidden_states = rearrange(hidden_states, "b c f h w -> b f c h w")
output_cond, output_uncond = hidden_states.chunk(2)
output = torch.cat([output_cond[0], output_uncond[0]], dim=0)
return output
class TemporalTransformer3DModel(nn.Module):
def __init__(
self,
in_channels,
num_attention_heads,
attention_head_dim,
num_layers,
attention_block_types=(
"Temporal_Self",
"Temporal_Self",
),
dropout=0.0,
norm_num_groups=32,
cross_attention_dim=768,
activation_fn="geglu",
attention_bias=False,
upcast_attention=False,
cross_frame_attention_mode=None,
temporal_position_encoding=False,
temporal_position_encoding_max_len=24,
):
super().__init__()
inner_dim = num_attention_heads * attention_head_dim
self.norm = torch.nn.GroupNorm(
num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True
)
self.proj_in = nn.Linear(in_channels, inner_dim)
self.transformer_blocks = nn.ModuleList(
[
TemporalTransformerBlock(
dim=inner_dim,
num_attention_heads=num_attention_heads,
attention_head_dim=attention_head_dim,
attention_block_types=attention_block_types,
dropout=dropout,
norm_num_groups=norm_num_groups,
cross_attention_dim=cross_attention_dim,
activation_fn=activation_fn,
attention_bias=attention_bias,
upcast_attention=upcast_attention,
cross_frame_attention_mode=cross_frame_attention_mode,
temporal_position_encoding=temporal_position_encoding,
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
)
for d in range(num_layers)
]
)
self.proj_out = nn.Linear(inner_dim, in_channels)
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None):
assert (
hidden_states.dim() == 5
), f"Expected hidden_states to have ndim=5, but got ndim={hidden_states.dim()}."
video_length = hidden_states.shape[2]
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
batch, channel, height, weight = hidden_states.shape
residual = hidden_states
hidden_states = self.norm(hidden_states)
inner_dim = hidden_states.shape[1]
hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(
batch, height * weight, inner_dim
)
hidden_states = self.proj_in(hidden_states)
# Transformer Blocks
for block in self.transformer_blocks:
hidden_states = block(
hidden_states,
encoder_hidden_states=encoder_hidden_states,
video_length=video_length,
)
# output
hidden_states = self.proj_out(hidden_states)
hidden_states = (
hidden_states.reshape(batch, height, weight, inner_dim)
.permute(0, 3, 1, 2)
.contiguous()
)
output = hidden_states + residual
output = rearrange(output, "(b f) c h w -> b c f h w", f=video_length)
return output
class TemporalTransformerBlock(nn.Module):
def __init__(
self,
dim,
num_attention_heads,
attention_head_dim,
attention_block_types=(
"Temporal_Self",
"Temporal_Self",
),
dropout=0.0,
norm_num_groups=32,
cross_attention_dim=768,
activation_fn="geglu",
attention_bias=False,
upcast_attention=False,
cross_frame_attention_mode=None,
temporal_position_encoding=False,
temporal_position_encoding_max_len=24,
):
super().__init__()
attention_blocks = []
norms = []
for block_name in attention_block_types:
attention_blocks.append(
VersatileAttention(
attention_mode=block_name.split("_")[0],
cross_attention_dim=cross_attention_dim
if block_name.endswith("_Cross")
else None,
query_dim=dim,
heads=num_attention_heads,
dim_head=attention_head_dim,
dropout=dropout,
bias=attention_bias,
upcast_attention=upcast_attention,
cross_frame_attention_mode=cross_frame_attention_mode,
temporal_position_encoding=temporal_position_encoding,
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
)
)
norms.append(nn.LayerNorm(dim))
self.attention_blocks = nn.ModuleList(attention_blocks)
self.norms = nn.ModuleList(norms)
self.ff = FeedForward(dim, dropout=dropout, glu=(activation_fn == "geglu"))
self.ff_norm = nn.LayerNorm(dim)
def forward(
self,
hidden_states,
encoder_hidden_states=None,
attention_mask=None,
video_length=None,
):
for attention_block, norm in zip(self.attention_blocks, self.norms):
norm_hidden_states = norm(hidden_states)
hidden_states = (
attention_block(
norm_hidden_states,
encoder_hidden_states=encoder_hidden_states
if attention_block.is_cross_attention
else None,
video_length=video_length,
)
+ hidden_states
)
hidden_states = self.ff(self.ff_norm(hidden_states)) + hidden_states
output = hidden_states
return output
class PositionalEncoding(nn.Module):
def __init__(self, d_model, dropout=0.0, max_len=24):
super().__init__()
self.dropout = nn.Dropout(p=dropout)
position = torch.arange(max_len).unsqueeze(1)
div_term = torch.exp(
torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model)
)
pe = torch.zeros(1, max_len, d_model)
pe[0, :, 0::2] = torch.sin(position * div_term)
pe[0, :, 1::2] = torch.cos(position * div_term)
self.register_buffer("pe", pe)
def forward(self, x):
x = x + self.pe[:, : x.size(1)]
return self.dropout(x)
class VersatileAttention(CrossAttention):
def __init__(
self,
attention_mode=None,
cross_frame_attention_mode=None,
temporal_position_encoding=False,
temporal_position_encoding_max_len=24,
*args,
**kwargs,
):
super().__init__(*args, **kwargs)
assert attention_mode == "Temporal"
self.attention_mode = attention_mode
self.is_cross_attention = kwargs["cross_attention_dim"] is not None
self.pos_encoder = (
PositionalEncoding(
kwargs["query_dim"],
dropout=0.0,
max_len=temporal_position_encoding_max_len,
)
if (temporal_position_encoding and attention_mode == "Temporal")
else None
)
def extra_repr(self):
return f"(Module Info) Attention_Mode: {self.attention_mode}, Is_Cross_Attention: {self.is_cross_attention}"
def forward(
self,
hidden_states,
encoder_hidden_states=None,
attention_mask=None,
video_length=None,
**cross_attention_kwargs,
):
if self.added_kv_proj_dim is not None:
raise NotImplementedError
if self.attention_mode != "Temporal":
raise NotImplementedError
d = hidden_states.shape[1]
hidden_states = rearrange(
hidden_states, "(b f) d c -> (b d) f c", f=video_length
)
if self.pos_encoder is not None:
hidden_states = self.pos_encoder(hidden_states)
encoder_hidden_states = (
repeat(encoder_hidden_states, "b n c -> (b d) n c", d=d)
if encoder_hidden_states is not None
else encoder_hidden_states
)
hidden_states = super().forward(
hidden_states,
encoder_hidden_states,
attention_mask,
**cross_attention_kwargs,
)
hidden_states = rearrange(hidden_states, "(b d) f c -> (b f) d c", d=d)
return hidden_states
+228
View File
@@ -0,0 +1,228 @@
import os
import hashlib
import torch
import numpy as np
from PIL import Image
from typing import Dict
import folder_paths
import comfy.ldm.modules.diffusionmodules.openaimodel as openaimodel
import comfy.model_management as model_management
from comfy.ldm.modules.attention import SpatialTransformer
from comfy.utils import load_torch_file
from comfy.sd import ModelPatcher, calculate_parameters
from .logger import logger
from .motion_module import MotionWrapper, VanillaTemporalModule
from .model_utils import MODEL_DIR, get_available_models
orig_forward_timestep_embed = openaimodel.forward_timestep_embed
def forward_timestep_embed(
ts, x, emb, context=None, transformer_options={}, output_shape=None
):
for layer in ts:
if isinstance(layer, openaimodel.TimestepBlock):
x = layer(x, emb)
elif isinstance(layer, VanillaTemporalModule):
x = layer(x, context)
elif isinstance(layer, SpatialTransformer):
x = layer(x, context, transformer_options)
transformer_options["current_index"] += 1
elif isinstance(layer, openaimodel.Upsample):
x = layer(x, output_shape=output_shape)
else:
x = layer(x)
return x
openaimodel.forward_timestep_embed = forward_timestep_embed
motion_module: MotionWrapper = None
class AnimateDiffLoader:
def __init__(self) -> None:
self.last_injected_model_hash = set()
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"model_name": (get_available_models(),),
"width": ("INT", {"default": 512, "min": 64, "max": 1024, "step": 8}),
"height": ("INT", {"default": 512, "min": 64, "max": 1024, "step": 8}),
"frame_number": (
"INT",
{"default": 16, "min": 2, "max": 24, "step": 1},
),
},
"optional": {
"init_latent": ("LATENT",),
},
}
RETURN_TYPES = ("MODEL", "LATENT")
CATEGORY = "Animate Diff"
FUNCTION = "inject_motion_modules"
def inject_motion_modules(
self,
model: ModelPatcher,
model_name: str,
width: int,
height: int,
frame_number=16,
init_latent: Dict[str, torch.Tensor] = None,
):
model = model.clone()
model_path = os.path.join(MODEL_DIR, model_name)
global motion_module
if motion_module is None:
logger.info(f"Loading motion module {model_name} from {model_path}")
mm_state_dict = load_torch_file(model_path)
motion_module = MotionWrapper()
parameters = calculate_parameters(mm_state_dict, "")
usefp16 = model_management.should_use_fp16(model_params=parameters)
if usefp16:
print("Using fp16, converting motion module to fp16")
motion_module.half()
offload_device = model_management.unet_offload_device()
motion_module = motion_module.to(offload_device)
motion_module.load_state_dict(mm_state_dict)
unet = model.model.diffusion_model
if self.calculate_model_hash(unet) in self.last_injected_model_hash:
logger.info(f"Motion module already injected, skipping injection.")
else:
logger.info(f"Injecting motion module into UNet input blocks.")
for mm_idx, unet_idx in enumerate([1, 2, 4, 5, 7, 8, 10, 11]):
mm_idx0, mm_idx1 = mm_idx // 2, mm_idx % 2
unet.input_blocks[unet_idx].append(
motion_module.down_blocks[mm_idx0].motion_modules[mm_idx1]
)
logger.info(f"Injecting motion module into UNet output blocks.")
for unet_idx in range(12):
mm_idx0, mm_idx1 = unet_idx // 3, unet_idx % 3
if unet_idx % 2 == 2:
unet.output_blocks[unet_idx].insert(
-1, motion_module.up_blocks[mm_idx0].motion_modules[mm_idx]
)
else:
unet.output_blocks[unet_idx].append(
motion_module.up_blocks[mm_idx0].motion_modules[mm_idx1]
)
self.last_injected_model_hash.add(self.calculate_model_hash(unet))
if init_latent is None:
latent = torch.zeros([frame_number, 4, width // 8, height // 8]).cpu()
else:
# clone value of first frame
latent = init_latent["samples"].clone().cpu()
# repeat for all frames
latent = latent.repeat(frame_number, 1, 1, 1)
return (model, {"samples": latent})
def calculate_model_hash(self, unet):
t = unet.input_blocks[1]
m = hashlib.sha256()
for buf in t.buffers():
m.update(buf.numpy().view(np.uint8))
return m.hexdigest()
class AnimateDiffCombine:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"images": ("IMAGE",),
"frame_rate": (
"INT",
{"default": 8, "min": 1, "max": 24, "step": 1},
),
"loop_count": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1}),
"save_image": (("Enabled", "Disabled"),),
"filename_prefix": ("STRING", {"default": "AnimateDiff"}),
}
}
RETURN_TYPES = ("IMAGE",)
CATEGORY = "Animate Diff"
OUTPUT_NODE = True
FUNCTION = "generate_gif"
def generate_gif(
self,
images,
frame_rate: int,
loop_count: int,
save_image="Enabled",
filename_prefix="AnimateDiff",
):
import imageio
# convert images to numpy
image_nps = []
for image in images:
img = 255.0 * image.cpu().numpy()
img = np.clip(img, 0, 255).astype(np.uint8)
image_nps.append(img)
# save image
output_dir = (
folder_paths.get_output_directory()
if save_image == "Enabled"
else folder_paths.get_temp_directory()
)
(
full_output_folder,
filename,
counter,
subfolder,
_,
) = folder_paths.get_save_image_path(filename_prefix, output_dir)
file = f"{filename}_{counter:05}.gif"
file_path = os.path.join(full_output_folder, file)
# save gif
imageio.mimsave(
file_path,
image_nps,
duration=round(1000 / frame_rate),
loop=loop_count,
)
# load saved image back as torch tensor
saved = Image.open(file_path)
saved = saved.convert("RGB")
saved = np.array(saved).astype(np.float32) / 255.0
saved = torch.from_numpy(saved)[None,]
previews = [
{
"filename": file,
"subfolder": subfolder,
"type": "output" if save_image == "Enabled" else "temp",
}
]
return {"ui": {"images": previews}, "result": (saved,)}
NODE_CLASS_MAPPINGS = {
"AnimateDiffLoader": AnimateDiffLoader,
"AnimateDiffCombile": AnimateDiffCombine,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"AnimateDiffLoader": "Animate Diff Loader",
"AnimateDiffCombile": "Animate Diff Combine",
}
+14
View File
@@ -0,0 +1,14 @@
#このコードはhttps://github.com/kohya-ss/sd-scripts/blob/main/finetune/tag_images_by_wd14_tagger.pyを参考にしていますというかパクっています。
from huggingface_hub import hf_hub_download
import cv2
import numpy as np
import os
IMAGE_SIZE = 448
TAGGER_REPO = "furusu/wd-v1-4-tagger-pytorch"
TAGGER_FILE = "wd-v1-4-vit-tagger-v2.ckpt"
def download(path):
if not os.path.exists(os.path.join(path, TAGGER_FILE)):
hf_hub_download(TAGGER_REPO, TAGGER_FILE, cache_dir=path, force_download=True, force_filename=TAGGER_FILE)
View File