From 7451bb58fa306f2f230b516ec7c6569ab23d259c Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 4 Apr 2024 14:51:49 -0500 Subject: [PATCH] Progress on CameraCtrl support --- ...v_adapter.py => adapter_animatelcm_i2v.py} | 0 animatediff/adapter_cameractrl.py | 195 ++++++++++++++++++ animatediff/motion_module_ad.py | 2 +- animatediff/nodes.py | 7 +- animatediff/nodes_gen2.py | 29 ++- 5 files changed, 230 insertions(+), 3 deletions(-) rename animatediff/{animatelcm_i2v_adapter.py => adapter_animatelcm_i2v.py} (100%) create mode 100644 animatediff/adapter_cameractrl.py diff --git a/animatediff/animatelcm_i2v_adapter.py b/animatediff/adapter_animatelcm_i2v.py similarity index 100% rename from animatediff/animatelcm_i2v_adapter.py rename to animatediff/adapter_animatelcm_i2v.py diff --git a/animatediff/adapter_cameractrl.py b/animatediff/adapter_cameractrl.py new file mode 100644 index 0000000..3829516 --- /dev/null +++ b/animatediff/adapter_cameractrl.py @@ -0,0 +1,195 @@ + +# Modified from https://github.com/hehao13/CameraCtrl/blob/main/cameractrl/models/pose_adaptor.py +# (whose parts were also taken from https://github.com/TencentARC/T2I-Adapter) +import torch +import torch.nn as nn +from torch import Tensor +from collections import OrderedDict +from einops import rearrange + +import comfy.ops + +from .motion_module_ad import TemporalTransformerBlock + + +def conv_nd(dims, *args, **kwargs): + """ + Create a 1D, 2D, or 3D convolution module. + """ + if dims == 1: + return nn.Conv1d(*args, **kwargs) + elif dims == 2: + return nn.Conv2d(*args, **kwargs) + elif dims == 3: + return nn.Conv3d(*args, **kwargs) + raise ValueError(f"unsupported dimensions: {dims}") + + +def avg_pool_nd(dims, *args, **kwargs): + """ + Create a 1D, 2D, or 3D average pooling module. + """ + if dims == 1: + return nn.AvgPool1d(*args, **kwargs) + elif dims == 2: + return nn.AvgPool2d(*args, **kwargs) + elif dims == 3: + return nn.AvgPool3d(*args, **kwargs) + raise ValueError(f"unsupported dimensions: {dims}") + + +# class PoseAdapter(nn.Module): +# def __init__(self, unet, pose_encoder): +# super().__init__() +# self.unet = unet +# self.pose_encoder = pose_encoder + +# def forward(self, noisy_latents: Tensor, timesteps, encoder_hidden_states: Tensor, pose_embedding): +# # original code needed to convert from 4 dims (bf c h w) to 5 dims (b c f h w), +# # but ComfyUI already deals with everything in 4 dims +# pose_embedding_features = self.pose_encoder(pose_embedding) +# noise_pred = self.unet(noisy_latents, +# timesteps, +# encoder_hidden_states, +# pose_embedding_features).sample +# return noise_pred + + +class CameraPoseEncoder(nn.Module): + def __init__(self, + downscale_factor=8, + channels=[320, 640, 1280, 1280], + nums_rb=3, + cin=64, + ksize=3, + sk=False, + use_conv=True, + compression_factor=1, + temporal_attention_nhead=8, + attention_block_types=("Temporal_Self", ), + temporal_position_encoding=False, + temporal_position_encoding_max_len=16, + rescale_output_factor=1.0, + ops=comfy.ops.disable_weight_init): + super(CameraPoseEncoder, self).__init__() + self.unshuffle = nn.PixelUnshuffle(downscale_factor) + self.channels = channels + self.nums_rb = nums_rb + self.encoder_conv_in = ops.Conv2d(cin, channels[0], 3, 1, 1) + self.encoder_down_conv_blocks = nn.ModuleList() + self.encoder_down_attention_blocks = nn.ModuleList() + for i in range(len(channels)): + conv_layers = nn.ModuleList() + temporal_attention_layers = nn.ModuleList() + for j in range(len(nums_rb)): + if j == 0 and i != 0: + in_dim = channels[i - 1] + out_dim = int(channels[i] / compression_factor) + conv_layer = ResnetBlock(in_dim, out_dim, down=True, ksize=ksize, sk=sk, use_conv=use_conv, ops=ops) + elif j == 0: + in_dim = channels[0] + out_dim = int(channels[i] / compression_factor) + conv_layer = ResnetBlock(in_dim, out_dim, down=False, ksize=ksize, sk=sk, use_conv=use_conv, ops=ops) + elif j == nums_rb - 1: + in_dim = channels[i] / compression_factor + out_dim = channels[i] + conv_layer = ResnetBlock(in_dim, out_dim, down=False, ksize=ksize, sk=sk, use_conv=use_conv, ops=ops) + else: + in_dim = int(channels[i] / compression_factor) + out_dim = int(channels[i] / compression_factor) + conv_layer = ResnetBlock(in_dim, out_dim, down=False, ksize=ksize, sk=sk, use_conv=use_conv, ops=ops) + temporal_attention_layer = TemporalTransformerBlock(dim=out_dim, + num_attention_heads=temporal_attention_nhead, + attention_head_dim=int(out_dim / temporal_attention_nhead), + attention_block_types=attention_block_types, + dropout=0.0, + cross_attention_dim=None, + temporal_position_encoding=temporal_position_encoding, + temporal_position_encoding_max_len=temporal_position_encoding_max_len, + rescale_output_factor=rescale_output_factor, + ops=ops) + conv_layers.append(conv_layer) + temporal_attention_layers.append(temporal_attention_layer) + self.encoder_down_conv_blocks.append(conv_layers) + self.encoder_down_attention_blocks.append(temporal_attention_layers) + + def forward(self, x: Tensor): + # unshuffle + x = self.unshuffle + # extract features + features = [] + x = self.encoder_conv_in(x) + for res_block, attention_block in zip(self.encoder_down_conv_blocks, self.encoder_down_attention_blocks): + for res_layer, attention_layer in zip(res_block, attention_block): + x = res_layer(x) + h, w = x.shape[-2:] + x = rearrange(x, 'b c h w -> (h w) b c') + x = attention_layer(x) + x = rearrange(x, '(h w) b c -> b c h w', h=h, w=w) + features.append(x) + return features + + +class ResnetBlock(nn.Module): + def __init__(self, in_c, out_c, down: bool, ksize=3, sk=False, use_conv=True, + ops=comfy.ops.disable_weight_init): + super().__init__() + ps = ksize // 2 # padding size + if in_c != out_c or sk == False: + self.in_conv = ops.Conv2d(in_c, out_c, ksize, 1, ps) + else: + self.in_conv = None + self.block1 = ops.Conv2d(out_c, out_c, 3, 1, 1) + self.act = nn.ReLU() + self.block2 = ops.Conv2d(out_c, out_c, ksize, 1, ps) + if sk == False: + self.skep = ops.Conv2d(in_c, out_c, ksize, 1, ps) + else: + self.skep = None + + self.down = down + if self.down == True: + self.down_opt = Downsample(in_c, use_conv=use_conv) + + def forward(self, x: Tensor): + if self.down == True: + x = self.down_opt(x) + if self.in_conv is not None: # edit + x = self.in_conv(x) + + h = self.block1(x) + h = self.act(h) + h = self.block2(h) + if self.skep is not None: + return h + self.skep(x) + else: + return h + x + + +class Downsample(nn.Module): + """ + A downsampling layer with an optional convolution. + :param channels: channels in the inputs and outputs. + :param use_conv: a bool determining if a convolution is applied. + :param dims: determines if the signal is 1D, 2D, or 3D. If 3D, then + downsampling occurs in the inner-two dimensions. + """ + + def __init__(self, channels, use_conv: bool, dims=2, out_channels=None, padding=1, + ops=comfy.ops.disable_weight_init): + super().__init__() + self.channels = channels + self.out_channels = out_channels or channels + self.use_conv = use_conv + self.dims = dims + stride = 2 if dims != 3 else (1, 2, 2) + if use_conv: + self.operation = ops.conv_nd(dims, in_channels=self.channels, out_channels=self.out_channels, + kernel_size=3, stride=stride, padding=padding) + else: + assert self.channels == self.out_channels + self.operation = avg_pool_nd(dims, kernel_size=stride, stride=stride) # both are stride value on purpose + + def forward(self, x: Tensor): + assert x.shape[1] == self.channels + return self.operation(x) diff --git a/animatediff/motion_module_ad.py b/animatediff/motion_module_ad.py index dbf002c..1dd04e7 100644 --- a/animatediff/motion_module_ad.py +++ b/animatediff/motion_module_ad.py @@ -16,7 +16,7 @@ import comfy.ops import comfy.model_management from .context import ContextFuseMethod, ContextOptions, get_context_weights, get_context_windows -from .animatelcm_i2v_adapter import AdapterEmbed +from .adapter_animatelcm_i2v import AdapterEmbed from .utils_motion import CrossAttentionMM, MotionCompatibilityError, DummyNNModule, extend_to_batch_size, prepare_mask_batch from .utils_model import BetaSchedules, ModelTypeSD from .logger import logger diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 09535f5..da7df61 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -4,7 +4,8 @@ from .sampling import motion_sample_factory from .nodes_gen1 import (AnimateDiffLoaderGen1, LegacyAnimateDiffLoaderWithContext) from .nodes_gen2 import (UseEvolvedSamplingNode, ApplyAnimateDiffModelNode, ApplyAnimateDiffModelBasicNode, ApplyAnimateLCMI2VModel, ADKeyframeNode, - LoadAnimateDiffModelNode, LoadAnimateLCMI2VModelNode, LoadAnimateDiffAndInjectI2VNode, UpscaleAndVaeEncode) + LoadAnimateDiffModelNode, LoadAnimateLCMI2VModelNode, LoadAnimateDiffAndInjectI2VNode, UpscaleAndVaeEncode, + LoadCameraCtrlAdapter) from .nodes_multival import MultivalDynamicNode, MultivalScaledMaskNode from .nodes_sample import (FreeInitOptionsNode, NoiseLayerAddWeightedNode, SampleSettingsNode, NoiseLayerAddNode, NoiseLayerReplaceNode, IterationOptionsNode, CustomCFGNode, CustomCFGKeyframeNode) @@ -92,6 +93,8 @@ NODE_CLASS_MAPPINGS = { "ADE_LoadAnimateLCMI2VModel": LoadAnimateLCMI2VModelNode, "ADE_UpscaleAndVAEEncode": UpscaleAndVaeEncode, "ADE_InjectI2VIntoAnimateDiffModel": LoadAnimateDiffAndInjectI2VNode, + # CameraCtrl Nodes + "ADE_LoadCameraCtrlAdapter": LoadCameraCtrlAdapter, # MaskedLoraLoader #"ADE_MaskedLoadLora": MaskedLoraLoader, # Deprecated Nodes @@ -165,6 +168,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ADE_LoadAnimateLCMI2VModel": "Load AnimateLCM-I2V Model πŸŽ­πŸ…πŸ…“β‘‘", "ADE_UpscaleAndVAEEncode": "Scale Ref Image and VAE Encode πŸŽ­πŸ…πŸ…“β‘‘", "ADE_InjectI2VIntoAnimateDiffModel": "πŸ§ͺInject I2V into AnimateDiff Model πŸŽ­πŸ…πŸ…“β‘‘", + # CameraCtrl Nodes + "ADE_LoadCameraCtrlAdapter": "Load CameraCtrl Adapter πŸŽ­πŸ…πŸ…“β‘‘", # MaskedLoraLoader #"ADE_MaskedLoadLora": "Load LoRA (Masked) πŸŽ­πŸ…πŸ…“", # Deprecated Nodes diff --git a/animatediff/nodes_gen2.py b/animatediff/nodes_gen2.py index c7829a5..a4404ec 100644 --- a/animatediff/nodes_gen2.py +++ b/animatediff/nodes_gen2.py @@ -10,7 +10,7 @@ from comfy.sd import VAE from .ad_settings import AnimateDiffSettings from .context import ContextOptionsGroup from .logger import logger -from .utils_model import BIGMAX, BetaSchedules, ScaleMethods, CropMethods, get_available_motion_models +from .utils_model import BIGMAX, BetaSchedules, ScaleMethods, CropMethods, get_available_motion_models, get_motion_model_path from .utils_motion import ADKeyframeGroup, ADKeyframe from .motion_lora import MotionLoraList from .model_injection import (InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelPatcher, create_fresh_motion_module, create_fresh_encoder_only_model, @@ -286,6 +286,33 @@ class LoadAnimateDiffAndInjectI2VNode: return (loaded_motion_model,) +class LoadCameraCtrlAdapter: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model_name": (get_available_motion_models(),), + }, + } + + RETURN_TYPES = ("MOTION_MODEL_ADE",)#("CAMERA_CTRL",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/β‘‘ Gen2 nodes β‘‘/CameraCtrl" + FUNCTION = "load_camera_ctrl" + + def load_camera_ctrl(self, model_name: str): + model_path = get_motion_model_path(model_name) + logger.info(f"Loading CameraCtrl Adapter {model_name}") + mm_state_dict: dict[str, torch.Tensor] = comfy.utils.load_torch_file(model_path, safe_load=True) + from pathlib import Path + with open(Path(__file__).parent.parent.parent / "cameractrl_keys.txt", "w") as cfile: + for key in mm_state_dict: + if type(mm_state_dict[key]) == torch.Tensor: + cfile.write(f"{key}: {list(mm_state_dict[key].shape)}\n") + else: + cfile.write(f"{key}: {mm_state_dict[key]}\n") + return (None,) + + class ADKeyframeNode: @classmethod def INPUT_TYPES(s):