diff --git a/animatediff/adapter_motionctrl.py b/animatediff/adapter_motionctrl.py new file mode 100644 index 0000000..6599f0f --- /dev/null +++ b/animatediff/adapter_motionctrl.py @@ -0,0 +1,92 @@ +# main code adapted from https://github.com/TencentARC/MotionCtrl/tree/animatediff +from __future__ import annotations +from torch import nn, Tensor + +from comfy.model_patcher import ModelPatcher +import comfy.model_management +import comfy.ops +import comfy.utils + +from .adapter_cameractrl import ResnetBlockCameraCtrl +from .motion_module_ad import AnimateDiffModel +from .utils_model import get_motion_model_path + +# cmcm (Camera Control) +def injection_motionctrl_cmcm(motion_model: AnimateDiffModel, cmcm_name: str): + pass + + +# omcm (Object Control) +def load_motionctrl_omcm(omcm_name: str): + omcm_path = get_motion_model_path(omcm_name) + state_dict = comfy.utils.load_torch_file(omcm_path, safe_load=True) + for key in list(state_dict.keys()): + # remove 'module.' prefix + if key.startswith('module.'): + new_key = key.replace('module.', '') + state_dict[new_key] = state_dict[key] + state_dict.pop(key) + + if comfy.model_management.unet_manual_cast(comfy.model_management.unet_dtype(), comfy.model_management.get_torch_device()) is None: + ops = comfy.ops.disable_weight_init + else: + ops = comfy.ops.manual_cast + adapter = MotionCtrlAdapter(ops=ops) + adapter.load_state_dict(state_dict=state_dict, strict=True) + adapter.to( + device = comfy.model_management.unet_offload_device(), + dtype = comfy.model_management.unet_dtype() + ) + omcm_modelpatcher = _create_OMCMModelPatcher(model=adapter, + load_device=comfy.model_management.get_torch_device(), + offload_device=comfy.model_management.unet_offload_device()) + return omcm_modelpatcher + + +def _create_OMCMModelPatcher(model, load_device, offload_device) -> ObjectControlModelPatcher: + patcher = ModelPatcher(model, load_device=load_device, offload_device=offload_device) + return patcher + + +class ObjectControlModelPatcher(ModelPatcher): + '''Class only used for type hints.''' + def __init__(self): + self.model: MotionCtrlAdapter + + +class MotionCtrlAdapter(nn.Module): + def __init__(self, + downscale_factor=8, + channels=[320, 640, 1280, 1280], + nums_rb=2, cin=128, # 2*8*8 + ksize=3, sk=True, + use_conv=False, + ops=comfy.ops.disable_weight_init): + super(MotionCtrlAdapter, self).__init__() + self.downscale_factor = downscale_factor + self.unshuffle = nn.PixelUnshuffle(downscale_factor) + self.channels = channels + self.nums_rb = nums_rb + self.body = [] + for i in range(len(channels)): + for j in range(nums_rb): + if (i != 0) and (j == 0): + self.body.append( + ResnetBlockCameraCtrl(channels[i - 1], channels[i], down=True, ksize=ksize, sk=sk, use_conv=use_conv, ops=ops)) + else: + self.body.append( + ResnetBlockCameraCtrl(channels[i], channels[i], down=False, ksize=ksize, sk=sk, use_conv=use_conv, ops=ops)) + self.body = nn.ModuleList(self.body) + self.conv_in = ops.Conv2d(cin, channels[0], 3, 1, 1) + + def forward(self, x: Tensor): + x = self.unshuffle(x) + # extract features + features = [] + x = self.conv_in(x) + for i in range(len(self.channels)): + for j in range(self.nums_rb): + idx = i * self.nums_rb + j + x = self.body[idx](x) + features.append(x) + return features diff --git a/animatediff/motion_module_ad.py b/animatediff/motion_module_ad.py index 21c265b..f14bae5 100644 --- a/animatediff/motion_module_ad.py +++ b/animatediff/motion_module_ad.py @@ -1313,6 +1313,8 @@ class TemporalTransformerBlock(nn.Module): self.ff = FeedForward(dim, dropout=dropout, glu=(activation_fn == "geglu"), operations=ops) self.ff_norm = ops.LayerNorm(dim) + # for MotionCtrl (CMCM) use + self.cc_projections: comfy.ops.disable_weight_init.Linear = None def set_scale_multiplier(self, idx: int, multiplier: Union[float, None]): self.attention_blocks[idx].set_scale_multiplier(multiplier) @@ -1346,6 +1348,7 @@ class TemporalTransformerBlock(nn.Module): elif view_options.context_length == video_length and not view_options.use_on_equal_length: view_options = None if not view_options: + count = 0 for attention_block, norm, scale_mask in zip(self.attention_blocks, self.norms, scale_masks): norm_hidden_states = norm(hidden_states).to(hidden_states.dtype) hidden_states = ( @@ -1362,6 +1365,15 @@ class TemporalTransformerBlock(nn.Module): transformer_options=transformer_options, ) + hidden_states ) + # do MotionCtrl-CMCM stuff if needed + if self.cc_projections is not None and count==0 and 'ADE_RT' in transformer_options: + RT: Tensor = transformer_options['ADE_RT'] + B, t, _ = RT.shape + RT = RT.reshape(B*t, 1, -1) + RT = RT.repeat(1, hidden_states.shape[1]) + hidden_states = torch.cat([hidden_states, RT], dim=-1) + hidden_states = self.cc_projections(hidden_states) + count += 1 else: # views idea gotten from diffusers AnimateDiff FreeNoise implementation: # https://github.com/arthur-qiu/FreeNoise-AnimateDiff/blob/main/animatediff/models/motion_module.py @@ -1381,6 +1393,7 @@ class TemporalTransformerBlock(nn.Module): sub_hidden_states = rearrange(hidden_states[:, sub_idxs], "b f d c -> (b f) d c") if has_camera_feature: mm_kwargs["camera_feature"] = orig_camera_feature[:, sub_idxs, :] + count = 0 for attention_block, norm, scale_mask in zip(self.attention_blocks, self.norms, scale_masks): norm_hidden_states = norm(sub_hidden_states).to(sub_hidden_states.dtype) sub_hidden_states = ( @@ -1397,6 +1410,7 @@ class TemporalTransformerBlock(nn.Module): transformer_options=transformer_options, ) + sub_hidden_states ) + count += 1 sub_hidden_states = rearrange(sub_hidden_states, "(b f) d c -> b f d c", f=len(sub_idxs)) weights = get_context_weights(len(sub_idxs), view_options.fuse_method) * batched_conds diff --git a/animatediff/nodes_motionctrl.py b/animatediff/nodes_motionctrl.py new file mode 100644 index 0000000..7c84753 --- /dev/null +++ b/animatediff/nodes_motionctrl.py @@ -0,0 +1,102 @@ +import torch +from torch import Tensor + +from .ad_settings import AnimateDiffSettings +from .adapter_motionctrl import injection_motionctrl_cmcm, load_motionctrl_omcm + +from .motion_module_ad import AllPerBlocks +from .model_injection import MotionModelPatcher, MotionModelGroup, load_motion_module_gen2 +from .motion_lora import MotionLoraList + +from .nodes_gen2 import ApplyAnimateDiffModelNode +from .utils_model import get_available_motion_models +from .utils_motion import ADKeyframeGroup + + +class LoadMotionCtrlCMCM: + NodeID = "ADE_LoadMotionCtrl_CMCMMOdel" + NodeName = "Load AnimateDiff+MotionCtrl Camera Model 🎭🅐🅓②" + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model_name": (get_available_motion_models(),), + "motionctrl_cmcm": (get_available_motion_models(),), + }, + "optional": { + "ad_settings": ("AD_SETTINGS",), + } + } + + RETURN_TYPES = ("MOTION_MODEL_ADE",) + RETURN_NAMES = ("MOTION_MODEL",) + CATEGORY = "Animate Diff 🎭🅐🅓/② Gen2 nodes ②/MotionCtrl" + FUNCTION = "load_motionctrl_cmcm" + + def load_motionctrl_cmcm(self, model_name: str, motionctrl_cmcm: str, ad_settings: AnimateDiffSettings=None): + motion_model = load_motion_module_gen2(model_name=model_name, motion_model_settings=ad_settings) + motion_model = injection_motionctrl_cmcm(motion_model, cmcm_name=motionctrl_cmcm) + return (motion_model,) + + +class LoadMotionCtrlOMCM: + NodeID = "ADE_LoadMotionCtrl_OMCMMOdel" + NodeName = "Load MotionCtrl Object Model 🎭🅐🅓②" + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "motionctrl_omcm": (get_available_motion_models(),), + } + } + + RETURN_TYPES = ("OMCM_MOTIONCTRL",) + CATEGORY = "Animate Diff 🎭🅐🅓/② Gen2 nodes ②/MotionCtrl" + FUNCTION = "load_motionctrl_omcm" + + def load_motionctrl_omcm(self, motionctrl_omcm: str): + omcm_modelpatcher = load_motionctrl_omcm(motionctrl_omcm) + return (omcm_modelpatcher,) + + +class ApplyAnimateDiffMotionCtrlModel: + NodeID = "ADE_ApplyAnimateDiffModelWithMotionCtrl" + NodeName = "Apply AnimateDiff+MotionCtrl Model 🎭🅐🅓②" + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "motion_model": ("MOTION_MODEL_ADE",), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), + "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}), + }, + "optional": { + "omcm_motionctrl": ("OMCM_MOTIONCTRL",), + "motion_lora": ("MOTION_LORA",), + "scale_multival": ("MULTIVAL",), + "effect_multival": ("MULTIVAL",), + "ad_keyframes": ("AD_KEYFRAMES",), + "prev_m_models": ("M_MODELS",), + "per_block": ("PER_BLOCK",), + }, + "hidden": { + "autosize": ("ADEAUTOSIZE", {"padding": 0}), + } + } + + RETURN_TYPES = ("M_MODELS",) + CATEGORY = "Animate Diff 🎭🅐🅓/② Gen2 nodes ②/MotionCtrl" + FUNCTION = "apply_motion_model" + + def apply_motion_model(self, motion_model: MotionModelPatcher, start_percent: float=0.0, end_percent: float=1.0, + motion_lora: MotionLoraList=None, ad_keyframes: ADKeyframeGroup=None, + scale_multival=None, effect_multival=None, per_block: AllPerBlocks=None, + prev_m_models: MotionModelGroup=None,): + (new_m_models,) = ApplyAnimateDiffModelNode.apply_motion_model(self, motion_model, start_percent=start_percent, end_percent=end_percent, + motion_lora=motion_lora, ad_keyframes=ad_keyframes, per_block=per_block, + scale_multival=scale_multival, effect_multival=effect_multival, prev_m_models=prev_m_models) + # most recent added model will always be first in list + curr_model = new_m_models.models[0] + # check if model has CMCM; if so, make sure something is provided for it + # check if OMCM is provided; if so, make sure something is provided for it + return (new_m_models,) diff --git a/animatediff/nodes_pia.py b/animatediff/nodes_pia.py index a7b9716..a61553a 100644 --- a/animatediff/nodes_pia.py +++ b/animatediff/nodes_pia.py @@ -126,6 +126,8 @@ class ApplyAnimateDiffPIAModel: "ad_keyframes": ("AD_KEYFRAMES",), "prev_m_models": ("M_MODELS",), "per_block": ("PER_BLOCK",), + }, + "hidden": { "autosize": ("ADEAUTOSIZE", {"padding": 0}), } } @@ -202,6 +204,8 @@ class PIA_ADKeyframeNode: "pia_input": ("PIA_INPUT",), "inherit_missing": ("BOOLEAN", {"default": True}, ), "guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}), + }, + "hidden": { "autosize": ("ADEAUTOSIZE", {"padding": 0}), } } @@ -254,8 +258,10 @@ class InputPIA_PaperPresetsNode: "optional": { "mult_multival": ("MULTIVAL",), "print_values": ("BOOLEAN", {"default": False},), - "autosize": ("ADEAUTOSIZE", {"padding": 0}), #"effect_multival": ("MULTIVAL",), + }, + "hidden": { + "autosize": ("ADEAUTOSIZE", {"padding": 0}), } }