Progress on MotionCtrl support, fixed PIA nodes autosize not being stored in hidden

This commit is contained in:
Jedrzej Kosinski
2024-12-27 17:52:30 -06:00
parent ebb5496001
commit 1431497ff8
4 changed files with 215 additions and 1 deletions
+92
View File
@@ -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
+14
View File
@@ -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
+102
View File
@@ -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,)
+7 -1
View File
@@ -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}),
}
}