Use internal view_options for CameraPoseEncoder to overcome 16 max context_length, made rearrange_hidden_shapes irrelevant by figuring out the proper rearrange inside CameraPoseEncoder's forward
This commit is contained in:
@@ -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])}")
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user