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