diff --git a/animatediff/logger.py b/animatediff/logger.py index 6aeccf4..2f9802d 100644 --- a/animatediff/logger.py +++ b/animatediff/logger.py @@ -28,10 +28,9 @@ 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") - ) + handler.setFormatter(ColoredFormatter("[%(name)s] - %(levelname)s - %(message)s")) logger.addHandler(handler) # Configure logger -logger.setLevel("INFO") +loglevel = logging.INFO +logger.setLevel(loglevel) diff --git a/animatediff/motion_module.py b/animatediff/motion_module.py index df32c67..0b73e61 100644 --- a/animatediff/motion_module.py +++ b/animatediff/motion_module.py @@ -17,14 +17,15 @@ def zero_module(module): class MotionWrapper(nn.Module): - def __init__(self): + def __init__(self, mm_type="mm_sd_v15.ckpt"): super().__init__() self.down_blocks = nn.ModuleList([]) self.up_blocks = nn.ModuleList([]) - for i, c in enumerate((320, 640, 1280, 1280)): + for c in (320, 640, 1280, 1280): self.down_blocks.append(MotionModule(c)) - for i, c in enumerate((1280, 1280, 640, 320)): + for c in (1280, 1280, 640, 320): self.up_blocks.append(MotionModule(c, is_up=True)) + self.mm_type = mm_type class MotionModule(nn.Module): @@ -75,18 +76,7 @@ class VanillaTemporalModule(nn.Module): ) 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 + return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask) class TemporalTransformer3DModel(nn.Module): @@ -142,11 +132,7 @@ class TemporalTransformer3DModel(nn.Module): 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") + video_length = hidden_states.shape[0] // 2 # TODO: config this value in scripts batch, channel, height, weight = hidden_states.shape residual = hidden_states @@ -175,7 +161,6 @@ class TemporalTransformer3DModel(nn.Module): ) output = hidden_states + residual - output = rearrange(output, "(b f) c h w -> b c f h w", f=video_length) return output @@ -342,4 +327,3 @@ class VersatileAttention(CrossAttention): hidden_states = rearrange(hidden_states, "(b d) f c -> (b f) d c", d=d) return hidden_states - diff --git a/animatediff/nodes.py b/animatediff/nodes.py index be9ed21..4c3480d 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -3,14 +3,16 @@ import json import hashlib import torch import numpy as np +from typing import Dict, List from PIL import Image from PIL.PngImagePlugin import PngInfo -from typing import Dict, List +from einops import rearrange 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.ldm.modules.diffusionmodules.util import GroupNorm32 from comfy.utils import load_torch_file from comfy.sd import ModelPatcher, calculate_parameters @@ -20,6 +22,7 @@ from .model_utils import MODEL_DIR, get_available_models orig_forward_timestep_embed = openaimodel.forward_timestep_embed +groupnorm32_original_forward = GroupNorm32.forward def forward_timestep_embed( @@ -95,10 +98,10 @@ class AnimateDiffLoader: model_path = os.path.join(MODEL_DIR, model_name) global motion_module - if motion_module is None: + if motion_module is None or motion_module.mm_type != model_name: logger.info(f"Loading motion module {model_name} from {model_path}") mm_state_dict = load_torch_file(model_path) - motion_module = MotionWrapper() + motion_module = MotionWrapper(model_name) parameters = calculate_parameters(mm_state_dict, "") usefp16 = model_management.should_use_fp16(model_params=parameters) @@ -109,6 +112,16 @@ class AnimateDiffLoader: motion_module = motion_module.to(offload_device) motion_module.load_state_dict(mm_state_dict) + logger.info(f"Hacking GroupNorm32 forward function.") + + 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 + + GroupNorm32.forward = groupnorm32_mm_forward + unet = model.model.diffusion_model if calculate_model_hash(unet) in injected_model_hashs: logger.info(f"Motion module already injected, skipping injection.") @@ -123,9 +136,9 @@ class AnimateDiffLoader: 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: + if unet_idx % 3 == 2 and unet_idx != 11: unet.output_blocks[unet_idx].insert( - -1, motion_module.up_blocks[mm_idx0].motion_modules[mm_idx] + -1, motion_module.up_blocks[mm_idx0].motion_modules[mm_idx1] ) else: unet.output_blocks[unet_idx].append( @@ -169,12 +182,15 @@ class AnimatedDiffUnload: logger.info(f"Unloading motion module from UNet output blocks.") for unet_idx in range(12): - if unet_idx % 2 == 2: + if unet_idx % 3 == 2 and unet_idx != 11: unet.output_blocks[unet_idx].pop(-2) else: unet.output_blocks[unet_idx].pop(-1) injected_model_hashs.remove(model_hash) + + logger.info(f"Restoring GroupNorm32 forward function.") + GroupNorm32.forward = groupnorm32_original_forward else: logger.info(f"Motion module not injected, skip unloading.")