chore: update new changes from original repo

This commit is contained in:
Tung Nguyen
2023-07-27 04:33:41 +07:00
parent 8849369ce0
commit 19220fea1a
3 changed files with 31 additions and 32 deletions
+3 -4
View File
@@ -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)
+6 -22
View File
@@ -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
View File
@@ -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.")