Generalized AnimateDiffModel initialization so mm_state_dict can be used to determine if down/up blocks should be present and the attention block count, added init_kwargs input to AnimateDiffModel so some behavior can be adapted by unusual models

This commit is contained in:
Jedrzej Kosinski
2024-11-29 01:33:49 -06:00
parent 5ee8cf677e
commit 95cfd12a81
+64 -43
View File
@@ -156,10 +156,16 @@ def is_fancyvideo(mm_state_dict: dict[str, Tensor]) -> bool:
def get_down_block_max(mm_state_dict: dict[str, Tensor]) -> int:
return get_block_max(mm_state_dict, "down_blocks")
def get_up_block_max(mm_state_dict: dict[str, Tensor]) -> int:
return get_block_max(mm_state_dict, "up_blocks")
def get_block_max(mm_state_dict: dict[str, Tensor], block_name: str) -> int:
# keep track of biggest down_block count in module
biggest_block = 0
biggest_block = -1
for key in mm_state_dict.keys():
if "down_blocks" in key:
if block_name in key:
try:
block_int = key.split(".")[1]
block_num = int(block_int)
@@ -169,7 +175,6 @@ def get_down_block_max(mm_state_dict: dict[str, Tensor]) -> int:
pass
return biggest_block
def has_mid_block(mm_state_dict: dict[str, Tensor]):
# check if keys contain mid_block
for key in mm_state_dict.keys():
@@ -177,6 +182,17 @@ def has_mid_block(mm_state_dict: dict[str, Tensor]):
return True
return False
_regex_attention_blocks_num = re.compile(r'\.attention_blocks\.(\d+)\.')
def get_attention_block_max_len(mm_state_dict: dict[str, Tensor]):
biggest_attention = -1
for key in mm_state_dict.keys():
found = _regex_attention_blocks_num.search(key)
if found:
attention_num = int(found.group(1))
if attention_num > biggest_attention:
biggest_attention = attention_num
return biggest_attention + 1
def get_position_encoding_max_len(mm_state_dict: dict[str, Tensor], mm_name: str, mm_format: str) -> Union[int, None]:
# use pos_encoder.pe entries to determine max length - [1, {max_length}, {320|640|1280}]
@@ -337,21 +353,33 @@ def convert_hellomeme_state_dict(mm_state_dict: dict[str, Tensor]):
del mm_state_dict[key]
class InitKwargs:
GET_UNET_FUNC = "get_unet_func"
ATTN_BLOCK_TYPE = "attn_block_type"
class BlockType:
UP = "up"
DOWN = "down"
MID = "mid"
def get_unet_default(wrapper: 'AnimateDiffModel', model: ModelPatcher):
return model.model.diffusion_model
class AnimateDiffModel(nn.Module):
def __init__(self, mm_state_dict: dict[str, Tensor], mm_info: AnimateDiffInfo):
def __init__(self, mm_state_dict: dict[str, Tensor], mm_info: AnimateDiffInfo, init_kwargs: dict[str]={}):
super().__init__()
self.mm_info = mm_info
self.down_blocks: Iterable[MotionModule] = nn.ModuleList([])
self.up_blocks: Iterable[MotionModule] = nn.ModuleList([])
self.down_blocks: list[MotionModule] = None
self.up_blocks: list[MotionModule] = None
self.mid_block: Union[MotionModule, None] = None
self.encoding_max_len = get_position_encoding_max_len(mm_state_dict, mm_info.mm_name, mm_info.mm_format)
self.has_position_encoding = self.encoding_max_len is not None
self.attn_len = get_attention_block_max_len(mm_state_dict)
self.attn_type = init_kwargs.get(InitKwargs.ATTN_BLOCK_TYPE, "Temporal_Self")
self.attn_block_types = tuple([self.attn_type] * self.attn_len)
# determine ops to use (to support fp8 properly)
if comfy.model_management.unet_manual_cast(comfy.model_management.unet_dtype(), comfy.model_management.get_torch_device()) is None:
ops = comfy.ops.disable_weight_init
@@ -364,16 +392,24 @@ class AnimateDiffModel(nn.Module):
else:
layer_channels = (320, 640, 1280, 1280)
self.layer_channels = layer_channels
self.middle_channel = 1280
# fill out down/up blocks and middle block, if present
for idx, c in enumerate(layer_channels):
self.down_blocks.append(MotionModule(c, temporal_pe=self.has_position_encoding,
temporal_pe_max_len=self.encoding_max_len, block_type=BlockType.DOWN, block_idx=idx, ops=ops))
for idx, c in enumerate(list(reversed(layer_channels))):
self.up_blocks.append(MotionModule(c, temporal_pe=self.has_position_encoding,
temporal_pe_max_len=self.encoding_max_len, block_type=BlockType.UP, block_idx=idx, ops=ops))
if get_down_block_max(mm_state_dict) > -1:
self.down_blocks = nn.ModuleList([])
for idx, c in enumerate(layer_channels):
self.down_blocks.append(MotionModule(c, temporal_pe=self.has_position_encoding,
temporal_pe_max_len=self.encoding_max_len, block_type=BlockType.DOWN, block_idx=idx,
attention_block_types=self.attn_block_types, ops=ops))
if get_up_block_max(mm_state_dict) > -1:
self.up_blocks = nn.ModuleList([])
for idx, c in enumerate(list(reversed(layer_channels))):
self.up_blocks.append(MotionModule(c, temporal_pe=self.has_position_encoding,
temporal_pe_max_len=self.encoding_max_len, block_type=BlockType.UP, block_idx=idx,
attention_block_types=self.attn_block_types, ops=ops))
if has_mid_block(mm_state_dict):
self.mid_block = MotionModule(1280, temporal_pe=self.has_position_encoding,
temporal_pe_max_len=self.encoding_max_len, block_type=BlockType.MID, ops=ops)
self.mid_block = MotionModule(self.middle_channel, temporal_pe=self.has_position_encoding,
temporal_pe_max_len=self.encoding_max_len, block_type=BlockType.MID,
attention_block_types=self.attn_block_types, ops=ops)
self.AD_video_length: int = 24
self.effect_model = 1.0
self.effect_per_block_list = None
@@ -395,6 +431,8 @@ class AnimateDiffModel(nn.Module):
self.init_fps_embedding(mm_state_dict)
if has_motion_embedding(mm_state_dict):
self.init_motion_embedding(mm_state_dict)
# get_unet_func initialization
self.get_unet_func = init_kwargs.get(InitKwargs.GET_UNET_FUNC, get_unet_default)
def init_img_encoder(self):
del self.img_encoder
@@ -501,7 +539,7 @@ class AnimateDiffModel(nn.Module):
self.img_encoder.cleanup()
def inject(self, model: ModelPatcher):
unet: openaimodel.UNetModel = model.model.diffusion_model
unet: openaimodel.UNetModel = self.get_unet_func(self, model)
# inject input (down) blocks
# SD15 mm contains 4 downblocks, each with 2 TemporalTransformers - 8 in total
# SDXL mm contains 3 downblocks, each with 2 TemporalTransformers - 6 in total
@@ -555,7 +593,7 @@ class AnimateDiffModel(nn.Module):
unet_idx += 1
def eject(self, model: ModelPatcher):
unet: openaimodel.UNetModel = model.model.diffusion_model
unet: openaimodel.UNetModel = self.get_unet_func(self, model)
# remove from input blocks (downblocks)
self._eject(unet.input_blocks)
# remove from output blocks (upblocks)
@@ -715,23 +753,24 @@ class MotionModule(nn.Module):
temporal_pe_max_len=24,
block_type: str=BlockType.DOWN,
block_idx: int=0,
attention_block_types=("Temporal_Self", "Temporal_Self"),
ops=comfy.ops.disable_weight_init
):
super().__init__()
if block_type == BlockType.MID:
# mid blocks contain only a single VanillaTemporalModule
self.motion_modules: Iterable[VanillaTemporalModule] = nn.ModuleList([get_motion_module(in_channels, block_type, block_idx, module_idx=0, temporal_pe=temporal_pe, temporal_pe_max_len=temporal_pe_max_len, ops=ops)])
self.motion_modules: list[VanillaTemporalModule] = nn.ModuleList([get_motion_module(in_channels, block_type, block_idx, module_idx=0, attention_block_types=attention_block_types, temporal_pe=temporal_pe, temporal_pe_max_len=temporal_pe_max_len, ops=ops)])
else:
# down blocks contain two VanillaTemporalModules
self.motion_modules: Iterable[VanillaTemporalModule] = nn.ModuleList(
self.motion_modules: list[VanillaTemporalModule] = nn.ModuleList(
[
get_motion_module(in_channels, block_type, block_idx, module_idx=0, temporal_pe=temporal_pe, temporal_pe_max_len=temporal_pe_max_len, ops=ops),
get_motion_module(in_channels, block_type, block_idx, module_idx=1, temporal_pe=temporal_pe, temporal_pe_max_len=temporal_pe_max_len, ops=ops)
get_motion_module(in_channels, block_type, block_idx, module_idx=0, attention_block_types=attention_block_types, temporal_pe=temporal_pe, temporal_pe_max_len=temporal_pe_max_len, ops=ops),
get_motion_module(in_channels, block_type, block_idx, module_idx=1, attention_block_types=attention_block_types, temporal_pe=temporal_pe, temporal_pe_max_len=temporal_pe_max_len, ops=ops)
]
)
# up blocks contain one additional VanillaTemporalModule
if block_type == BlockType.UP:
self.motion_modules.append(get_motion_module(in_channels, block_type, block_idx, module_idx=2, temporal_pe=temporal_pe, temporal_pe_max_len=temporal_pe_max_len, ops=ops))
self.motion_modules.append(get_motion_module(in_channels, block_type, block_idx, module_idx=2, attention_block_types=attention_block_types, temporal_pe=temporal_pe, temporal_pe_max_len=temporal_pe_max_len, ops=ops))
def set_video_length(self, video_length: int, full_length: int):
for motion_module in self.motion_modules:
@@ -772,8 +811,10 @@ class MotionModule(nn.Module):
def get_motion_module(in_channels, block_type: str, block_idx: int, module_idx: int,
attention_block_types: list[str],
temporal_pe, temporal_pe_max_len, ops=comfy.ops.disable_weight_init):
return VanillaTemporalModule(in_channels=in_channels, block_type=block_type, block_idx=block_idx, module_idx=module_idx,
attention_block_types=attention_block_types,
temporal_pe=temporal_pe, temporal_pe_max_len=temporal_pe_max_len, ops=ops)
@@ -1324,7 +1365,6 @@ class TemporalTransformerBlock(nn.Module):
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) // video_length
# store original camera_feature, if present
has_camera_feature = False
@@ -1354,23 +1394,6 @@ class TemporalTransformerBlock(nn.Module):
)
sub_hidden_states = rearrange(sub_hidden_states, "(b f) d c -> b f d c", f=len(sub_idxs))
# if view_options.fuse_method == ContextFuseMethod.RELATIVE:
# for pos, idx in enumerate(sub_idxs):
# # bias is the influence of a specific index in relation to the whole context window
# bias = 1 - abs(idx - (sub_idxs[0] + sub_idxs[-1]) / 2) / ((sub_idxs[-1] - sub_idxs[0] + 1e-2) / 2)
# bias = max(1e-2, bias)
# # take weighted averate relative to total bias of current idx
# bias_total = bias_final[idx]
# prev_weight = torch.tensor([bias_total / (bias_total + bias)],
# dtype=value_final.dtype, device=value_final.device).unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
# #prev_weight = torch.cat([prev_weight]*value_final.shape[1], dim=1)
# new_weight = torch.tensor([bias / (bias_total + bias)],
# dtype=value_final.dtype, device=value_final.device).unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
# #new_weight = torch.cat([new_weight]*value_final.shape[1], dim=1)
# test = value_final[:, idx:idx+1, :, :]
# value_final[:, idx:idx+1, :, :] = value_final[:, idx:idx+1, :, :] * prev_weight + sub_hidden_states[:, pos:pos+1, : ,:] * new_weight
# bias_final[idx] = bias_total + bias
# else:
weights = get_context_weights(len(sub_idxs), view_options.fuse_method) * batched_conds
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
@@ -1379,13 +1402,11 @@ class TemporalTransformerBlock(nn.Module):
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:
# get weighted average of sub_hidden_states
hidden_states = value_final / count_final
hidden_states = rearrange(hidden_states, "b f d c -> (b f) d c")
del value_final
del count_final
# del bias_final
hidden_states = self.ff(self.ff_norm(hidden_states)) + hidden_states
@@ -1521,7 +1542,7 @@ class VersatileAttention(CrossAttentionMM):
class EncoderOnlyAnimateDiffModel(AnimateDiffModel):
def __init__(self, mm_state_dict: dict[str, Tensor], mm_info: AnimateDiffInfo):
super().__init__(mm_state_dict=mm_state_dict, mm_info=mm_info)
self.down_blocks: Iterable[EncoderOnlyMotionModule] = nn.ModuleList([])
self.down_blocks: list[EncoderOnlyMotionModule] = nn.ModuleList([])
self.up_blocks = None
self.mid_block = None
# fill out down/up blocks and middle block, if present