chore: update new changes from original repo
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
+22
-6
@@ -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.")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user