Progress on MotionCtrl support, fixed PIA nodes autosize not being stored in hidden
This commit is contained in:
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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,)
|
||||
@@ -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}),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user