Author SHA1 Message Date
Tung Nguyen e1fb362804 fix missing block_type when init MotionModule 2023-09-16 06:25:36 +07:00
Tung Nguyen 61fc641564 refactor & apply some changes from Kosinkadink fork 2023-09-16 06:21:10 +07:00
Tung Nguyen 0c20aa71fc add workflow.json file 2023-09-15 21:58:38 +07:00
Tung Nguyen b8381c96ef update README 2023-09-15 21:56:35 +07:00
ArtVenture 1fa365a372 Merge pull request #22 from ArtVentureX/feat/support_v2
Support animatediff v2 & improve image quality
2023-09-15 17:43:57 +07:00
5 changed files with 671 additions and 352 deletions
+37 -15
View File
@@ -7,18 +7,31 @@
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.
#### Update 2023/09/15
- You can now use community models from [manshoety/AD_Stabilized_Motion](https://huggingface.co/manshoety/AD_Stabilized_Motion) or [CiaraRowles/TemporalDiff](https://huggingface.co/CiaraRowles/TemporalDiff)
- Supports AnimateDiff v2 [mm_sd_v15_v2.ckpt](https://huggingface.co/guoyww/animatediff/blob/main/mm_sd_v15_v2.ckpt) model
- Fix image is grayed out.
- New node: **AnimateDiffSampler** and **AnimateDiffLoader**
- Mostly the same with `KSampler`
- Use `AnimateDiffLoader` to load the motion module
- `inject_method`: should left default. See [this issue](https://github.com/ArtVentureX/comfyui-animatediff#gif-has-wartermark-after-update-to-the-latest-version) for more details.
- `frame_number`: animation length
<img width="506" alt="image" src="https://github.com/ArtVentureX/comfyui-animatediff/assets/133728487/f22d6b36-ce36-44cc-80e8-dffe6f77b296">
#### Example Workflow
<img width="1311" alt="image" src="https://github.com/ArtVentureX/comfyui-animatediff/assets/133728487/b7164539-bc58-4ef9-b178-d914e833805e">
Workflow file: https://github.com/ArtVentureX/comfyui-animatediff/blob/main/workflow.json
## Samples
### txt2img
![23b44c29-29e8-4f48-ab3c-4df87c90c13f](https://github.com/ArtVentureX/comfyui-animatediff/assets/133728487/97efb96f-3d3d-4976-8789-78b88f89b2eb)
<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)
![25f6c60c-f8ac-4abe-984f-1559c355d7f6](https://github.com/ArtVentureX/comfyui-animatediff/assets/133728487/c39b26f7-a2af-4dc4-902f-c363e2e6f39a)
## Known Issues
@@ -26,11 +39,23 @@
![AnimateDiff_00007_](https://github.com/ArtVentureX/comfyui-animatediff/assets/8894763/e6cd53cb-9878-45da-a58a-a15851882386)
This is usually due to memory (VRAM) is not enough to process the whole image batch at the same time. Try reduce the image size and frame number.
See: https://github.com/continue-revolution/sd-webui-animatediff/issues/38
Main reasons:
- Promt are too long (more than 75 tokens)
- Resolution are too high
- Number of frame too high
Work around:
- Shorter your prompt and negative prompt
- Reduce resolution. AnimateDiff is trained on 512x512 images so it works best with 512x512 output.
- Shouldn't generate longer than 16 frames. AnimateDiff is trained to output the best results with 16 frames.
### GIF has Wartermark after update to the latest version
See https://github.com/continue-revolution/sd-webui-animatediff/issues/31
See: https://github.com/continue-revolution/sd-webui-animatediff/issues/31
As mentioned in the issue thread, it seems to be due to the training dataset. The new version is the correct implementation and produces smoother GIFs compared to the older version.
@@ -49,7 +74,4 @@ As mentioned in the issue thread, it seems to be due to the training dataset. Th
</tr>
</table>
I played around with both version and found that the watermark only present in some models, not always. So I've brought back the old method and also created a new node with the new method. You can try both to find the best fit for each model.
![Screenshot 2023-07-28 at 18 14 14](https://github.com/ArtVentureX/comfyui-animatediff/assets/133728487/25cf6092-3e67-435e-86cc-43614ca7d6aa)
I played around with both version and found that the watermark only present in some models, not always. To use the **old (legacy)** method, change `injection_method` to `legacy` in the `AnimateDiffSampler` node.
+1 -32
View File
@@ -3,8 +3,6 @@ import hashlib
import folder_paths
from .logger import logger
folder_paths.folder_names_and_paths["AnimateDiff"] = (
[
@@ -14,17 +12,6 @@ folder_paths.folder_names_and_paths["AnimateDiff"] = (
folder_paths.supported_pt_extensions,
)
known_models = {
"aa7fd8a200a89031edd84487e2a757c5315460eca528fa70d4b3885c399bffd5": "mm_sd_v14.ckpt",
"cf16ea656cb16124990c8e2c70a29c793f9841f3a2223073fac8bd89ebd9b69a": "mm_sd_v15.ckpt",
"0aaf157b9c51a0ae07cb5d9ea7c51299f07bddc6f52025e1f9bb81cd763631df": "mm-Stabilized_high.pth",
"39de8b71b1c09f10f4602f5d585d82771a60d3cf282ba90215993e06afdfe875": "mm-Stabilized_mid.pth",
"3cb569f7ce3dc6a10aa8438e666265cb9be3120d8f205de6a456acf46b6c99f4": "temporaldiff-v1-animatediff.ckpt",
"69ed0f5fef82b110aca51bcab73b21104242bc65d6ab4b8b2a2a94d31cad1bf0": "mm_sd_v15_v2.ckpt",
}
v2_models = ["69ed0f5fef82b110aca51bcab73b21104242bc65d6ab4b8b2a2a94d31cad1bf0"]
def get_available_models():
return folder_paths.get_filename_list("AnimateDiff")
@@ -34,25 +21,7 @@ def get_model_path(model_name):
return folder_paths.get_full_path("AnimateDiff", model_name)
def sha256_file(file_path):
def get_model_hash(file_path):
with open(file_path, "rb") as f:
bytes = f.read() # read entire file as bytes
return hashlib.sha256(bytes).hexdigest()
def validate_mm_model(model_name):
model_path = get_model_path(model_name)
model_hash = sha256_file(model_path)
if model_hash in known_models:
logger.info(f"You are using {model_name}, which has been tested and supported.")
else:
logger.warn(
f"Your model {model_name} has not been tested and supported."
"Either your download is incomplete or your model has not been tested. "
"Please use at your own risk."
)
using_v2 = model_hash in v2_models
return (model_hash, using_v2)
+105 -31
View File
@@ -1,12 +1,12 @@
import os
import torch
import torch.nn.functional as F
from torch import nn
from torch import Tensor, nn
import math
from einops import rearrange, repeat
from comfy.ldm.modules.attention import FeedForward
from .attention_processor import Attention as CrossAttention
from comfy.utils import load_torch_file
from comfy.ldm.modules.attention import FeedForward, CrossAttention
def zero_module(module):
@@ -15,41 +15,108 @@ def zero_module(module):
p.detach().zero_()
return module
# Merge from https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved
def get_encoding_max_len(mm_state_dict: dict[str, Tensor]) -> int:
# use pos_encoder.pe entries to determine max length - [1, {max_length}, {320|640|1280}]
for key in mm_state_dict.keys():
if key.endswith("pos_encoder.pe"):
return mm_state_dict[key].size(1) # get middle dim
raise ValueError(f"No pos_encoder.pe found in mm_state_dict")
def has_mid_block(mm_state_dict: dict[str, Tensor]):
# check if keys contain mid_block
for key in mm_state_dict.keys():
if key.startswith("mid_block."):
return True
return False
class MotionWrapper(nn.Module):
def __init__(self, mm_hash, is_v2 = False):
def __init__(self, mm_type: str, encoding_max_len: int = 24, is_v2=False):
super().__init__()
if is_v2:
max_len = 32
else:
max_len = 24
self.mm_type = mm_type
self.is_v2 = is_v2
self.down_blocks = nn.ModuleList([])
self.up_blocks = nn.ModuleList([])
self.mid_block = None
for c in (320, 640, 1280, 1280):
self.down_blocks.append(MotionModule(c, max_len=max_len))
self.down_blocks.append(
MotionModule(c, BlockType.DOWN, encoding_max_len=encoding_max_len)
)
for c in (1280, 1280, 640, 320):
self.up_blocks.append(MotionModule(c, is_up=True, max_len=max_len))
self.up_blocks.append(
MotionModule(c, BlockType.UP, encoding_max_len=encoding_max_len)
)
if is_v2:
self.mid_block = MotionModule(1280, max_len=max_len, is_mid=is_v2)
self.mm_hash = mm_hash
self.is_v2 = is_v2
self.mid_block = MotionModule(
1280, BlockType.MID, encoding_max_len=encoding_max_len
)
@classmethod
def from_pretrained(cls, checkpoint_path: str):
mm_state_dict = load_torch_file(checkpoint_path)
mm_type = os.path.basename(checkpoint_path)
encoding_max_len = get_encoding_max_len(mm_state_dict)
is_v2 = has_mid_block(mm_state_dict)
mm = cls(mm_type, encoding_max_len=encoding_max_len, is_v2=is_v2)
mm.load_state_dict(mm_state_dict)
return mm
def set_video_length(self, video_length: int):
for block in self.down_blocks:
block.set_video_length(video_length)
for block in self.up_blocks:
block.set_video_length(video_length)
if self.mid_block is not None:
self.mid_block.set_video_length(video_length)
class BlockType:
UP = "up"
DOWN = "down"
MID = "mid"
class MotionModule(nn.Module):
def __init__(self, in_channels, is_up=False, is_mid=False, max_len=24):
def __init__(
self,
in_channels,
block_type: BlockType,
encoding_max_len=24,
):
super().__init__()
if is_mid:
self.motion_modules = nn.ModuleList([get_motion_module(in_channels, max_len)])
self.block_type = block_type
if block_type == BlockType.MID:
self.motion_modules = nn.ModuleList(
[get_motion_module(in_channels, encoding_max_len)]
)
else:
self.motion_modules = nn.ModuleList(
[get_motion_module(in_channels, max_len), get_motion_module(in_channels, max_len)]
[
get_motion_module(in_channels, encoding_max_len),
get_motion_module(in_channels, encoding_max_len),
]
)
if is_up:
self.motion_modules.append(get_motion_module(in_channels, max_len))
if block_type == BlockType.UP:
self.motion_modules.append(
get_motion_module(in_channels, encoding_max_len)
)
def set_video_length(self, video_length: int):
for motion_module in self.motion_modules:
motion_module.set_video_length(video_length)
def get_motion_module(in_channels, max_len):
return VanillaTemporalModule(in_channels=in_channels, temporal_position_encoding_max_len=max_len)
return VanillaTemporalModule(
in_channels=in_channels, temporal_position_encoding_max_len=max_len
)
class VanillaTemporalModule(nn.Module):
@@ -85,8 +152,13 @@ class VanillaTemporalModule(nn.Module):
self.temporal_transformer.proj_out
)
def set_video_length(self, video_length: int):
self.temporal_transformer.set_video_length(video_length)
def forward(self, input_tensor, encoder_hidden_states, attention_mask=None):
return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask)
return self.temporal_transformer(
input_tensor, encoder_hidden_states, attention_mask
)
class TemporalTransformer3DModel(nn.Module):
@@ -140,10 +212,12 @@ class TemporalTransformer3DModel(nn.Module):
]
)
self.proj_out = nn.Linear(inner_dim, in_channels)
self.video_length = 16
def set_video_length(self, video_length: int):
self.video_length = video_length
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None):
video_length = hidden_states.shape[0] // 2 # TODO: config this value in scripts
batch, channel, height, weight = hidden_states.shape
residual = hidden_states
@@ -159,7 +233,7 @@ class TemporalTransformer3DModel(nn.Module):
hidden_states = block(
hidden_states,
encoder_hidden_states=encoder_hidden_states,
video_length=video_length,
video_length=self.video_length,
)
# output
@@ -204,15 +278,15 @@ class TemporalTransformerBlock(nn.Module):
attention_blocks.append(
VersatileAttention(
attention_mode=block_name.split("_")[0],
cross_attention_dim=cross_attention_dim
context_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,
# bias=attention_bias, # remove for Comfy CrossAttention
# upcast_attention=upcast_attention, # remove for Comfy CrossAttention
cross_frame_attention_mode=cross_frame_attention_mode,
temporal_position_encoding=temporal_position_encoding,
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
@@ -284,7 +358,7 @@ class VersatileAttention(CrossAttention):
assert attention_mode == "Temporal"
self.attention_mode = attention_mode
self.is_cross_attention = kwargs["cross_attention_dim"] is not None
self.is_cross_attention = kwargs["context_dim"] is not None
self.pos_encoder = (
PositionalEncoding(
@@ -327,8 +401,8 @@ class VersatileAttention(CrossAttention):
hidden_states = super().forward(
hidden_states,
encoder_hidden_states,
attention_mask,
**cross_attention_kwargs,
value=None,
mask=attention_mask,
)
hidden_states = rearrange(hidden_states, "(b d) f c -> (b f) d c", d=d)
+77 -274
View File
@@ -1,9 +1,10 @@
import os
import json
import hashlib
import torch
import numpy as np
from typing import Dict, List, Tuple
from typing import Dict, List
from torch import Tensor
from torch.nn.functional import group_norm
from PIL import Image
from PIL.PngImagePlugin import PngInfo
from einops import rearrange
@@ -11,19 +12,14 @@ from einops import rearrange
import folder_paths
import comfy.ldm.modules.diffusionmodules.openaimodel as openaimodel
import comfy.model_management as model_management
from comfy.model_base import BaseModel
from comfy.ldm.modules.attention import SpatialTransformer
from comfy.ldm.modules.diffusionmodules.util import GroupNorm32
from comfy.utils import load_torch_file, calculate_parameters
from comfy.model_patcher import ModelPatcher
from comfy.cli_args import args as cli_args
from nodes import KSampler
from .logger import logger
from .motion_module import MotionWrapper, VanillaTemporalModule
from .model_utils import get_available_models, get_model_path, validate_mm_model
orig_forward_timestep_embed = openaimodel.forward_timestep_embed
groupnorm32_original_forward = GroupNorm32.forward
from .model_utils import get_available_models, get_model_path, get_model_hash
def forward_timestep_embed(
@@ -44,44 +40,38 @@ def forward_timestep_embed(
return x
def groupnorm32_mm_forward(self, x):
x = rearrange(x, "(b f) c h w -> b c f h w", b=2)
x = groupnorm32_original_forward(self, x)
x = rearrange(x, "b c f h w -> (b f) c h w", b=2)
return x
def groupnorm_mm_factory(video_length: int):
def groupnorm_mm_forward(self, input: Tensor) -> Tensor:
# axes_factor normalizes batch based on total conds and unconds passed in batch;
# the conds and unconds per batch can change based on VRAM optimizations that may kick in
axes_factor = input.size(0) // video_length
input = rearrange(input, "(b f) c h w -> b c f h w", b=axes_factor)
input = group_norm(input, self.num_groups, self.weight, self.bias, self.eps)
input = rearrange(input, "b c f h w -> (b f) c h w", b=axes_factor)
return input
return groupnorm_mm_forward
orig_forward_timestep_embed = openaimodel.forward_timestep_embed
orig_maximum_batch_area = model_management.maximum_batch_area
orig_groupnorm_forward = torch.nn.GroupNorm.forward
openaimodel.forward_timestep_embed = forward_timestep_embed
motion_modules: Dict[str, MotionWrapper] = {}
original_model_hashs = set()
injected_model_hashs: Dict[str, Tuple[str, str]] = {}
def calculate_model_hash(unet):
t = unet.input_blocks[1]
m = hashlib.sha256()
for buf in t.buffers():
m.update(buf.cpu().numpy().view(np.uint8))
return m.hexdigest()
def load_motion_module(model_name: str):
model_path = get_model_path(model_name)
model_hash, is_v2 = validate_mm_model(model_name)
model_hash = get_model_hash(model_path)
if model_hash not in motion_modules:
logger.info(f"Loading motion module {model_name}")
mm_state_dict = load_torch_file(model_path)
motion_module = MotionWrapper(model_name, is_v2=is_v2)
parameters = calculate_parameters(mm_state_dict, "")
usefp16 = model_management.should_use_fp16(model_params=parameters)
if usefp16:
logger.info("Using fp16, converting motion module to fp16")
motion_module = MotionWrapper.from_pretrained(model_path)
if not cli_args.force_fp32:
logger.info(f"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)
motion_modules[model_hash] = motion_module
return motion_modules[model_hash]
@@ -176,81 +166,6 @@ ejectors = {
}
class AnimateDiffLoaderLegacy:
def __init__(self) -> None:
self.version = "legacy"
@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",),
},
}
@classmethod
def IS_CHANGED(s, model: ModelPatcher):
unet = model.model.diffusion_model
# return calculate_model_hash(unet) not in injected_model_hashs
return hasattr(unet, "motion_module") and unet.motion_module is not None
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,
):
motion_module = load_motion_module(model_name)
model = model.clone()
unet = model.model.diffusion_model
unet_hash = calculate_model_hash(unet)
need_inject = unet_hash not in injected_model_hashs
if unet_hash in injected_model_hashs:
(mm_hash, version) = injected_model_hashs[unet_hash]
if version != self.version or mm_hash != motion_module.mm_hash:
# injected by another motion module, unload first
logger.info(f"Ejecting motion module {mm_hash} version {version}.")
ejectors[version](unet)
need_inject = True
else:
logger.info(f"Motion module already injected, skipping injection.")
if need_inject:
logger.info(f"Injecting motion module {model_name} version {self.version}.")
injectors[self.version](unet, motion_module)
unet_hash = calculate_model_hash(unet)
injected_model_hashs[unet_hash] = (motion_module.mm_hash, self.version)
if init_latent is None:
latent = torch.zeros([frame_number, 4, height // 8, width // 8]).cpu()
else:
# clone value of first frame
latent = init_latent["samples"][:1, :, :, :].clone().cpu()
# repeat for all frames
latent = latent.repeat(frame_number, 1, 1, 1)
return (model, {"samples": latent})
class AnimateDiffModuleLoader:
@classmethod
def INPUT_TYPES(s):
@@ -273,105 +188,6 @@ class AnimateDiffModuleLoader:
return (motion_module,)
class AnimateDiffLoader:
def __init__(self) -> None:
self.version = "v1"
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"init_latent": ("LATENT",),
"model_name": (get_available_models(),),
"frame_number": (
"INT",
{"default": 16, "min": 2, "max": 32, "step": 1},
),
},
}
@classmethod
def IS_CHANGED(s, model: ModelPatcher, _):
unet = model.model.diffusion_model
# return calculate_model_hash(unet) not in injected_model_hashs
return hasattr(unet, "motion_module") and unet.motion_module is not None
RETURN_TYPES = ("MODEL", "LATENT")
CATEGORY = "Animate Diff"
FUNCTION = "inject_motion_modules"
def inject_motion_modules(
self,
model: ModelPatcher,
init_latent: Dict[str, torch.Tensor],
model_name: str,
frame_number=16,
):
motion_module = load_motion_module(model_name)
model = model.clone()
unet = model.model.diffusion_model
unet_hash = calculate_model_hash(unet)
need_inject = unet_hash not in injected_model_hashs
if unet_hash in injected_model_hashs:
(mm_type, version) = injected_model_hashs[unet_hash]
if version != self.version or mm_type != motion_module.mm_hash:
# injected by another motion module, unload first
logger.info(f"Ejecting motion module {mm_type} version {version}.")
ejectors[version](unet)
need_inject = True
else:
logger.info(f"Motion module already injected, skipping injection.")
if need_inject:
logger.info(f"Injecting motion module {model_name} version {self.version}.")
injectors[self.version](unet, motion_module)
unet_hash = calculate_model_hash(unet)
injected_model_hashs[unet_hash] = (motion_module.mm_hash, self.version)
init_frames = len(init_latent["samples"])
samples = init_latent["samples"][:init_frames, :, :, :].clone().cpu()
if init_frames < frame_number:
last_frame = samples[-1].unsqueeze(0)
repeated_last_frames = last_frame.repeat(
frame_number - init_frames, 1, 1, 1
)
samples = torch.cat((samples, repeated_last_frames), dim=0)
return (model, {"samples": samples})
class AnimateDiffUnload:
@classmethod
def INPUT_TYPES(s):
return {"required": {"model": ("MODEL",)}}
@classmethod
def IS_CHANGED(s, model: ModelPatcher):
unet = model.model.diffusion_model
return calculate_model_hash(unet) in injected_model_hashs
RETURN_TYPES = ("MODEL",)
CATEGORY = "Animate Diff"
FUNCTION = "unload_motion_modules"
def unload_motion_modules(self, model: ModelPatcher):
model = model.clone()
unet = model.model.diffusion_model
model_hash = calculate_model_hash(unet)
if model_hash in injected_model_hashs:
(model_name, version) = injected_model_hashs[model_hash]
logger.info(f"Ejecting motion module {model_name} version {version}.")
ejectors[version](unet)
else:
logger.info(f"Motion module not injected, skip unloading.")
return (model,)
class AnimateDiffSampler(KSampler):
@classmethod
def INPUT_TYPES(s):
@@ -394,66 +210,56 @@ class AnimateDiffSampler(KSampler):
def __init__(self) -> None:
super().__init__()
self.prev_beta = None
self.prev_alpha_cumprod = None
self.prev_alpha_cumprod_prev = None
self.prev_linear_start = None
self.prev_linear_end = None
def override_ddim_alpha(self, model):
logger.info(f"Setting DDIM alpha.")
device = model_management.unet_offload_device()
beta_start = 0.00085
beta_end = 0.012
betas = torch.linspace(
beta_start,
beta_end,
model.num_timesteps,
dtype=torch.float32,
device=device,
def override_beta_schedule(self, model: BaseModel):
logger.info(f"Override beta schedule.")
self.prev_beta = model.get_buffer("betas")
self.prev_linear_start = model.linear_start
self.prev_linear_end = model.linear_end
model.register_schedule(
given_betas=None,
beta_schedule="sqrt_linear",
timesteps=1000,
linear_start=0.00085,
linear_end=0.012,
cosine_s=8e-3,
)
alphas = 1.0 - betas
alphas_cumprod = torch.cumprod(alphas, dim=0)
alphas_cumprod_prev = torch.cat(
(
torch.tensor([1.0], dtype=torch.float32, device=device),
alphas_cumprod[:-1],
)
)
self.prev_beta = model.betas
model.betas = betas
self.prev_alpha_cumprod = model.alphas_cumprod
model.alphas_cumprod = alphas_cumprod
self.prev_alpha_cumprod_prev = model.alphas_cumprod_prev
model.alphas_cumprod_prev = alphas_cumprod_prev
def restore_ddim_alpha(self, model):
logger.info(f"Restoring DDIM alpha.")
model.betas = self.prev_beta
model.alphas_cumprod = self.prev_alpha_cumprod
model.alphas_cumprod_prev = self.prev_alpha_cumprod_prev
def restore_beta_schedule(self, model: BaseModel):
logger.info(f"Restoring beta schedule.")
model.register_schedule(
given_betas=self.prev_beta,
linear_start=self.prev_linear_start,
linear_end=self.prev_linear_end,
)
self.prev_beta = None
self.prev_alpha_cumprod = None
self.prev_alpha_cumprod_prev = None
self.prev_linear_start = None
self.prev_linear_end = None
def inject_motion_module(self, model, motion_module, inject_method):
def inject_motion_module(
self, model, motion_module: MotionWrapper, inject_method: str, frame_number: int
):
model = model.clone()
unet = model.model.diffusion_model
logger.info(f"Injecting motion module with method {inject_method}.")
injectors[inject_method](unet, motion_module)
self.override_ddim_alpha(model.model)
self.override_beta_schedule(model.model)
if not motion_module.is_v2:
logger.info(f"Hacking GroupNorm32 forward function.")
GroupNorm32.forward = groupnorm32_mm_forward
logger.info(f"Hacking GroupNorm.forward function.")
torch.nn.GroupNorm.forward = groupnorm_mm_factory(frame_number)
return model
def eject_motion_module(self, model, inject_method):
unet = model.model.diffusion_model
self.restore_ddim_alpha(model.model)
self.restore_beta_schedule(model.model)
if not unet.motion_module.is_v2:
logger.info(f"Restore GroupNorm32 forward function.")
GroupNorm32.forward = groupnorm32_original_forward
torch.nn.GroupNorm.forward = orig_groupnorm_forward
logger.info(f"Ejecting motion module with method {inject_method}.")
ejectors[inject_method](unet)
@@ -474,7 +280,9 @@ class AnimateDiffSampler(KSampler):
latent_image,
denoise=1.0,
):
model = self.inject_motion_module(model, motion_module, inject_method)
model = self.inject_motion_module(
model, motion_module, inject_method, frame_number
)
init_frames = len(latent_image["samples"])
samples = latent_image["samples"][:init_frames, :, :, :].clone().cpu()
@@ -488,22 +296,23 @@ class AnimateDiffSampler(KSampler):
latent_image = {"samples": samples}
results = super().sample(
model,
seed,
steps,
cfg,
sampler_name,
scheduler,
positive,
negative,
latent_image,
denoise=1.0,
)
self.eject_motion_module(model, inject_method)
return results
try:
return super().sample(
model,
seed,
steps,
cfg,
sampler_name,
scheduler,
positive,
negative,
latent_image,
denoise=denoise,
)
except:
raise
finally:
self.eject_motion_module(model, inject_method)
class AnimateDiffCombine:
@@ -603,17 +412,11 @@ class AnimateDiffCombine:
NODE_CLASS_MAPPINGS = {
# "AnimateDiffLoader": AnimateDiffLoaderLegacy,
# "AnimateDiffLoader_v2": AnimateDiffLoader,
# "AnimateDiffUnload": AnimateDiffUnload,
"AnimateDiffModuleLoader": AnimateDiffModuleLoader,
"AnimateDiffCombine": AnimateDiffCombine,
"AnimateDiffSampler": AnimateDiffSampler,
}
NODE_DISPLAY_NAME_MAPPINGS = {
# "AnimateDiffLoader": "[DEPRECATED] Animate Diff Loader Legacy",
# "AnimateDiffLoader_v2": "[DEPRECATED] Animate Diff Loader",
# "AnimateDiffUnload": "[DEPRECATED] Animate Diff Unload",
"AnimateDiffModuleLoader": "Animate Diff Module Loader",
"AnimateDiffSampler": "Animate Diff Sampler",
"AnimateDiffCombine": "Animate Diff Combine",
+451
View File
@@ -0,0 +1,451 @@
{
"last_node_id": 20,
"last_link_id": 35,
"nodes": [
{
"id": 6,
"type": "CLIPTextEncode",
"pos": [
415,
186
],
"size": {
"0": 422.84503173828125,
"1": 164.31304931640625
},
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 3
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
29
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"masterpiece, best quality, 1girl, solo, cherry blossoms, hanami, pink flower, white flower, spring season, wisteria, petals, flower, plum blossoms, outdoors, falling petals, white hair, black eyes"
]
},
{
"id": 8,
"type": "VAEDecode",
"pos": [
1253,
191
],
"size": {
"0": 210,
"1": 46
},
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "samples",
"type": "LATENT",
"link": 28
},
{
"name": "vae",
"type": "VAE",
"link": 20
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
19
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VAEDecode"
}
},
{
"id": 12,
"type": "AnimateDiffCombine",
"pos": [
1254,
290
],
"size": {
"0": 315,
"1": 342
},
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 19
}
],
"properties": {
"Node name for S&R": "AnimateDiffCombine"
},
"widgets_values": [
8,
0,
"Enabled",
"AnimateDiff"
]
},
{
"id": 16,
"type": "AnimateDiffModuleLoader",
"pos": [
27,
345
],
"size": {
"0": 315,
"1": 58
},
"flags": {},
"order": 0,
"mode": 0,
"outputs": [
{
"name": "MOTION_MODULE",
"type": "MOTION_MODULE",
"links": [
24
],
"shape": 3
}
],
"properties": {
"Node name for S&R": "AnimateDiffModuleLoader"
},
"widgets_values": [
"mm-Stabilized_mid.pth"
]
},
{
"id": 4,
"type": "CheckpointLoaderSimple",
"pos": [
26,
474
],
"size": {
"0": 315,
"1": 98
},
"flags": {},
"order": 1,
"mode": 0,
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": [
25
],
"slot_index": 0
},
{
"name": "CLIP",
"type": "CLIP",
"links": [
3,
5
],
"slot_index": 1
},
{
"name": "VAE",
"type": "VAE",
"links": [],
"slot_index": 2
}
],
"properties": {
"Node name for S&R": "CheckpointLoaderSimple"
},
"widgets_values": [
"AnimeLike25D_v11.safetensors"
]
},
{
"id": 13,
"type": "VAELoader",
"pos": [
28,
223
],
"size": {
"0": 315,
"1": 58
},
"flags": {},
"order": 2,
"mode": 0,
"outputs": [
{
"name": "VAE",
"type": "VAE",
"links": [
20
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VAELoader"
},
"widgets_values": [
"klF8Anime2.ckpt"
]
},
{
"id": 15,
"type": "AnimateDiffSampler",
"pos": [
882,
192
],
"size": {
"0": 315,
"1": 330
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "motion_module",
"type": "MOTION_MODULE",
"link": 24,
"slot_index": 0
},
{
"name": "model",
"type": "MODEL",
"link": 25,
"slot_index": 1
},
{
"name": "positive",
"type": "CONDITIONING",
"link": 29
},
{
"name": "negative",
"type": "CONDITIONING",
"link": 30
},
{
"name": "latent_image",
"type": "LATENT",
"link": 35
}
],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
28
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "AnimateDiffSampler"
},
"widgets_values": [
"default",
16,
345029849956677,
"fixed",
20,
8,
"euler",
"normal",
0.8
]
},
{
"id": 7,
"type": "CLIPTextEncode",
"pos": [
413,
389
],
"size": {
"0": 425.27801513671875,
"1": 180.6060791015625
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 5
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
30
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"embedding:easynegative, embedding:badhandv4, "
]
},
{
"id": 20,
"type": "EmptyLatentImage",
"pos": [
522,
621
],
"size": {
"0": 315,
"1": 106
},
"flags": {},
"order": 3,
"mode": 0,
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
35
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "EmptyLatentImage"
},
"widgets_values": [
512,
512,
1
]
}
],
"links": [
[
3,
4,
1,
6,
0,
"CLIP"
],
[
5,
4,
1,
7,
0,
"CLIP"
],
[
19,
8,
0,
12,
0,
"IMAGE"
],
[
20,
13,
0,
8,
1,
"VAE"
],
[
24,
16,
0,
15,
0,
"MOTION_MODULE"
],
[
25,
4,
0,
15,
1,
"MODEL"
],
[
28,
15,
0,
8,
0,
"LATENT"
],
[
29,
6,
0,
15,
2,
"CONDITIONING"
],
[
30,
7,
0,
15,
3,
"CONDITIONING"
],
[
35,
20,
0,
15,
4,
"LATENT"
]
],
"groups": [],
"config": {},
"extra": {},
"version": 0.4
}