diff --git a/animatediff/motion_module_ad.py b/animatediff/motion_module_ad.py index 89f9a3d..2379479 100644 --- a/animatediff/motion_module_ad.py +++ b/animatediff/motion_module_ad.py @@ -345,7 +345,7 @@ class AnimateDiffModel(nn.Module): def cleanup(self): self._reset_sub_idxs() - self._reset_scale_multiplier() + self._reset_scale() self._reset_temp_vars() if self.img_encoder is not None: self.img_encoder.cleanup() @@ -465,15 +465,15 @@ class AnimateDiffModel(nn.Module): if self.mid_block is not None: self.mid_block.set_video_length(video_length, full_length) - def set_scale(self, multival: Union[float, Tensor]): - if multival is None: - multival = 1.0 - if type(multival) == Tensor: - self._set_scale_multiplier(1.0) - self._set_scale_mask(multival) - else: - self._set_scale_multiplier(multival) - self._set_scale_mask(None) + def set_scale(self, scale: Union[float, Tensor, None], per_block_list: Union[list[PerBlock], None]=None): + if self.down_blocks is not None: + for block in self.down_blocks: + block.set_scale(scale, per_block_list) + if self.up_blocks is not None: + for block in self.up_blocks: + block.set_scale(scale, per_block_list) + if self.mid_block is not None: + self.mid_block.set_scale(scale, per_block_list) def set_effect(self, multival: Union[float, Tensor]): # keep track of if model is in effect @@ -539,26 +539,6 @@ class AnimateDiffModel(nn.Module): if self.up_blocks is not None: for block in self.up_blocks: block.set_camera_features(camera_features=list(reversed(camera_features))) - - def _set_scale_multiplier(self, multiplier: Union[float, None]): - if self.down_blocks is not None: - for block in self.down_blocks: - block.set_scale_multiplier(multiplier) - if self.up_blocks is not None: - for block in self.up_blocks: - block.set_scale_multiplier(multiplier) - if self.mid_block is not None: - self.mid_block.set_scale_multiplier(multiplier) - - def _set_scale_mask(self, mask: Tensor): - if self.down_blocks is not None: - for block in self.down_blocks: - block.set_scale_mask(mask) - if self.up_blocks is not None: - for block in self.up_blocks: - block.set_scale_mask(mask) - if self.mid_block is not None: - self.mid_block.set_scale_mask(mask) def _reset_temp_vars(self): if self.down_blocks is not None: @@ -570,8 +550,8 @@ class AnimateDiffModel(nn.Module): if self.mid_block is not None: self.mid_block.reset_temp_vars() - def _reset_scale_multiplier(self): - self._set_scale_multiplier(None) + def _reset_scale(self): + self.set_scale(None) def _reset_sub_idxs(self): self.set_sub_idxs(None) @@ -606,14 +586,10 @@ class MotionModule(nn.Module): for motion_module in self.motion_modules: motion_module.set_video_length(video_length, full_length) - def set_scale_multiplier(self, multiplier: Union[float, None]): + def set_scale(self, scale: Union[float, Tensor, None], per_block_list: Union[list[PerBlock], None]=None): for motion_module in self.motion_modules: - motion_module.set_scale_multiplier(multiplier) - - def set_scale_mask(self, mask: Tensor): - for motion_module in self.motion_modules: - motion_module.set_scale_mask(mask) - + motion_module.set_scale(scale, per_block_list) + def set_effect(self, multival: Union[float, Tensor]): for motion_module in self.motion_modules: motion_module.set_effect(multival) @@ -711,15 +687,9 @@ class VanillaTemporalModule(nn.Module): self.video_length = video_length self.full_length = full_length self.temporal_transformer.set_video_length(video_length, full_length) - - def set_scale_multiplier(self, multiplier: Union[float, None]): - self.temporal_transformer.set_scale_multiplier(multiplier) - def set_scale_mask(self, mask: Tensor): - self.temporal_transformer.set_scale_mask(mask) - - def set_scale(self, multiplier: Union[float, Tensor, None], per_block_list: Union[list[PerBlock], None]=None): - self.temporal_transformer.set_scale(multiplier) + def set_scale(self, scale: Union[float, Tensor, None], per_block_list: Union[list[PerBlock], None]=None): + self.temporal_transformer.set_scale(scale) def set_effect(self, multival: Union[float, Tensor], per_block_list: Union[list[PerBlock], None]=None): if per_block_list is not None: @@ -848,6 +818,17 @@ class TemporalTransformer3DModel(nn.Module): ops=comfy.ops.disable_weight_init, ): super().__init__() + self.video_length = 16 + self.full_length = 16 + self.raw_scale_mask: Union[Tensor, None] = None + self.temp_scale_mask: Union[Tensor, None] = None + self.sub_idxs: Union[list[int], None] = None + self.prev_hidden_states_batch = 0 + + # cameractrl stuff + self.raw_cameractrl_effect: Union[float, Tensor] = None + self.temp_cameractrl_effect: Union[float, Tensor] = None + self.prev_cameractrl_hidden_states_batch = 0 inner_dim = num_attention_heads * attention_head_dim @@ -880,18 +861,24 @@ class TemporalTransformer3DModel(nn.Module): self.proj_out = ops.Linear(inner_dim, in_channels) def set_video_length(self, video_length: int, full_length: int): - for block in self.transformer_blocks: - block.set_video_length(video_length, full_length) - + self.video_length = video_length + self.full_length = full_length + def set_scale_multiplier(self, multiplier: Union[float, None]): for block in self.transformer_blocks: - block.set_scale_multiplier(multiplier) + mult = multiplier + block.set_scale_multiplier(mult) def set_scale_mask(self, mask: Tensor): - for block in self.transformer_blocks: - block.set_scale_mask(mask) + self.raw_scale_mask = mask + self.temp_scale_mask = None - def set_scale(self, scale: Union[float, Tensor, None]): + def set_scale(self, scale: Union[float, Tensor, None], per_attn_list: Union[list[PerAttn], None]=None): + # if per_attn_list is not None: + # for per_attn in per_attn_list: + # if per_attn.attn_idx == idx: + # mult = per_attn.scale + # break if type(scale) == Tensor: self.set_scale_mask(scale) self.set_scale_multiplier(None) @@ -899,141 +886,12 @@ class TemporalTransformer3DModel(nn.Module): self.set_scale_mask(None) self.set_scale_multiplier(scale) - def set_cameractrl_effect(self, multival: Union[float, Tensor]): - for block in self.transformer_blocks: - block.set_cameractrl_effect(multival) - - def set_sub_idxs(self, sub_idxs: list[int]): - for block in self.transformer_blocks: - block.set_sub_idxs(sub_idxs) - - def reset_temp_vars(self): - for block in self.transformer_blocks: - block.reset_temp_vars() - - def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, view_options: ContextOptions=None, mm_kwargs: dict[str]=None): - batch, channel, height, width = hidden_states.shape - residual = hidden_states - # add some casts for fp8 purposes - does not affect speed otherwise - hidden_states = self.norm(hidden_states).to(hidden_states.dtype) - inner_dim = hidden_states.shape[1] - hidden_states = hidden_states.permute(0, 2, 3, 1).reshape( - batch, height * width, inner_dim - ) - hidden_states = self.proj_in(hidden_states).to(hidden_states.dtype) - - # Transformer Blocks - for block in self.transformer_blocks: - hidden_states = block( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - view_options=view_options, - mm_kwargs=mm_kwargs - ) - - # output - hidden_states = self.proj_out(hidden_states) - hidden_states = ( - hidden_states.reshape(batch, height, width, inner_dim) - .permute(0, 3, 1, 2) - .contiguous() - ) - - output = hidden_states + residual - - return output - - -class TemporalTransformerBlock(nn.Module): - def __init__( - self, - dim, - num_attention_heads, - attention_head_dim, - attention_block_types=( - "Temporal_Self", - "Temporal_Self", - ), - dropout=0.0, - norm_num_groups=32, - cross_attention_dim=768, - activation_fn="geglu", - attention_bias=False, - upcast_attention=False, - cross_frame_attention_mode=None, - temporal_pe=False, - temporal_pe_max_len=24, - ops=comfy.ops.disable_weight_init, - ): - super().__init__() - self.video_length = 16 - self.full_length = 16 - self.raw_scale_mask: Union[Tensor, None] = None - self.temp_scale_mask: Union[Tensor, None] = None - self.sub_idxs: Union[list[int], None] = None - self.prev_hidden_states_batch = 0 - - # cameractrl stuff - self.raw_cameractrl_effect: Union[float, Tensor] = None - self.temp_cameractrl_effect: Union[float, Tensor] = None - self.prev_cameractrl_hidden_states_batch = 0 - - attention_blocks: Iterable[VersatileAttention] = [] - norms = [] - - for block_name in attention_block_types: - attention_blocks.append( - VersatileAttention( - attention_mode=block_name.split("_")[0], - context_dim=cross_attention_dim # called context_dim for ComfyUI impl - if block_name.endswith("_Cross") - else None, - query_dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - dropout=dropout, - #bias=attention_bias, # remove for Comfy CrossAttention - #upcast_attention=upcast_attention, # remove for Comfy CrossAttention - cross_frame_attention_mode=cross_frame_attention_mode, - temporal_pe=temporal_pe, - temporal_pe_max_len=temporal_pe_max_len, - ops=ops, - ) - ) - norms.append(ops.LayerNorm(dim)) - - attention_blocks[0].camera_feature_enabled = True - self.attention_blocks: Iterable[VersatileAttention] = nn.ModuleList(attention_blocks) - self.norms = nn.ModuleList(norms) - - self.ff = FeedForward(dim, dropout=dropout, glu=(activation_fn == "geglu"), operations=ops) - self.ff_norm = ops.LayerNorm(dim) - - def set_video_length(self, video_length: int, full_length: int): - self.video_length = video_length - self.full_length = full_length - - def set_scale_multiplier(self, multiplier: Union[float, None], per_attn_list: Union[list[PerAttn], None]=None): - for idx, block in enumerate(self.attention_blocks): - mult = multiplier - if per_attn_list is not None: - for per_attn in per_attn_list: - if per_attn.attn_idx == idx: - mult = per_attn.scale - break - block.set_scale_multiplier(mult) - - def set_scale_mask(self, mask: Tensor): - self.raw_scale_mask = mask - self.temp_scale_mask = None - def set_cameractrl_effect(self, multival: Union[float, Tensor]): self.raw_cameractrl_effect = multival self.temp_cameractrl_effect = None def set_sub_idxs(self, sub_idxs: list[int]): - for block in self.attention_blocks: + for block in self.transformer_blocks: block.set_sub_idxs(sub_idxs) def reset_temp_vars(self): @@ -1043,9 +901,9 @@ class TemporalTransformerBlock(nn.Module): del self.temp_cameractrl_effect self.temp_cameractrl_effect = None self.prev_cameractrl_hidden_states_batch = 0 - for block in self.attention_blocks: + for block in self.transformer_blocks: block.reset_temp_vars() - + def get_scale_mask(self, hidden_states: Tensor) -> Union[Tensor, None]: # if no raw mask, return None if self.raw_scale_mask is None: @@ -1135,21 +993,128 @@ class TemporalTransformerBlock(nn.Module): return self.temp_cameractrl_effect[:, self.sub_idxs, :] return self.temp_cameractrl_effect + def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, view_options: ContextOptions=None, mm_kwargs: dict[str]=None): + batch, channel, height, width = hidden_states.shape + residual = hidden_states + scale_mask = self.get_scale_mask(hidden_states) + cameractrl_effect = self.get_cameractrl_effect(hidden_states) + # add some casts for fp8 purposes - does not affect speed otherwise + hidden_states = self.norm(hidden_states).to(hidden_states.dtype) + inner_dim = hidden_states.shape[1] + hidden_states = hidden_states.permute(0, 2, 3, 1).reshape( + batch, height * width, inner_dim + ) + hidden_states = self.proj_in(hidden_states).to(hidden_states.dtype) + + # Transformer Blocks + for block in self.transformer_blocks: + hidden_states = block( + hidden_states, + encoder_hidden_states=encoder_hidden_states, + attention_mask=attention_mask, + video_length=self.video_length, + scale_mask=scale_mask, + cameractrl_effect=cameractrl_effect, + view_options=view_options, + mm_kwargs=mm_kwargs + ) + + # output + hidden_states = self.proj_out(hidden_states) + hidden_states = ( + hidden_states.reshape(batch, height, width, inner_dim) + .permute(0, 3, 1, 2) + .contiguous() + ) + + output = hidden_states + residual + + return output + + +class TemporalTransformerBlock(nn.Module): + def __init__( + self, + dim, + num_attention_heads, + attention_head_dim, + attention_block_types=( + "Temporal_Self", + "Temporal_Self", + ), + dropout=0.0, + norm_num_groups=32, + cross_attention_dim=768, + activation_fn="geglu", + attention_bias=False, + upcast_attention=False, + cross_frame_attention_mode=None, + temporal_pe=False, + temporal_pe_max_len=24, + ops=comfy.ops.disable_weight_init, + ): + super().__init__() + + attention_blocks: Iterable[VersatileAttention] = [] + norms = [] + + for block_name in attention_block_types: + attention_blocks.append( + VersatileAttention( + attention_mode=block_name.split("_")[0], + context_dim=cross_attention_dim # called context_dim for ComfyUI impl + if block_name.endswith("_Cross") + else None, + query_dim=dim, + heads=num_attention_heads, + dim_head=attention_head_dim, + dropout=dropout, + #bias=attention_bias, # remove for Comfy CrossAttention + #upcast_attention=upcast_attention, # remove for Comfy CrossAttention + cross_frame_attention_mode=cross_frame_attention_mode, + temporal_pe=temporal_pe, + temporal_pe_max_len=temporal_pe_max_len, + ops=ops, + ) + ) + norms.append(ops.LayerNorm(dim)) + + attention_blocks[0].camera_feature_enabled = True + self.attention_blocks: Iterable[VersatileAttention] = nn.ModuleList(attention_blocks) + self.norms = nn.ModuleList(norms) + + self.ff = FeedForward(dim, dropout=dropout, glu=(activation_fn == "geglu"), operations=ops) + self.ff_norm = ops.LayerNorm(dim) + + def set_scale_multiplier(self, multiplier: Union[float, None]): + for idx, block in enumerate(self.attention_blocks): + mult = multiplier + block.set_scale_multiplier(mult) + + def set_sub_idxs(self, sub_idxs: list[int]): + for block in self.attention_blocks: + block.set_sub_idxs(sub_idxs) + + def reset_temp_vars(self): + for block in self.attention_blocks: + block.reset_temp_vars() + def forward( self, hidden_states: Tensor, encoder_hidden_states: Tensor=None, attention_mask: Tensor=None, + video_length: int=None, + scale_mask: Tensor=None, + cameractrl_effect: Union[float, Tensor] = None, view_options: Union[ContextOptions, None]=None, mm_kwargs: dict[str]=None, ): - scale_mask = self.get_scale_mask(hidden_states) - cameractrl_effect = self.get_cameractrl_effect(hidden_states) # make view_options None if context_length > video_length, or if equal and equal not allowed if view_options: - if view_options.context_length > self.video_length: + if view_options.context_length > video_length: view_options = None - elif view_options.context_length == self.video_length and not view_options.use_on_equal_length: + elif view_options.context_length == video_length and not view_options.use_on_equal_length: view_options = None if not view_options: for attention_block, norm in zip(self.attention_blocks, self.norms): @@ -1161,7 +1126,7 @@ class TemporalTransformerBlock(nn.Module): if attention_block.is_cross_attention else None, attention_mask=attention_mask, - video_length=self.video_length, + video_length=video_length, scale_mask=scale_mask, cameractrl_effect=cameractrl_effect, mm_kwargs=mm_kwargs @@ -1171,12 +1136,12 @@ class TemporalTransformerBlock(nn.Module): # views idea gotten from diffusers AnimateDiff FreeNoise implementation: # https://github.com/arthur-qiu/FreeNoise-AnimateDiff/blob/main/animatediff/models/motion_module.py # apply sliding context windows (views) - views = get_context_windows(num_frames=self.video_length, opts=view_options) - hidden_states = rearrange(hidden_states, "(b f) d c -> b f d c", f=self.video_length) + views = get_context_windows(num_frames=video_length, opts=view_options) + hidden_states = rearrange(hidden_states, "(b f) d c -> b f d c", f=video_length) value_final = torch.zeros_like(hidden_states) count_final = torch.zeros_like(hidden_states) # bias_final = [0.0] * video_length - batched_conds = hidden_states.size(1) // self.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: