diff --git a/animatediff/context.py b/animatediff/context.py index e87513c..accc207 100644 --- a/animatediff/context.py +++ b/animatediff/context.py @@ -12,7 +12,8 @@ class ContextFuseMethod: PYRAMID = "pyramid" RELATIVE = "relative" - LIST = [PYRAMID, RELATIVE, FLAT] + LIST = [PYRAMID, FLAT] + LIST_STATIC = [PYRAMID, RELATIVE, FLAT] class ContextType: @@ -332,6 +333,7 @@ def create_weights_pyramid(length: int, **kwargs) -> list[float]: FUSE_MAPPING = { ContextFuseMethod.FLAT: create_weights_flat, ContextFuseMethod.PYRAMID: create_weights_pyramid, + ContextFuseMethod.RELATIVE: create_weights_pyramid, } diff --git a/animatediff/motion_module_ad.py b/animatediff/motion_module_ad.py index 8d582c8..c1b489e 100644 --- a/animatediff/motion_module_ad.py +++ b/animatediff/motion_module_ad.py @@ -15,7 +15,7 @@ from comfy.utils import repeat_to_batch_size import comfy.ops import comfy.model_management -from .context import ContextOptions, get_context_weights, get_context_windows +from .context import ContextFuseMethod, ContextOptions, get_context_weights, get_context_windows from .utils_motion import CrossAttentionMM, MotionCompatibilityError, extend_to_batch_size, prepare_mask_batch from .utils_model import BetaSchedules, ModelTypeSD from .logger import logger @@ -813,11 +813,9 @@ 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 for sub_idxs in views: - 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) - sub_hidden_states = rearrange(hidden_states[:, sub_idxs], "b f d c -> (b f) d c") for attention_block, norm in zip(self.attention_blocks, self.norms): norm_hidden_states = norm(sub_hidden_states).to(sub_hidden_states.dtype) @@ -834,14 +832,35 @@ 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 count_final[:, sub_idxs] += weights_tensor - # get weighted average of sub_hidden_states + # 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 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 diff --git a/animatediff/nodes_context.py b/animatediff/nodes_context.py index f56dd9b..fd86d0f 100644 --- a/animatediff/nodes_context.py +++ b/animatediff/nodes_context.py @@ -145,7 +145,7 @@ class StandardStaticContextOptionsNode: "context_overlap": ("INT", {"default": 4, "min": 0, "max": OVERLAP_MAX}), }, "optional": { - "fuse_method": (ContextFuseMethod.LIST,), + "fuse_method": (ContextFuseMethod.LIST_STATIC,), "use_on_equal_length": ("BOOLEAN", {"default": False},), "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), "guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}), diff --git a/animatediff/nodes_deprecated.py b/animatediff/nodes_deprecated.py index fb5b072..ecf88a5 100644 --- a/animatediff/nodes_deprecated.py +++ b/animatediff/nodes_deprecated.py @@ -12,7 +12,7 @@ from PIL.PngImagePlugin import PngInfo import folder_paths from comfy.model_patcher import ModelPatcher -from .context import ContextSchedules, ContextOptions +from .context import ContextOptionsGroup, ContextOptions, ContextSchedules from .logger import logger from .utils_model import Folders, BetaSchedules, get_available_motion_models from .model_injection import ModelPatcherAndInjector, InjectionParams, MotionModelGroup, load_motion_module_gen1 @@ -109,16 +109,18 @@ class AnimateDiffLoaderAdvanced_Deprecated: model_name=model_name, apply_v2_properly=False, ) - # set context settings - params.set_context( + context_group = ContextOptionsGroup() + context_group.add( ContextOptions( context_length=context_length, context_stride=context_stride, context_overlap=context_overlap, context_schedule=context_schedule, closed_loop=closed_loop, + ) ) - ) + # set context settings + params.set_context(context_options=context_group) # inject for use in sampling code model = ModelPatcherAndInjector(model) model.motion_models = MotionModelGroup(motion_model) diff --git a/animatediff/sampling.py b/animatediff/sampling.py index e265f47..66b36f1 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -492,7 +492,7 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep, bias_total = bias_final[(full_length*n)+idx] prev_weight = (bias_total / (bias_total + bias)) new_weight = (bias / (bias_total + bias)) - cond_final[(full_length*n)+idx] = cond_final[(full_length*n)+idx] * prev_weight + sub_cond_out[(full_length*n)+pos] * new_weight + cond_final[(full_length*n)+idx] = cond_final[(full_length*n)+idx] * prev_weight + sub_cond_out[(full_length*n)+pos] * new_weight uncond_final[(full_length*n)+idx] = uncond_final[(full_length*n)+idx] * prev_weight + sub_uncond_out[(full_length*n)+pos] * new_weight bias_final[(full_length*n)+idx] = bias_total + bias else: