first commit
This commit is contained in:
@@ -0,0 +1,4 @@
|
|||||||
|
__pycache__
|
||||||
|
/venv
|
||||||
|
.vscode
|
||||||
|
*.ckpt
|
||||||
@@ -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">
|
||||||
|
|
||||||
|

|
||||||
|
|
||||||
|
### 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">
|
||||||
|
|
||||||
|

|
||||||
@@ -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
@@ -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")
|
||||||
@@ -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,
|
||||||
|
)
|
||||||
@@ -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
|
||||||
|
|
||||||
@@ -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",
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
Reference in New Issue
Block a user