Moved scale_mask and cameractrl_effec to TemporalTransformerBlock
This commit is contained in:
+144
-139
@@ -848,17 +848,6 @@ 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
|
||||
|
||||
@@ -891,16 +880,16 @@ class TemporalTransformer3DModel(nn.Module):
|
||||
self.proj_out = ops.Linear(inner_dim, in_channels)
|
||||
|
||||
def set_video_length(self, video_length: int, full_length: int):
|
||||
self.video_length = video_length
|
||||
self.full_length = full_length
|
||||
for block in self.transformer_blocks:
|
||||
block.set_video_length(video_length, full_length)
|
||||
|
||||
def set_scale_multiplier(self, multiplier: Union[float, None]):
|
||||
for block in self.transformer_blocks:
|
||||
block.set_scale_multiplier(multiplier)
|
||||
|
||||
def set_scale_mask(self, mask: Tensor):
|
||||
self.raw_scale_mask = mask
|
||||
self.temp_scale_mask = None
|
||||
for block in self.transformer_blocks:
|
||||
block.set_scale_mask(mask)
|
||||
|
||||
def set_scale(self, scale: Union[float, Tensor, None]):
|
||||
if type(scale) == Tensor:
|
||||
@@ -910,13 +899,141 @@ 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]):
|
||||
self.sub_idxs = sub_idxs
|
||||
for block in self.transformer_blocks:
|
||||
for block in self.attention_blocks:
|
||||
block.set_sub_idxs(sub_idxs)
|
||||
|
||||
def reset_temp_vars(self):
|
||||
@@ -926,9 +1043,9 @@ class TemporalTransformer3DModel(nn.Module):
|
||||
del self.temp_cameractrl_effect
|
||||
self.temp_cameractrl_effect = None
|
||||
self.prev_cameractrl_hidden_states_batch = 0
|
||||
for block in self.transformer_blocks:
|
||||
for block in self.attention_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:
|
||||
@@ -1018,133 +1135,21 @@ class TemporalTransformer3DModel(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], 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_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 > video_length:
|
||||
if view_options.context_length > self.video_length:
|
||||
view_options = None
|
||||
elif view_options.context_length == video_length and not view_options.use_on_equal_length:
|
||||
elif view_options.context_length == self.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):
|
||||
@@ -1156,7 +1161,7 @@ class TemporalTransformerBlock(nn.Module):
|
||||
if attention_block.is_cross_attention
|
||||
else None,
|
||||
attention_mask=attention_mask,
|
||||
video_length=video_length,
|
||||
video_length=self.video_length,
|
||||
scale_mask=scale_mask,
|
||||
cameractrl_effect=cameractrl_effect,
|
||||
mm_kwargs=mm_kwargs
|
||||
@@ -1166,12 +1171,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=video_length, opts=view_options)
|
||||
hidden_states = rearrange(hidden_states, "(b f) d c -> b f d c", f=video_length)
|
||||
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)
|
||||
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) // video_length
|
||||
batched_conds = hidden_states.size(1) // self.video_length
|
||||
# store original camera_feature, if present
|
||||
has_camera_feature = False
|
||||
if mm_kwargs is not None:
|
||||
|
||||
Reference in New Issue
Block a user