From 93cafde180924deec5852051c60c7192fbbc58f1 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sun, 29 Dec 2024 14:42:54 -0600 Subject: [PATCH] Got MotionCtrl's camere control working, fixed frame_length affecting CameraCtrl Pose speed; still need to fix CameraCtrl Poses not adding to each other as expected --- animatediff/adapter_motionctrl.py | 11 +++++++ animatediff/model_injection.py | 47 ++++++++++++++++++++++++++++ animatediff/motion_module_ad.py | 22 ++++++++++--- animatediff/nodes.py | 3 +- animatediff/nodes_cameractrl.py | 34 ++++++++++++++++++--- animatediff/nodes_motionctrl.py | 51 +++++++++++++++++++++++++++++-- animatediff/sampling.py | 5 +-- 7 files changed, 159 insertions(+), 14 deletions(-) diff --git a/animatediff/adapter_motionctrl.py b/animatediff/adapter_motionctrl.py index fef7d1d..01bf922 100644 --- a/animatediff/adapter_motionctrl.py +++ b/animatediff/adapter_motionctrl.py @@ -1,6 +1,7 @@ # main code adapted from https://github.com/TencentARC/MotionCtrl/tree/animatediff from __future__ import annotations from torch import nn, Tensor +import torch from comfy.model_patcher import ModelPatcher import comfy.model_management @@ -83,6 +84,16 @@ def _remove_module_prefix(state_dict: dict[str, Tensor]): state_dict.pop(key) +def convert_cameractrl_poses_to_RT(poses: list[list[float]]): + tensors = [] + for pose in poses: + new_tensor = torch.tensor(pose[7:]) + new_tensor = new_tensor.unsqueeze(0) + tensors.append(new_tensor) + RT = torch.cat(tensors, dim=0) + return RT + + class ObjectControlModelPatcher(ModelPatcher): '''Class only used for type hints.''' def __init__(self): diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 954d7b5..0a4aca7 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -291,6 +291,12 @@ class MotionModelAttachment: self.prev_fancy_latents_shape: tuple = None self.fancy_multival: Union[float, Tensor] = None + # MotionCtrl + self.orig_RT: Tensor = None + self.RT: Tensor = None + self.prev_RT_shape: tuple = None + self.prev_RT_uuids: list = None + # temporary variables self.current_used_steps = 0 self.current_keyframe: ADKeyframe = None @@ -460,6 +466,43 @@ class MotionModelAttachment: self.prev_sub_idxs = sub_idxs self.prev_batched_number = batched_number + def prepare_motionctrl_camera(self, patcher: MotionModelPatcher, x: Tensor, transformer_options: dict[str]): + '''Used for MotionCtrl''' + # if no cc enabled, done + if not patcher.model.is_motionctrl_cc_enabled(): + if "ADE_RT" in transformer_options: + transformer_options.pop("ADE_RT") + return + cond_or_uncond: list[int] = transformer_options["cond_or_uncond"] + uuids: list = transformer_options["uuids"] + batched_number = len(cond_or_uncond) + ad_params = transformer_options["ad_params"] + full_length = ad_params["full_length"] + sub_idxs = ad_params["sub_idxs"] + goal_length = x.size(0) // batched_number + if self.prev_RT_shape != x.shape or sub_idxs != self.prev_sub_idxs or uuids != self.prev_RT_uuids: + real_RT = self.orig_RT.clone().to(dtype=x.dtype, device=x.device) # [t, 12] + # make sure RT is of the valid length + real_RT = extend_to_batch_size(real_RT, full_length) + if sub_idxs is not None: + real_RT = real_RT[sub_idxs] + real_RT = real_RT.unsqueeze(0) # [1, t, 12] + # match batch length - conds get real_RT, unconds get empty + if batched_number > 1: + batched_RTs = [] + for condtype in cond_or_uncond: + if condtype == 0: # cond + batched_RTs.append(real_RT) + else: # uncond + batched_RTs.append(torch.zeros_like(real_RT)) + real_RT = torch.cat(batched_RTs, dim=0) + self.RT = real_RT.to(dtype=x.dtype, device=x.device) + self.prev_RT_shape = x.shape + transformer_options["ADE_RT"] = self.RT + self.prev_sub_idxs = sub_idxs + self.prev_batched_number = batched_number + + def get_pia_c_concat(self, model: BaseModel, x: Tensor) -> Tensor: '''Used for PIA''' # if have cached shape, check if matches - if so, return cached pia_latents @@ -583,6 +626,10 @@ class MotionModelAttachment: # PIA self.combined_pia_mask = None self.combined_pia_effect = None + # MotionCtrl + self.RT = None + self.prev_RT_shape = None + self.prev_RT_uuids = None # Default self.current_used_steps = 0 self.current_keyframe = None diff --git a/animatediff/motion_module_ad.py b/animatediff/motion_module_ad.py index 0fefe1d..d6af306 100644 --- a/animatediff/motion_module_ad.py +++ b/animatediff/motion_module_ad.py @@ -1,3 +1,4 @@ +from __future__ import annotations import math from typing import Iterable, Tuple, Union, TYPE_CHECKING import re @@ -366,7 +367,7 @@ class AnimateDiffModel(nn.Module): if has_img_encoder(mm_state_dict): self.init_img_encoder() # CameraCtrl stuff - self.camera_encoder: 'CameraPoseEncoder' = None + self.camera_encoder: CameraPoseEncoder = None # PIA/FancyVideo stuff - create conv_in if keys are present for it self.conv_in: comfy.ops.disable_weight_init.Conv2d = None self.orig_conv_in: comfy.ops.disable_weight_init.Conv2d = None @@ -382,11 +383,15 @@ class AnimateDiffModel(nn.Module): # get_unet_func initialization self.get_unet_func = init_kwargs.get(InitKwargs.GET_UNET_FUNC, get_unet_default) + def needs_apply_model_wrapper(self): + '''Returns true of AnimateLCM-I2V, CameraCtrl, or MotionCtrl is in use.''' + return self.img_encoder is not None or self.camera_encoder is not None or self.is_motionctrl_cc_enabled() + def init_img_encoder(self): del self.img_encoder self.img_encoder = AdapterEmbed(cin=4, channels=self.layer_channels, nums_rb=2, ksize=1, sk=True, use_conv=False, ops=self.ops) - def set_camera_encoder(self, camera_encoder: 'CameraPoseEncoder'): + def set_camera_encoder(self, camera_encoder: CameraPoseEncoder): del self.camera_encoder self.camera_encoder = camera_encoder @@ -428,6 +433,13 @@ class AnimateDiffModel(nn.Module): ttb: TemporalTransformerBlock = comfy.utils.get_attr(self, ttb_key) ttb.init_cc_projection(in_features=in_features, out_features=out_features, ops=self.ops) + def is_motionctrl_cc_enabled(self): + '''Used for MotionCtrl''' + if self.down_blocks: + ttb: TemporalTransformerBlock = self.down_blocks[0].motion_modules[0].temporal_transformer.transformer_blocks[0] + return ttb.cc_projection is not None + return False + def get_fancyvideo_emb_patches(self, dtype, device, fps=25, motion_score=3.0): patches = [] if self.fps_embedding is not None: @@ -1320,12 +1332,12 @@ class TemporalTransformerBlock(nn.Module): ) # do MotionCtrl-CMCM stuff if needed if self.cc_projection is not None and count==0 and 'ADE_RT' in transformer_options: - RT: Tensor = transformer_options['ADE_RT'] + RT: Tensor = transformer_options['ADE_RT'].to(dtype=hidden_states.dtype) B, t, _ = RT.shape RT = RT.reshape(B*t, 1, -1) - RT = RT.repeat(1, hidden_states.shape[1]) + RT = RT.repeat(1, hidden_states.shape[1], 1) hidden_states = torch.cat([hidden_states, RT], dim=-1) - hidden_states = self.cc_projection(hidden_states) + hidden_states = self.cc_projection(hidden_states).to(dtype=hidden_states.dtype) count += 1 else: # views idea gotten from diffusers AnimateDiff FreeNoise implementation: diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 64befdc..f83b32a 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -8,7 +8,7 @@ from .nodes_cameractrl import (LoadAnimateDiffModelWithCameraCtrl, ApplyAnimateD LoadCameraPosesFromFile, LoadCameraPosesFromPath, CameraCtrlPoseBasic, CameraCtrlPoseCombo, CameraCtrlPoseAdvanced, CameraCtrlManualAppendPose, CameraCtrlReplaceCameraParameters, CameraCtrlSetOriginalAspectRatio) -from .nodes_motionctrl import (LoadMotionCtrlCMCM, LoadMotionCtrlOMCM, ApplyAnimateDiffMotionCtrlModel) +from .nodes_motionctrl import (LoadMotionCtrlCMCM, LoadMotionCtrlOMCM, ApplyAnimateDiffMotionCtrlModel, LoadMotionCtrlCameraPosesFromFile) from .nodes_pia import (ApplyAnimateDiffPIAModel, LoadAnimateDiffAndInjectPIANode, InputPIA_MultivalNode, InputPIA_PaperPresetsNode, PIA_ADKeyframeNode) from .nodes_fancyvideo import (ApplyAnimateDiffFancyVideo,) from .nodes_hellomeme import (TestHMRefNetInjection,) @@ -197,6 +197,7 @@ NODE_CLASS_MAPPINGS = { LoadMotionCtrlCMCM.NodeID: LoadMotionCtrlCMCM, LoadMotionCtrlOMCM.NodeID: LoadMotionCtrlOMCM, ApplyAnimateDiffMotionCtrlModel.NodeID: ApplyAnimateDiffMotionCtrlModel, + LoadMotionCtrlCameraPosesFromFile.NodeID: LoadMotionCtrlCameraPosesFromFile, # CameraCtrl Nodes "ADE_ApplyAnimateDiffModelWithCameraCtrl": ApplyAnimateDiffWithCameraCtrl, "ADE_LoadAnimateDiffModelWithCameraCtrl": LoadAnimateDiffModelWithCameraCtrl, diff --git a/animatediff/nodes_cameractrl.py b/animatediff/nodes_cameractrl.py index 7b7bb9a..b91c033 100644 --- a/animatediff/nodes_cameractrl.py +++ b/animatediff/nodes_cameractrl.py @@ -115,13 +115,13 @@ def compute_R_from_rad_angle(angles: np.ndarray): R = np.dot(Rz, np.dot(Ry, Rx)) return R -def get_camera_motion(angle: np.ndarray, T: np.ndarray, speed: float, n=16): +def get_camera_motion(angle: np.ndarray, T: np.ndarray, speed: float, n=16, base=16): RT = [] for i in range(n): - _angle = (i/n)*speed*(CAM.BASE_ANGLE)*angle + _angle = (i/base)*speed*(CAM.BASE_ANGLE)*angle R = compute_R_from_rad_angle(_angle) # _T = (i/n)*speed*(T.reshape(3,1)) - _T=(i/n)*speed*(CAM.BASE_T_NORM)*(T.reshape(3,1)) + _T=(i/base)*speed*(CAM.BASE_T_NORM)*(T.reshape(3,1)) _RT = np.concatenate([R,_T], axis=1) RT.append(_RT) RT = np.stack(RT) @@ -143,6 +143,22 @@ def combine_RTs(RT_0: np.ndarray, RT_1: np.ndarray): return np.concatenate([RT_0, RT_1], axis=0) +def stack_RTs(RT_0: np.ndarray, RT_1: np.ndarray): + RT_target = copy.deepcopy(RT_1) + static_motion = CAM.get(CAM.STATIC) + RT_static = get_camera_motion(static_motion.rotate, static_motion.translate, 1.0, 1) + RT_offset = RT_0[-1] - RT_static[-1] + + temp = [] + for sub_RT in RT_target: + temp.append(sub_RT + RT_offset) + + RT_1 = np.stack(temp) + RT_0 = RT_0[:-1] + + return np.concatenate([RT_0, RT_1], axis=0) + + def set_original_pose_dims(poses: list[list[float]], pose_width, pose_height): # indexes 5 and 6 are not used for anything in the poses, so can use 5 and 6 to set original pose width/height new_poses = copy.deepcopy(poses) @@ -157,7 +173,17 @@ def combine_poses(poses0: list[list[float]], poses1: list[list[float]]): inter_poses = ndarray_to_poses(new_RT) # maintain fx, fy, cx, and cy values by pasting only the movement portion of poses for i in range(len(new_poses)): - new_poses[7:] = inter_poses[7:] + new_poses[i][7:] = inter_poses[i][7:] + return new_poses + + +def combine_poses_redux(poses0: list[list[float]], poses1: list[list[float]]): + new_poses = copy.deepcopy(poses0[:-1]) + copy.deepcopy(poses1) + new_RT = stack_RTs(poses_to_ndarray(poses0), poses_to_ndarray(poses1)) + inter_poses = ndarray_to_poses(new_RT) + # maintain fx, fy, cx, and cy values by pasting only the movement portion of poses + for i in range(len(new_poses)): + new_poses[i][7:] = inter_poses[i][7:] return new_poses diff --git a/animatediff/nodes_motionctrl.py b/animatediff/nodes_motionctrl.py index 3d02e1f..81c057f 100644 --- a/animatediff/nodes_motionctrl.py +++ b/animatediff/nodes_motionctrl.py @@ -1,10 +1,16 @@ import torch from torch import Tensor +import numpy as np +import os +import json + +import folder_paths from .ad_settings import AnimateDiffSettings -from .adapter_motionctrl import inject_motionctrl_cmcm, load_motionctrl_omcm +from .adapter_motionctrl import (ObjectControlModelPatcher, inject_motionctrl_cmcm, load_motionctrl_omcm, + convert_cameractrl_poses_to_RT) -from .model_injection import MotionModelPatcher, MotionModelGroup, load_motion_module_gen2 +from .model_injection import MotionModelPatcher, MotionModelGroup, load_motion_module_gen2, get_mm_attachment from .motion_lora import MotionLoraList from .nodes_gen2 import ApplyAnimateDiffModelNode @@ -60,6 +66,33 @@ class LoadMotionCtrlOMCM: return (omcm_modelpatcher,) +class LoadMotionCtrlCameraPosesFromFile: + NodeID = "ADE_LoadMotionCtrlCameraPosesFromFile" + NodeName = "Load MotionCtrl Camera Poses 🎭🅐🅓" + @classmethod + def INPUT_TYPES(s): + input_dir = folder_paths.get_input_directory() + files = [f for f in os.listdir(input_dir) if os.path.isfile(os.path.join(input_dir, f))] + files = [f for f in files if f.endswith(".json")] + return { + "required": { + "pose_filename": (sorted(files),), + } + } + + RETURN_TYPES = ("CAMERA_MOTIONCTRL",) + CATEGORY = "Animate Diff 🎭🅐🅓/② Gen2 nodes ②/MotionCtrl" + FUNCTION = "load_camera_poses" + + def load_camera_poses(self, pose_filename): + file_path = folder_paths.get_annotated_filepath(pose_filename) + with open(file_path, 'r') as f: + RT = json.load(f) + RT = np.array(RT) + RT = torch.tensor(RT).float() # [t, 12] + return (RT,) + + class ApplyAnimateDiffMotionCtrlModel: NodeID = "ADE_ApplyAnimateDiffModelWithMotionCtrl" NodeName = "Apply AnimateDiff+MotionCtrl Model 🎭🅐🅓②" @@ -73,6 +106,7 @@ class ApplyAnimateDiffMotionCtrlModel: }, "optional": { "omcm_motionctrl": ("OMCM_MOTIONCTRL",), + "cameractrl_poses": ("CAMERACTRL_POSES",), "motion_lora": ("MOTION_LORA",), "scale_multival": ("MULTIVAL",), "effect_multival": ("MULTIVAL",), @@ -90,6 +124,8 @@ class ApplyAnimateDiffMotionCtrlModel: FUNCTION = "apply_motion_model" def apply_motion_model(self, motion_model: MotionModelPatcher, start_percent: float=0.0, end_percent: float=1.0, + omcm_motionctrl: ObjectControlModelPatcher=None, + cameractrl_poses: list[list[float]]=None, motion_lora: MotionLoraList=None, ad_keyframes: ADKeyframeGroup=None, scale_multival=None, effect_multival=None, per_block: AllPerBlocks=None, prev_m_models: MotionModelGroup=None,): @@ -99,5 +135,16 @@ class ApplyAnimateDiffMotionCtrlModel: # 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 + if curr_model.model.is_motionctrl_cc_enabled(): + attachment = get_mm_attachment(curr_model) + if cameractrl_poses is not None: + RT = convert_cameractrl_poses_to_RT(cameractrl_poses) + attachment.orig_RT = RT + else: + attachment.orig_RT = torch.zeros((1, 12)) + # attachment.orig_RT = cameractrl_poses + # else: + # attachment.orig_RT = torch.zeros([]) + # check if OMCM is provided; if so, make sure something is provided for it return (new_m_models,) diff --git a/animatediff/sampling.py b/animatediff/sampling.py index f0fdf84..a48af04 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -193,6 +193,7 @@ def _apply_model_wrapper(executor, *args, **kwargs): attachment = get_mm_attachment(motion_model) attachment.prepare_alcmi2v_features(motion_model, x=x, cond_or_uncond=cond_or_uncond, ad_params=ad_params, latent_format=executor.class_obj.latent_format) attachment.prepare_camera_features(motion_model, x=x, cond_or_uncond=cond_or_uncond, ad_params=ad_params) + attachment.prepare_motionctrl_camera(motion_model, x=x, transformer_options=transformer_options) del x return executor(*args, **kwargs) @@ -285,9 +286,9 @@ class FunctionInjectionHolder: helper.model.model.memory_required = unlimited_memory_required except Exception: pass - # if img_encoder or camera_encoder present, inject apply_model to handle correctly + # if AnimateLCM-I2V, CameraCtrl, or MotionCtrl present, inject apply_model to handle correctly for motion_model in helper.get_motion_models(): - if (motion_model.model.img_encoder is not None) or (motion_model.model.camera_encoder is not None): + if motion_model.model.needs_apply_model_wrapper(): create_special_model_apply_model_wrapper(model_options) break del info