Moved scale_mask and cameractrl_effec to TemporalTransformerBlock

This commit is contained in:
Jedrzej Kosinski
2024-08-15 23:01:56 -05:00
parent e34c410ceb
commit 2a520db00c
+144 -139
View File
@@ -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: