diff --git a/animatediff/adapter_cameractrl.py b/animatediff/adapter_cameractrl.py index 8339741..45966d3 100644 --- a/animatediff/adapter_cameractrl.py +++ b/animatediff/adapter_cameractrl.py @@ -8,9 +8,11 @@ from einops import rearrange import comfy.ops -from .motion_module_ad import TemporalTransformerBlock +from .context import ContextOptions, ContextFuseMethod, ContextSchedules +from .motion_module_ad import TemporalTransformerBlock, get_position_encoding_max_len from .logger import logger + def conv_nd(dims, *args, **kwargs): """ Create a 1D, 2D, or 3D convolution module. @@ -209,12 +211,12 @@ class CameraPoseEncoder(nn.Module): cross_attention_dim=None, temporal_pe=temporal_position_encoding, temporal_pe_max_len=temporal_position_encoding_max_len, - rearrange_hidden_shapes=False, # different from AD 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) + self.temporal_pe_max_len = 16 def forward(self, x: Tensor, video_length: int, batched_number: int=1): # rearrange to match expected format @@ -224,6 +226,13 @@ class CameraPoseEncoder(nn.Module): x = self.unshuffle(x) # extract features features = [] + # prepare view_options, if needed + view_options = ContextOptions( + context_length=self.temporal_pe_max_len, + context_overlap=self.temporal_pe_max_len//2, # at 16 max_len, context_overlap will be 8 + context_schedule=ContextSchedules.STATIC_STANDARD, + fuse_method=ContextFuseMethod.PYRAMID, + ) # logger.warn(f"x dtype: {x.dtype}, device: {x.device}") # logger.warn(f"dtype: {get_parameter_dtype(self)}, device: {get_parameter_device(self)}") x = self.encoder_conv_in(x.to(dtype=get_parameter_dtype(self), device=get_parameter_device(self))) @@ -231,9 +240,9 @@ class CameraPoseEncoder(nn.Module): 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, video_length=video_length) - x = rearrange(x, '(h w) b c -> b c h w', h=h, w=w) + x = rearrange(x, 'b c h w -> b (h w) c') # h w are in middle instead of beginning like in diffusers + x = attention_layer(x, video_length=video_length, view_options=view_options) + x = rearrange(x, 'b (h w) c -> b c h w', h=h, w=w) # h w are in middle instead of beginning like in diffusers features.append(x) # for idx, feature in enumerate(features): # logger.info(f"{idx}: {feature.shape}, {float(feature[0][0][0][0])}") diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 83dde0e..127b1a5 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -14,7 +14,8 @@ from comfy.model_base import BaseModel from .ad_settings import AnimateDiffSettings, AdjustPE, AdjustWeight from .adapter_cameractrl import CameraPoseEncoder, CameraEntry, prepare_pose_embedding from .context import ContextOptions, ContextOptions, ContextOptionsGroup -from .motion_module_ad import AnimateDiffModel, AnimateDiffFormat, EncoderOnlyAnimateDiffModel, VersatileAttention, has_mid_block, normalize_ad_state_dict +from .motion_module_ad import (AnimateDiffModel, AnimateDiffFormat, EncoderOnlyAnimateDiffModel, VersatileAttention, + has_mid_block, normalize_ad_state_dict, get_position_encoding_max_len) from .logger import logger from .utils_motion import ADKeyframe, ADKeyframeGroup, MotionCompatibilityError, get_combined_multival, ade_broadcast_image_to, normalize_min_max from .motion_lora import MotionLoraInfo, MotionLoraList @@ -555,6 +556,7 @@ def inject_camera_encoder_into_model(motion_model: MotionModelPatcher, camera_ct dtype=comfy.model_management.unet_dtype() ) camera_encoder.load_state_dict(camera_state_dict) + camera_encoder.temporal_pe_max_len = get_position_encoding_max_len(camera_state_dict, mm_name=camera_ctrl_name, mm_format=AnimateDiffFormat.ANIMATEDIFF) motion_model.model.set_camera_encoder(camera_encoder=camera_encoder) # initialize qkv_merge on specific attention blocks, and load keys for key in attention_state_dict: diff --git a/animatediff/motion_module_ad.py b/animatediff/motion_module_ad.py index 468adcd..53ba251 100644 --- a/animatediff/motion_module_ad.py +++ b/animatediff/motion_module_ad.py @@ -708,7 +708,6 @@ class TemporalTransformer3DModel(nn.Module): cross_frame_attention_mode=None, temporal_pe=False, temporal_pe_max_len=24, - rearrange_hidden_shapes=True, ops=comfy.ops.disable_weight_init, ): super().__init__() @@ -747,7 +746,6 @@ class TemporalTransformer3DModel(nn.Module): cross_frame_attention_mode=cross_frame_attention_mode, temporal_pe=temporal_pe, temporal_pe_max_len=temporal_pe_max_len, - rearrange_hidden_shapes=rearrange_hidden_shapes, ops=ops, ) for d in range(num_layers) @@ -931,7 +929,6 @@ class TemporalTransformerBlock(nn.Module): cross_frame_attention_mode=None, temporal_pe=False, temporal_pe_max_len=24, - rearrange_hidden_shapes=True, ops=comfy.ops.disable_weight_init, ): super().__init__() @@ -955,7 +952,6 @@ class TemporalTransformerBlock(nn.Module): cross_frame_attention_mode=cross_frame_attention_mode, temporal_pe=temporal_pe, temporal_pe_max_len=temporal_pe_max_len, - rearrange_hidden_shapes=rearrange_hidden_shapes, ops=ops, ) ) @@ -1019,8 +1015,16 @@ class TemporalTransformerBlock(nn.Module): count_final = torch.zeros_like(hidden_states) # bias_final = [0.0] * video_length batched_conds = hidden_states.size(1) // video_length + # store original camera_feature, if present + has_camera_feature = False + if mm_kwargs is not None: + has_camera_feature = True + orig_camera_feature = mm_kwargs["camera_feature"] + # perform view options for sub_idxs in views: 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, :] for attention_block, norm in zip(self.attention_blocks, self.norms): norm_hidden_states = norm(sub_hidden_states).to(sub_hidden_states.dtype) sub_hidden_states = ( @@ -1059,7 +1063,10 @@ class TemporalTransformerBlock(nn.Module): weights_tensor = torch.Tensor(weights).to(device=hidden_states.device).unsqueeze(0).unsqueeze(-1).unsqueeze(-1) value_final[:, sub_idxs] += sub_hidden_states * weights_tensor count_final[:, sub_idxs] += weights_tensor - + # restore original camera_feature + if has_camera_feature: + mm_kwargs["camera_feature"] = orig_camera_feature + del orig_camera_feature # get weighted average of sub_hidden_states, if fuse method requires it # if view_options.fuse_method != ContextFuseMethod.RELATIVE: hidden_states = value_final / count_final @@ -1107,7 +1114,6 @@ class VersatileAttention(CrossAttentionMM): cross_frame_attention_mode=None, temporal_pe=False, temporal_pe_max_len=24, - rearrange_hidden_shapes=True, ops=comfy.ops.disable_weight_init, *args, **kwargs, @@ -1121,7 +1127,6 @@ class VersatileAttention(CrossAttentionMM): self.query_dim: int = kwargs["query_dim"] self.qkv_merge: comfy.ops.disable_weight_init.Linear = None self.camera_feature_enabled = False - self.rearrange_hidden_shapes = rearrange_hidden_shapes self.pos_encoder = ( PositionalEncoding( @@ -1162,11 +1167,10 @@ class VersatileAttention(CrossAttentionMM): if self.attention_mode != "Temporal": raise NotImplementedError - if self.rearrange_hidden_shapes: - d = hidden_states.shape[1] - hidden_states = rearrange( - hidden_states, "(b f) d c -> (b d) f c", f=video_length - ) + d = hidden_states.shape[1] + hidden_states = rearrange( + hidden_states, "(b f) d c -> (b d) f c", f=video_length + ) if self.pos_encoder is not None: hidden_states = self.pos_encoder(hidden_states).to(hidden_states.dtype) @@ -1189,8 +1193,7 @@ class VersatileAttention(CrossAttentionMM): scale_mask=scale_mask, ) - if self.rearrange_hidden_shapes: - hidden_states = rearrange(hidden_states, "(b d) f c -> (b f) d c", d=d) + hidden_states = rearrange(hidden_states, "(b d) f c -> (b f) d c", d=d) return hidden_states