From dcfbee1454a18b2f57c2f59c7e889b8d59878755 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sat, 13 Jan 2024 10:28:59 -0600 Subject: [PATCH] Views, start of Context scheduling, use_on_equal_length for Contexts, --- __init__.py | 2 +- animatediff/context.py | 340 ++++++++---------- animatediff/model_injection.py | 19 +- animatediff/motion_module_ad.py | 105 ++++-- animatediff/nodes.py | 47 ++- animatediff/nodes_context.py | 235 +++++++++++- animatediff/nodes_deprecated.py | 2 +- animatediff/nodes_extras.py | 2 +- animatediff/nodes_gen1.py | 8 +- animatediff/nodes_gen2.py | 8 +- animatediff/nodes_multival.py | 2 +- animatediff/nodes_sample.py | 2 +- animatediff/sample_settings.py | 12 +- animatediff/sampling.py | 55 +-- .../{model_utils.py => utils_model.py} | 0 .../{motion_utils.py => utils_motion.py} | 22 +- 16 files changed, 567 insertions(+), 294 deletions(-) rename animatediff/{model_utils.py => utils_model.py} (100%) rename animatediff/{motion_utils.py => utils_motion.py} (89%) diff --git a/__init__.py b/__init__.py index c184225..a5635e3 100644 --- a/__init__.py +++ b/__init__.py @@ -1,6 +1,6 @@ import folder_paths from .animatediff.logger import logger -from .animatediff.model_utils import get_available_motion_models, Folders +from .animatediff.utils_model import get_available_motion_models, Folders from .animatediff.nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS if len(get_available_motion_models()) == 0: diff --git a/animatediff/context.py b/animatediff/context.py index c94560e..3e388d2 100644 --- a/animatediff/context.py +++ b/animatediff/context.py @@ -1,7 +1,9 @@ -from typing import Callable, Optional +from typing import Callable, Optional, Union import numpy as np +from .utils_motion import get_sorted_list_via_attr + class ContextFuseMethod: FLAT = "flat" PYRAMID = "pyramid" @@ -15,22 +17,128 @@ class ContextType: class ContextOptions: def __init__(self, context_length: int=None, context_stride: int=None, context_overlap: int=None, - context_schedule: str=None, closed_loop: bool=False, fuse_method: str=ContextFuseMethod.FLAT): + context_schedule: str=None, closed_loop: bool=False, fuse_method: str=ContextFuseMethod.FLAT, + use_on_equal_length: bool=False, view_options: 'ContextOptions'=None, + start_percent=0.0, guarantee_steps=1): + # permanent settings self.context_length = context_length self.context_stride = context_stride self.context_overlap = context_overlap self.context_schedule = context_schedule self.closed_loop = closed_loop self.fuse_method = fuse_method - self.sync_context_to_pe = False + self.sync_context_to_pe = False # this feature is likely bad and stay unused, so I might remove this + self.use_on_equal_length = use_on_equal_length + self.view_options = view_options.clone() if view_options else view_options + # scheduling + self.start_percent = float(start_percent) + self.guarantee_steps = guarantee_steps + # temporary vars + self._step: int = 0 + @property + def step(self): + return self._step + @step.setter + def step(self, value: int): + self._step = value + if self.view_options: + self.view_options.step = value + def clone(self): n = ContextOptions(context_length=self.context_length, context_stride=self.context_stride, context_overlap=self.context_overlap, context_schedule=self.context_schedule, - closed_loop=self.closed_loop, fuse_method=self.fuse_method) + closed_loop=self.closed_loop, fuse_method=self.fuse_method, + use_on_equal_length=self.use_on_equal_length, view_options=self.view_options, + start_percent=self.start_percent, guarantee_steps=self.guarantee_steps) return n +class ContextOptionsGroup: + def __init__(self): + self.contexts: list[ContextOptions] = [] + self._current_context: ContextOptions = None + self._current_guaranteed_steps: int = 0 + self.step = 0 + + def reset(self): + self._current_context: ContextOptions = None + self._current_guaranteed_steps: int = 0 + self.step = 0 + self._set_first_as_current() + + @classmethod + def default(cls): + def_context = ContextOptions() + new_group = ContextOptionsGroup() + new_group.add(def_context) + return new_group + + def add(self, context: ContextOptions): + # add to end of list, then sort + self.contexts.append(context) + self.contexts = get_sorted_list_via_attr(self.contexts, "start_percent") + self._set_first_as_current() + + def add_to_start(self, context: ContextOptions): + # add to start of list, then sort + self.contexts.insert(0, context) + self.contexts = get_sorted_list_via_attr(self.contexts, "start_percent") + self._set_first_as_current() + + def is_empty(self) -> bool: + return len(self.contexts) == 0 + + def clone(self): + cloned = ContextOptionsGroup() + for context in self.contexts: + cloned.contexts.append(context) + cloned._set_first_as_current() + return cloned + + def update_current_context(self, t: float): + self._current_context = self.contexts[0] + # based on t + current_steps, determine which context to use + pass + + def _set_first_as_current(self): + if len(self.contexts) > 0: + self._current_context = self.contexts[0] + + # properties shadow those of ContextOptions + @property + def context_length(self): + return self._current_context.context_length + + @property + def context_overlap(self): + return self._current_context.context_overlap + + @property + def context_stride(self): + return self._current_context.context_stride + + @property + def context_schedule(self): + return self._current_context.context_schedule + + @property + def closed_loop(self): + return self._current_context.closed_loop + + @property + def fuse_method(self): + return self._current_context.fuse_method + + @property + def use_on_equal_length(self): + return self._current_context.use_on_equal_length + + @property + def view_options(self): + return self._current_context.view_options + + class ContextSchedules: UNIFORM_LOOPED = "uniform" UNIFORM_STANDARD = "uniform_standard" @@ -39,23 +147,25 @@ class ContextSchedules: BATCHED = "batched" - UNIFORM_SCHEDULE_LIST = [UNIFORM_LOOPED] # only include somewhat functional contexts here + VIEW_AS_CONTEXT = "view_as_context" + + UNIFORM_SCHEDULE_LIST = [UNIFORM_LOOPED] STATIC_SCHEDULE_LIST = [STATIC_STANDARD] # from https://github.com/neggles/animatediff-cli/blob/main/src/animatediff/pipelines/context.py -def create_windows_uniform_looped(step: int, num_frames: int, opts: ContextOptions): +def create_windows_uniform_looped(num_frames: int, opts: Union[ContextOptionsGroup, ContextOptions]): windows = [] - if num_frames <= opts.context_length: + if num_frames < opts.context_length: windows.append(list(range(num_frames))) return windows context_stride = min(opts.context_stride, int(np.ceil(np.log2(num_frames / opts.context_length))) + 1) # obtain uniform windows as normal, looping and all for context_step in 1 << np.arange(context_stride): - pad = int(round(num_frames * ordered_halving(step))) + pad = int(round(num_frames * ordered_halving(opts.step))) for j in range( - int(ordered_halving(step) * context_step) + pad, + int(ordered_halving(opts.step) * context_step) + pad, num_frames + pad + (0 if opts.closed_loop else -opts.context_overlap), (opts.context_length * context_step - opts.context_overlap), ): @@ -64,7 +174,7 @@ def create_windows_uniform_looped(step: int, num_frames: int, opts: ContextOptio return windows -def create_windows_uniform_standard(step: int, num_frames: int, opts: ContextOptions): +def create_windows_uniform_standard(num_frames: int, opts: Union[ContextOptionsGroup, ContextOptions]): # unlike looped, uniform_straight does NOT allow windows that loop back to the beginning; # instead, they get shifted to the corresponding end of the frames. # in the case that a window (shifted or not) is identical to the previous one, it gets skipped. @@ -76,9 +186,9 @@ def create_windows_uniform_standard(step: int, num_frames: int, opts: ContextOpt context_stride = min(opts.context_stride, int(np.ceil(np.log2(num_frames / opts.context_length))) + 1) # first, obtain uniform windows as normal, looping and all for context_step in 1 << np.arange(context_stride): - pad = int(round(num_frames * ordered_halving(step))) + pad = int(round(num_frames * ordered_halving(opts.step))) for j in range( - int(ordered_halving(step) * context_step) + pad, + int(ordered_halving(opts.step) * context_step) + pad, num_frames + pad + (-opts.context_overlap), (opts.context_length * context_step - opts.context_overlap), ): @@ -104,7 +214,7 @@ def create_windows_uniform_standard(step: int, num_frames: int, opts: ContextOpt break win_i += 1 - # reverse delete_idxs so that they will be deleted in an order that does break idx correlation + # reverse delete_idxs so that they will be deleted in an order that doesn't break idx correlation delete_idxs.reverse() for i in delete_idxs: windows.pop(i) @@ -112,7 +222,7 @@ def create_windows_uniform_standard(step: int, num_frames: int, opts: ContextOpt return windows -def create_windows_static_standard(step: int, num_frames: int, opts: ContextOptions): +def create_windows_static_standard(num_frames: int, opts: Union[ContextOptionsGroup, ContextOptions]): windows = [] if num_frames <= opts.context_length: windows.append(list(range(num_frames))) @@ -131,7 +241,7 @@ def create_windows_static_standard(step: int, num_frames: int, opts: ContextOpti return windows -def create_windows_batched(step: int, num_frames: int, opts: ContextOptions): +def create_windows_batched(num_frames: int, opts: Union[ContextOptionsGroup, ContextOptions]): windows = [] if num_frames <= opts.context_length: windows.append(list(range(num_frames))) @@ -144,11 +254,15 @@ def create_windows_batched(step: int, num_frames: int, opts: ContextOptions): return windows -def get_context_windows(step: int, num_frames: int, opts: ContextOptions): +def create_windows_default(num_frames: int, opts: Union[ContextOptionsGroup, ContextOptions]): + return [list(range(num_frames))] + + +def get_context_windows(num_frames: int, opts: Union[ContextOptionsGroup, ContextOptions]): context_func = CONTEXT_MAPPING.get(opts.context_schedule, None) if not context_func: - raise ValueError(f"Unknown context_schedule '{opts.context_schedule}'") - return context_func(step, num_frames, opts) + raise ValueError(f"Unknown context_schedule '{opts.context_schedule}'.") + return context_func(num_frames, opts) CONTEXT_MAPPING = { @@ -156,19 +270,40 @@ CONTEXT_MAPPING = { ContextSchedules.UNIFORM_STANDARD: create_windows_uniform_standard, ContextSchedules.STATIC_STANDARD: create_windows_static_standard, ContextSchedules.BATCHED: create_windows_batched, + ContextSchedules.VIEW_AS_CONTEXT: create_windows_default, # just return all to allow Views to do all the work } -def generate_distance_weight(n): - if n % 2 == 0: - max_weight = n // 2 +def get_context_weights(num_frames: int, fuse_method: str): + weights_func = FUSE_MAPPING.get(fuse_method, None) + if not weights_func: + raise ValueError(f"Unknown fuse_method '{fuse_method}'.") + return weights_func(num_frames) + + +def create_weights_flat(length: int, **kwargs) -> list[float]: + # weight is the same for all + return [1.0] * length + + +def create_weights_pyramid(length: int, **kwargs) -> list[float]: + # weight is based on the distance away from the edge of the context window; + # based on weighted average concept in FreeNoise paper + if length % 2 == 0: + max_weight = length // 2 weight_sequence = list(range(1, max_weight + 1, 1)) + list(range(max_weight, 0, -1)) else: - max_weight = (n + 1) // 2 + max_weight = (length + 1) // 2 weight_sequence = list(range(1, max_weight, 1)) + [max_weight] + list(range(max_weight - 1, 0, -1)) return weight_sequence +FUSE_MAPPING = { + ContextFuseMethod.FLAT: create_weights_flat, + ContextFuseMethod.PYRAMID: create_weights_pyramid, +} + + # Returns fraction that has denominator that is a power of 2 def ordered_halving(val): # get binary value, padded with 0s for 64 bits @@ -219,164 +354,3 @@ def shift_window_to_end(window: list[int], num_frames: int): for i in range(len(window)): # 2) add end_delta to each val to slide windows to end window[i] = window[i] + end_delta - - - - - - -################################################################################################ - -# Generator that returns lists of latent indeces to diffuse on -def uniform( - step: int, - num_frames: int, - opts: ContextOptions, - print_final: bool = False, -): - if num_frames <= opts.context_length: - yield list(range(num_frames)) - return - - context_stride = min(opts.context_stride, int(np.ceil(np.log2(num_frames / opts.context_length))) + 1) - - for context_step in 1 << np.arange(context_stride): - pad = int(round(num_frames * ordered_halving(step, print_final))) - for j in range( - int(ordered_halving(step) * context_step) + pad, - num_frames + pad + (0 if opts.closed_loop else -opts.context_overlap), - (opts.context_length * context_step - opts.context_overlap), - ): - yield [e % num_frames for e in range(j, j + opts.context_length * context_step, context_step)] - - -################################# -# helper funcs for testing -def get_total_steps( - scheduler, - timesteps: list[int], - num_steps: Optional[int] = None, - num_frames: int = ..., - context_size: Optional[int] = None, - context_stride: int = 3, - context_overlap: int = 4, - closed_loop: bool = True, -): - return sum( - len( - list( - scheduler( - i, - num_steps, - num_frames, - context_size, - context_stride, - context_overlap, - ) - ) - ) - for i in range(len(timesteps)) - ) - - -def get_total_steps_fixed( - scheduler, - timesteps: list[int], - num_steps: Optional[int] = None, - num_frames: int = ..., - context_size: Optional[int] = None, - context_stride: int = 3, - context_overlap: int = 4, - closed_loop: bool = True, -): - total_loops = 0 - for i, t in enumerate(timesteps): - for context in scheduler(i, num_steps, num_frames, context_size, context_stride, context_overlap, closed_loop=closed_loop): - total_loops += 1 - return total_loops - - -def uniform_v2( - step: int = ..., - num_frames: int = ..., - context_size: Optional[int] = None, - context_stride: int = 3, - context_overlap: int = 4, - closed_loop: bool = True, - print_final: bool = False, -): - if num_frames <= context_size: - yield list(range(num_frames)) - return - - context_stride = min(context_stride, int(np.ceil(np.log2(num_frames / context_size))) + 1) - - pad = int(round(num_frames * ordered_halving(step, print_final))) - for context_step in 1 << np.arange(context_stride): - j_initial = int(ordered_halving(step) * context_step) + pad - for j in range( - j_initial, - num_frames + pad - context_overlap, - (context_size * context_step - context_overlap), - ): - if context_size * context_step > num_frames: - # On the final context_step, - # ensure no frame appears in the window twice - yield [e % num_frames for e in range(j, j + num_frames, context_step)] - continue - j = j % num_frames - if j > (j + context_size * context_step) % num_frames and not closed_loop: - yield [e for e in range(j, num_frames, context_step)] - j_stop = (j + context_size * context_step) % num_frames - # When ((num_frames % (context_size - context_overlap)+context_overlap) % context_size != 0, - # This can cause 'superflous' runs where all frames in - # a context window have already been processed during - # the first context window of this stride and step. - # While the following commented if should prevent this, - # I believe leaving it in is more correct as it maintains - # the total conditional passes per frame over a large total steps - # if j_stop > context_overlap: - yield [e for e in range(0, j_stop, context_step)] - continue - yield [e % num_frames for e in range(j, j + context_size * context_step, context_step)] - - -def uniform_constant( - step: int = ..., - num_frames: int = ..., - context_size: Optional[int] = None, - context_stride: int = 3, - context_overlap: int = 4, - closed_loop: bool = True, - print_final: bool = False, -): - if num_frames <= context_size: - yield list(range(num_frames)) - return - - context_stride = min(context_stride, int(np.ceil(np.log2(num_frames / context_size))) + 1) - - # want to avoid loops that connect end to beginning - - for context_step in 1 << np.arange(context_stride): - pad = int(round(num_frames * ordered_halving(step, print_final))) - for j in range( - int(ordered_halving(step) * context_step) + pad, - num_frames + pad + (0 if closed_loop else -context_overlap), - (context_size * context_step - context_overlap), - ): - skip_this_window = False - prev_val = -1 - to_yield = [] - for e in range(j, j + context_size * context_step, context_step): - e = e % num_frames - # if not a closed loop and loops back on itself, should be skipped - if not closed_loop and e < prev_val: - skip_this_window = True - break - to_yield.append(e) - prev_val = e - if skip_this_window: - continue - # yield if not skipped - yield to_yield diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index f7af9f1..ec9bad1 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -11,12 +11,12 @@ import comfy.utils from comfy.model_patcher import ModelPatcher from comfy.model_base import BaseModel -from .context import ContextOptions, ContextOptions +from .context import ContextOptions, ContextOptions, ContextOptionsGroup from .motion_module_ad import AnimateDiffModel, has_mid_block, normalize_ad_state_dict from .logger import logger -from .motion_utils import ADKeyframe, ADKeyframeGroup, MotionCompatibilityError, get_combined_multival, normalize_min_max +from .utils_motion import ADKeyframe, ADKeyframeGroup, MotionCompatibilityError, get_combined_multival, normalize_min_max from .motion_lora import MotionLoraInfo, MotionLoraList -from .model_utils import get_motion_lora_path, get_motion_model_path, get_sd_model_type +from .utils_model import get_motion_lora_path, get_motion_model_path, get_sd_model_type from .sample_settings import SampleSettings, SeedNoiseGeneration @@ -185,7 +185,6 @@ class MotionModelPatcher(ModelPatcher): self.model.set_effect(self.combined_effect) self.was_within_range = True - def cleanup(self): if self.model is not None: self.model.cleanup() @@ -245,6 +244,10 @@ class MotionModelGroup: for motion_model in self.models: motion_model.model.set_sub_idxs(sub_idxs=sub_idxs) + def set_view_options(self, view_options: ContextOptions): + for motion_model in self.models: + motion_model.model.set_view_options(view_options) + def set_video_length(self, video_length: int, full_length: int): for motion_model in self.models: motion_model.model.set_video_length(video_length=video_length, full_length=full_length) @@ -500,15 +503,15 @@ class InjectionParams: self.apply_mm_groupnorm_hack = apply_mm_groupnorm_hack self.model_name = model_name self.apply_v2_properly = apply_v2_properly - self.context_options: ContextOptions = ContextOptions() + self.context_options: ContextOptionsGroup = ContextOptionsGroup.default() self.motion_model_settings = MotionModelSettings() # Gen1 self.sub_idxs = None # value should NOT be included in clone, so it will auto reset def set_noise_extra_args(self, noise_extra_args: dict): noise_extra_args["context_options"] = self.context_options.clone() - def set_context(self, context_options: ContextOptions): - self.context_options = context_options.clone() if context_options else ContextOptions() + def set_context(self, context_options: ContextOptionsGroup): + self.context_options = context_options.clone() if context_options else ContextOptionsGroup.default() def is_using_sliding_context(self) -> bool: return self.context_options.context_length is not None @@ -520,7 +523,7 @@ class InjectionParams: self.motion_model_settings = motion_model_settings def reset_context(self): - self.context_options = ContextOptions() + self.context_options = ContextOptionsGroup.default() def clone(self) -> 'InjectionParams': new_params = InjectionParams( diff --git a/animatediff/motion_module_ad.py b/animatediff/motion_module_ad.py index 2f0969f..11b8eaf 100644 --- a/animatediff/motion_module_ad.py +++ b/animatediff/motion_module_ad.py @@ -13,8 +13,9 @@ from comfy.ldm.modules.diffusionmodules.openaimodel import SpatialTransformer from comfy.controlnet import broadcast_image_to from comfy.utils import repeat_to_batch_size -from .motion_utils import GroupNormAD, CrossAttentionMM, MotionCompatibilityError, extend_to_batch_size, prepare_mask_batch -from .model_utils import ModelTypeSD +from .context import ContextOptions, get_context_weights, get_context_windows +from .utils_motion import GroupNormAD, CrossAttentionMM, MotionCompatibilityError, extend_to_batch_size, prepare_mask_batch +from .utils_model import ModelTypeSD from .logger import logger @@ -281,6 +282,14 @@ class AnimateDiffModel(nn.Module): if self.mid_block is not None: self.mid_block.set_sub_idxs(sub_idxs) + def set_view_options(self, view_options: ContextOptions): + for block in self.down_blocks: + block.set_view_options(view_options) + for block in self.up_blocks: + block.set_view_options(view_options) + if self.mid_block is not None: + self.mid_block.set_view_options(view_options) + def reset(self): self._reset_sub_idxs() self._reset_scale_multiplier() @@ -355,6 +364,10 @@ class MotionModule(nn.Module): for motion_module in self.motion_modules: motion_module.set_sub_idxs(sub_idxs) + def set_view_options(self, view_options: ContextOptions): + for motion_module in self.motion_modules: + motion_module.set_view_options(view_options=view_options) + def reset_temp_vars(self): for motion_module in self.motion_modules: motion_module.reset_temp_vars() @@ -382,6 +395,7 @@ class VanillaTemporalModule(nn.Module): self.video_length = 16 self.full_length = 16 self.sub_idxs = None + self.view_options = None self.effect = None self.temp_effect_mask: Tensor = None @@ -429,8 +443,12 @@ class VanillaTemporalModule(nn.Module): self.sub_idxs = sub_idxs self.temporal_transformer.set_sub_idxs(sub_idxs) + def set_view_options(self, view_options: ContextOptions): + self.view_options = view_options + def reset_temp_vars(self): self.set_effect(None) + self.set_view_options(None) self.temporal_transformer.reset_temp_vars() def get_effect_mask(self, input_tensor: Tensor): @@ -459,7 +477,7 @@ class VanillaTemporalModule(nn.Module): def forward(self, input_tensor: Tensor, encoder_hidden_states=None, attention_mask=None): if self.effect is None: - return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask) + return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask, self.view_options) # return weighted average of input_tensor and AD output if type(self.effect) != Tensor: effect = self.effect @@ -468,7 +486,7 @@ class VanillaTemporalModule(nn.Module): return input_tensor else: effect = self.get_effect_mask(input_tensor) - return input_tensor*(1.0-effect) + self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask)*effect + return input_tensor*(1.0-effect) + self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask, self.view_options)*effect class TemporalTransformer3DModel(nn.Module): @@ -595,7 +613,7 @@ class TemporalTransformer3DModel(nn.Module): return self.temp_scale_mask[:, self.sub_idxs, :] return self.temp_scale_mask - def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None): + def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, view_options: ContextOptions=None): batch, channel, height, width = hidden_states.shape residual = hidden_states scale_mask = self.get_scale_mask(hidden_states) @@ -614,7 +632,8 @@ class TemporalTransformer3DModel(nn.Module): encoder_hidden_states=encoder_hidden_states, attention_mask=attention_mask, video_length=self.video_length, - scale_mask=scale_mask + scale_mask=scale_mask, + view_options=view_options ) # output @@ -691,26 +710,64 @@ class TemporalTransformerBlock(nn.Module): def forward( self, - hidden_states, - encoder_hidden_states=None, - attention_mask=None, - video_length=None, - scale_mask=None + hidden_states: Tensor, + encoder_hidden_states: Tensor=None, + attention_mask: Tensor=None, + video_length: int=None, + scale_mask: Tensor=None, + view_options: ContextOptions=None, ): - for attention_block, norm in zip(self.attention_blocks, self.norms): - norm_hidden_states = norm(hidden_states).to(hidden_states.dtype) - hidden_states = ( - attention_block( - norm_hidden_states, - encoder_hidden_states=encoder_hidden_states - if attention_block.is_cross_attention - else None, - attention_mask=attention_mask, - video_length=video_length, - scale_mask=scale_mask + if not view_options: + for attention_block, norm in zip(self.attention_blocks, self.norms): + norm_hidden_states = norm(hidden_states).to(hidden_states.dtype) + hidden_states = ( + attention_block( + norm_hidden_states, + encoder_hidden_states=encoder_hidden_states + if attention_block.is_cross_attention + else None, + attention_mask=attention_mask, + video_length=video_length, + scale_mask=scale_mask + ) + hidden_states ) - + hidden_states - ) + else: + # 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) + value_final = torch.zeros_like(hidden_states) + count_final = torch.zeros_like(hidden_states) + 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) + sub_hidden_states = ( + attention_block( + norm_hidden_states, + encoder_hidden_states=encoder_hidden_states # do these need to be changed for sub_idxs too? + if attention_block.is_cross_attention + else None, + attention_mask=attention_mask, + video_length=len(sub_idxs), + scale_mask=scale_mask[:, sub_idxs, :] if scale_mask is not None else scale_mask + ) + sub_hidden_states + ) + sub_hidden_states = rearrange(sub_hidden_states, "(b f) d c -> b f d c", f=len(sub_idxs)) + + value_final[:, sub_idxs] += sub_hidden_states * weights_tensor + count_final[:, sub_idxs] += weights_tensor + + # 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 hidden_states = self.ff(self.ff_norm(hidden_states)) + hidden_states diff --git a/animatediff/nodes.py b/animatediff/nodes.py index ee57eb4..fa4ee1e 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -5,7 +5,7 @@ import comfy.sample as comfy_sample from comfy.model_patcher import ModelPatcher from .logger import logger -from .model_utils import BetaSchedules, get_available_motion_loras, get_available_motion_models, get_motion_lora_path +from .utils_model import BetaSchedules, get_available_motion_loras, get_available_motion_models, get_motion_lora_path from .motion_lora import MotionLoraInfo, MotionLoraList from .model_injection import InjectionParams, ModelPatcherAndInjector, MotionModelSettings, load_motion_module from .sample_settings import SampleSettings, SeedNoiseGeneration @@ -15,7 +15,8 @@ from .nodes_gen1 import AnimateDiffLoaderWithContext from .nodes_gen2 import UseEvolvedSamplingNode, ApplyAnimateDiffModelNode, ApplyAnimateDiffModelBasicNode, LoadAnimateDiffModelNode, ADKeyframeNode from .nodes_multival import MultivalDynamicNode, MultivalFloatNode, MultivalScaledMaskNode from .nodes_sample import FreeInitOptionsNode, NoiseLayerAddWeightedNode, SampleSettingsNode, NoiseLayerAddNode, NoiseLayerReplaceNode, IterationOptionsNode -from .nodes_context import LoopedUniformContextOptionsNode, StandardUniformContextOptionsNode, StandardStaticContextOptionsNode, BatchedContextOptionsNode +from .nodes_context import (LoopedUniformContextOptionsNode, LoopedUniformViewOptionsNode, StandardUniformContextOptionsNode, StandardStaticContextOptionsNode, BatchedContextOptionsNode, + StandardStaticViewOptionsNode, StandardUniformViewOptionsNode, ViewAsContextOptionsNode) from .nodes_extras import AnimateDiffUnload, EmptyLatentImageLarge, CheckpointLoaderSimpleWithNoiseSelect from .nodes_experimental import AnimateDiffModelSettingsSimple, AnimateDiffModelSettingsAdvanced, AnimateDiffModelSettingsAdvancedAttnStrengths from .nodes_deprecated import AnimateDiffLoader_Deprecated, AnimateDiffLoaderAdvanced_Deprecated, AnimateDiffCombine_Deprecated @@ -97,18 +98,23 @@ NODE_CLASS_MAPPINGS = { # Multival Nodes "ADE_MultivalDynamic": MultivalDynamicNode, "ADE_MultivalScaledMask": MultivalScaledMaskNode, + # Context Opts + "ADE_StandardStaticContextOptions": StandardStaticContextOptionsNode, + "ADE_StandardUniformContextOptions": StandardUniformContextOptionsNode, + "ADE_AnimateDiffUniformContextOptions": LoopedUniformContextOptionsNode, + "ADE_ViewsOnlyContextOptions": ViewAsContextOptionsNode, + "ADE_BatchedContextOptions": BatchedContextOptionsNode, + # View Opts + "ADE_StandardStaticViewOptions": StandardStaticViewOptionsNode, + "ADE_StandardUniformViewOptions": StandardUniformViewOptionsNode, + "ADE_LoopedUniformViewOptions": LoopedUniformViewOptionsNode, + # Iteration Opts + "ADE_IterationOptsDefault": IterationOptionsNode, + "ADE_IterationOptsFreeInit": FreeInitOptionsNode, # Noise Layer Nodes "ADE_NoiseLayerAdd": NoiseLayerAddNode, "ADE_NoiseLayerAddWeighted": NoiseLayerAddWeightedNode, "ADE_NoiseLayerReplace": NoiseLayerReplaceNode, - # Context Opts - "ADE_AnimateDiffUniformContextOptions": LoopedUniformContextOptionsNode, - "ADE_StandardUniformContextOptions": StandardUniformContextOptionsNode, - "ADE_StandardStaticContextOptions": StandardStaticContextOptionsNode, - "ADE_BatchedContextOptions": BatchedContextOptionsNode, - # Iteration Opts - "ADE_IterationOptsDefault": IterationOptionsNode, - "ADE_IterationOptsFreeInit": FreeInitOptionsNode, # Extras Nodes "ADE_AnimateDiffUnload": AnimateDiffUnload, "ADE_EmptyLatentImageLarge": EmptyLatentImageLarge, @@ -136,18 +142,23 @@ NODE_DISPLAY_NAME_MAPPINGS = { # Multival Nodes "ADE_MultivalDynamic": "Multival Dynamic πŸŽ­πŸ…πŸ…“", "ADE_MultivalScaledMask": "Multival Scaled Mask πŸŽ­πŸ…πŸ…“", + # Context Opts + "ADE_StandardStaticContextOptions": "Context Optionsβ—†Standard Static πŸŽ­πŸ…πŸ…“", + "ADE_StandardUniformContextOptions": "Context Optionsβ—†Standard Uniform πŸŽ­πŸ…πŸ…“", + "ADE_AnimateDiffUniformContextOptions": "Context Optionsβ—†Looped Uniform πŸŽ­πŸ…πŸ…“", + "ADE_ViewsOnlyContextOptions": "Context Optionsβ—†Views Only [VRAMβ‡ˆ] πŸŽ­πŸ…πŸ…“", + "ADE_BatchedContextOptions": "Context Optionsβ—†Batched [Non-AD] πŸŽ­πŸ…πŸ…“", + # View Opts + "ADE_StandardStaticViewOptions": "View Optionsβ—†Standard Static πŸŽ­πŸ…πŸ…“", + "ADE_StandardUniformViewOptions": "View Optionsβ—†Standard Uniform πŸŽ­πŸ…πŸ…“", + "ADE_LoopedUniformViewOptions": "View Optionsβ—†Looped Uniform πŸŽ­πŸ…πŸ…“", + # Iteration Opts + "ADE_IterationOptsDefault": "Default Iteration Options πŸŽ­πŸ…πŸ…“", + "ADE_IterationOptsFreeInit": "FreeInit Iteration Options πŸŽ­πŸ…πŸ…“", # Noise Layer Nodes "ADE_NoiseLayerAdd": "Noise Layer [Add] πŸŽ­πŸ…πŸ…“", "ADE_NoiseLayerAddWeighted": "Noise Layer [Add Weighted] πŸŽ­πŸ…πŸ…“", "ADE_NoiseLayerReplace": "Noise Layer [Replace] πŸŽ­πŸ…πŸ…“", - # Context Opts - "ADE_AnimateDiffUniformContextOptions": "Looped Uniform Context Options πŸŽ­πŸ…πŸ…“", - "ADE_StandardUniformContextOptions": "Standard Uniform Context Options πŸŽ­πŸ…πŸ…“", - "ADE_StandardStaticContextOptions": "Standard Static Context Options πŸŽ­πŸ…πŸ…“", - "ADE_BatchedContextOptions": "[Non-AD] Batched Context Options πŸŽ­πŸ…πŸ…“", - # Iteration Opts - "ADE_IterationOptsDefault": "Default Iteration Options πŸŽ­πŸ…πŸ…“", - "ADE_IterationOptsFreeInit": "FreeInit Iteration Options πŸŽ­πŸ…πŸ…“", # Extras Nodes "ADE_AnimateDiffUnload": "AnimateDiff Unload πŸŽ­πŸ…πŸ…“", "ADE_EmptyLatentImageLarge": "Empty Latent Image (Big Batch) πŸŽ­πŸ…πŸ…“", diff --git a/animatediff/nodes_context.py b/animatediff/nodes_context.py index e38dfbf..e839c8f 100644 --- a/animatediff/nodes_context.py +++ b/animatediff/nodes_context.py @@ -1,4 +1,10 @@ -from .context import ContextFuseMethod, ContextOptions, ContextSchedules +from .context import ContextFuseMethod, ContextOptions, ContextOptionsGroup, ContextSchedules +from .utils_model import BIGMAX + + +LENGTH_MAX = 128 # keep an eye on these max values; +STRIDE_MAX = 32 # would need to be updated +OVERLAP_MAX = 128 # if new motion modules come out class LoopedUniformContextOptionsNode: @@ -6,24 +12,35 @@ class LoopedUniformContextOptionsNode: def INPUT_TYPES(s): return { "required": { - "context_length": ("INT", {"default": 16, "min": 1, "max": 128}), # keep an eye on these max values - "context_stride": ("INT", {"default": 1, "min": 1, "max": 32}), # would need to be updated - "context_overlap": ("INT", {"default": 4, "min": 0, "max": 128}), # if new motion modules come out + "context_length": ("INT", {"default": 16, "min": 1, "max": LENGTH_MAX}), + "context_stride": ("INT", {"default": 1, "min": 1, "max": STRIDE_MAX}), + "context_overlap": ("INT", {"default": 4, "min": 0, "max": OVERLAP_MAX}), "context_schedule": (ContextSchedules.UNIFORM_SCHEDULE_LIST,), "closed_loop": ("BOOLEAN", {"default": False},), #"sync_context_to_pe": ("BOOLEAN", {"default": False},), }, "optional": { "fuse_method": (ContextFuseMethod.LIST,), + "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}), + "prev_context": ("CONTEXT_OPTIONS",), + "view_opts": ("VIEW_OPTS",), } } RETURN_TYPES = ("CONTEXT_OPTIONS",) + RETURN_NAMES = ("CONTEXT_OPTS",) CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/context opts" FUNCTION = "create_options" def create_options(self, context_length: int, context_stride: int, context_overlap: int, context_schedule: int, closed_loop: bool, - fuse_method: str=ContextFuseMethod.FLAT): + fuse_method: str=ContextFuseMethod.FLAT, use_on_equal_length=False, start_percent: float=0.0, guarantee_steps: int=1, + view_opts: ContextOptions=None, prev_context: ContextOptionsGroup=None): + if prev_context is None: + prev_context = ContextOptionsGroup() + prev_context = prev_context.clone() + context_options = ContextOptions( context_length=context_length, context_stride=context_stride, @@ -31,9 +48,14 @@ class LoopedUniformContextOptionsNode: context_schedule=context_schedule, closed_loop=closed_loop, fuse_method=fuse_method, + use_on_equal_length=use_on_equal_length, + start_percent=start_percent, + guarantee_steps=guarantee_steps, + view_options=view_opts, ) #context_options.set_sync_context_to_pe(sync_context_to_pe) - return (context_options,) + prev_context.add(context_options) + return (prev_context,) class StandardUniformContextOptionsNode: @@ -41,21 +63,32 @@ class StandardUniformContextOptionsNode: def INPUT_TYPES(s): return { "required": { - "context_length": ("INT", {"default": 16, "min": 1, "max": 128}), # keep an eye on these max values - "context_stride": ("INT", {"default": 1, "min": 1, "max": 32}), # would need to be updated - "context_overlap": ("INT", {"default": 4, "min": 0, "max": 128}), # if new motion modules come out + "context_length": ("INT", {"default": 16, "min": 1, "max": LENGTH_MAX}), + "context_stride": ("INT", {"default": 1, "min": 1, "max": STRIDE_MAX}), + "context_overlap": ("INT", {"default": 4, "min": 0, "max": OVERLAP_MAX}), }, "optional": { "fuse_method": (ContextFuseMethod.LIST,), + "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}), + "prev_context": ("CONTEXT_OPTIONS",), + "view_opts": ("VIEW_OPTS",), } } RETURN_TYPES = ("CONTEXT_OPTIONS",) + RETURN_NAMES = ("CONTEXT_OPTS",) CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/context opts" FUNCTION = "create_options" def create_options(self, context_length: int, context_stride: int, context_overlap: int, - fuse_method: str=ContextFuseMethod.FLAT): + fuse_method: str=ContextFuseMethod.FLAT, use_on_equal_length=False, start_percent: float=0.0, guarantee_steps: int=1, + view_opts: ContextOptions=None, prev_context: ContextOptionsGroup=None): + if prev_context is None: + prev_context = ContextOptionsGroup() + prev_context = prev_context.clone() + context_options = ContextOptions( context_length=context_length, context_stride=context_stride, @@ -63,8 +96,13 @@ class StandardUniformContextOptionsNode: context_schedule=ContextSchedules.UNIFORM_STANDARD, closed_loop=False, fuse_method=fuse_method, + use_on_equal_length=use_on_equal_length, + start_percent=start_percent, + guarantee_steps=guarantee_steps, + view_options=view_opts, ) - return (context_options,) + prev_context.add(context_options) + return (prev_context,) class StandardStaticContextOptionsNode: @@ -72,28 +110,44 @@ class StandardStaticContextOptionsNode: def INPUT_TYPES(s): return { "required": { - "context_length": ("INT", {"default": 16, "min": 1, "max": 128}), # keep an eye on these max values - "context_overlap": ("INT", {"default": 4, "min": 0, "max": 128}), # if new motion modules come out + "context_length": ("INT", {"default": 16, "min": 1, "max": LENGTH_MAX}), + "context_overlap": ("INT", {"default": 4, "min": 0, "max": OVERLAP_MAX}), }, "optional": { "fuse_method": (ContextFuseMethod.LIST, {"default": ContextFuseMethod.PYRAMID}), + "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}), + "prev_context": ("CONTEXT_OPTIONS",), + "view_opts": ("VIEW_OPTS",), } } RETURN_TYPES = ("CONTEXT_OPTIONS",) + RETURN_NAMES = ("CONTEXT_OPTS",) CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/context opts" FUNCTION = "create_options" def create_options(self, context_length: int, context_overlap: int, - fuse_method: str=ContextFuseMethod.FLAT): + fuse_method: str=ContextFuseMethod.FLAT, use_on_equal_length=False, start_percent: float=0.0, guarantee_steps: int=1, + view_opts: ContextOptions=None, prev_context: ContextOptionsGroup=None): + if prev_context is None: + prev_context = ContextOptionsGroup() + prev_context = prev_context.clone() + context_options = ContextOptions( context_length=context_length, context_stride=None, context_overlap=context_overlap, context_schedule=ContextSchedules.STATIC_STANDARD, fuse_method=fuse_method, + use_on_equal_length=use_on_equal_length, + start_percent=start_percent, + guarantee_steps=guarantee_steps, + view_options=view_opts, ) - return (context_options,) + prev_context.add(context_options) + return (prev_context,) class BatchedContextOptionsNode: @@ -101,18 +155,163 @@ class BatchedContextOptionsNode: def INPUT_TYPES(s): return { "required": { - "context_length": ("INT", {"default": 16, "min": 1, "max": 128}), # keep an eye on these max values + "context_length": ("INT", {"default": 16, "min": 1, "max": LENGTH_MAX}), }, + "optional": { + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), + "guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}), + "prev_context": ("CONTEXT_OPTIONS",), + } } RETURN_TYPES = ("CONTEXT_OPTIONS",) + RETURN_NAMES = ("CONTEXT_OPTS",) CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/context opts" FUNCTION = "create_options" - def create_options(self, context_length: int): + def create_options(self, context_length: int, start_percent: float=0.0, guarantee_steps: int=1, + prev_context: ContextOptionsGroup=None): + if prev_context is None: + prev_context = ContextOptionsGroup() + prev_context = prev_context.clone() + context_options = ContextOptions( context_length=context_length, context_overlap=0, context_schedule=ContextSchedules.BATCHED, + start_percent=start_percent, + guarantee_steps=guarantee_steps, ) - return (context_options,) + prev_context.add(context_options) + return (prev_context,) + + +class ViewAsContextOptionsNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "view_opts_req": ("VIEW_OPTS",), + }, + "optional": { + "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}), + "prev_context": ("CONTEXT_OPTIONS",), + } + } + + RETURN_TYPES = ("CONTEXT_OPTIONS",) + RETURN_NAMES = ("CONTEXT_OPTS",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/context opts" + FUNCTION = "create_options" + + def create_options(self, view_opts_req: ContextOptions, start_percent: float=0.0, guarantee_steps: int=1, + prev_context: ContextOptionsGroup=None): + if prev_context is None: + prev_context = ContextOptionsGroup() + prev_context = prev_context.clone() + context_options = ContextOptions( + context_schedule=ContextSchedules.VIEW_AS_CONTEXT, + start_percent=start_percent, + guarantee_steps=guarantee_steps, + view_options=view_opts_req, + use_on_equal_length=True + ) + prev_context.add(context_options) + return (prev_context,) + + +######################### +# View Options +class StandardStaticViewOptionsNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "view_length": ("INT", {"default": 16, "min": 1, "max": LENGTH_MAX}), + "view_overlap": ("INT", {"default": 4, "min": 0, "max": OVERLAP_MAX}), + }, + "optional": { + "fuse_method": (ContextFuseMethod.LIST, {"default": ContextFuseMethod.PYRAMID}), + } + } + + RETURN_TYPES = ("VIEW_OPTS",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/context opts/view opts" + FUNCTION = "create_options" + + def create_options(self, view_length: int, view_overlap: int, + fuse_method: str=ContextFuseMethod.FLAT,): + view_options = ContextOptions( + context_length=view_length, + context_stride=None, + context_overlap=view_overlap, + context_schedule=ContextSchedules.STATIC_STANDARD, + fuse_method=fuse_method, + ) + return (view_options,) + + +class StandardUniformViewOptionsNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "view_length": ("INT", {"default": 16, "min": 1, "max": LENGTH_MAX}), + "view_stride": ("INT", {"default": 1, "min": 1, "max": STRIDE_MAX}), + "view_overlap": ("INT", {"default": 4, "min": 0, "max": OVERLAP_MAX}), + }, + "optional": { + "fuse_method": (ContextFuseMethod.LIST, {"default": ContextFuseMethod.PYRAMID}), + } + } + + RETURN_TYPES = ("VIEW_OPTS",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/context opts/view opts" + FUNCTION = "create_options" + + def create_options(self, view_length: int, view_overlap: int, view_stride: int, + fuse_method: str=ContextFuseMethod.FLAT,): + view_options = ContextOptions( + context_length=view_length, + context_stride=view_stride, + context_overlap=view_overlap, + context_schedule=ContextSchedules.UNIFORM_STANDARD, + fuse_method=fuse_method, + ) + return (view_options,) + + +class LoopedUniformViewOptionsNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "view_length": ("INT", {"default": 16, "min": 1, "max": LENGTH_MAX}), + "view_stride": ("INT", {"default": 1, "min": 1, "max": STRIDE_MAX}), + "view_overlap": ("INT", {"default": 4, "min": 0, "max": OVERLAP_MAX}), + "closed_loop": ("BOOLEAN", {"default": False},), + }, + "optional": { + "fuse_method": (ContextFuseMethod.LIST, {"default": ContextFuseMethod.PYRAMID}), + } + } + + RETURN_TYPES = ("VIEW_OPTS",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/context opts/view opts" + FUNCTION = "create_options" + + def create_options(self, view_length: int, view_overlap: int, view_stride: int, closed_loop: bool, + fuse_method: str=ContextFuseMethod.FLAT,): + view_options = ContextOptions( + context_length=view_length, + context_stride=view_stride, + context_overlap=view_overlap, + context_schedule=ContextSchedules.UNIFORM_LOOPED, + closed_loop=closed_loop, + fuse_method=fuse_method, + ) + return (view_options,) + + diff --git a/animatediff/nodes_deprecated.py b/animatediff/nodes_deprecated.py index 3428644..aa087da 100644 --- a/animatediff/nodes_deprecated.py +++ b/animatediff/nodes_deprecated.py @@ -14,7 +14,7 @@ from comfy.model_patcher import ModelPatcher from .context import ContextSchedules, ContextOptions from .logger import logger -from .model_utils import Folders, BetaSchedules, get_available_motion_models +from .utils_model import Folders, BetaSchedules, get_available_motion_models from .model_injection import ModelPatcherAndInjector, InjectionParams, MotionModelGroup, load_motion_module diff --git a/animatediff/nodes_extras.py b/animatediff/nodes_extras.py index 52f60df..40bb6f5 100644 --- a/animatediff/nodes_extras.py +++ b/animatediff/nodes_extras.py @@ -6,7 +6,7 @@ from comfy.model_patcher import ModelPatcher from comfy.sd import load_checkpoint_guess_config from .logger import logger -from .model_utils import IsChangedHelper, BetaSchedules +from .utils_model import IsChangedHelper, BetaSchedules from .model_injection import get_vanilla_model_patcher diff --git a/animatediff/nodes_gen1.py b/animatediff/nodes_gen1.py index 5887789..c5928dc 100644 --- a/animatediff/nodes_gen1.py +++ b/animatediff/nodes_gen1.py @@ -4,10 +4,10 @@ import torch import comfy.sample as comfy_sample from comfy.model_patcher import ModelPatcher -from .context import ContextOptions, ContextSchedules +from .context import ContextOptions, ContextOptionsGroup, ContextSchedules from .logger import logger -from .model_utils import BetaSchedules, get_available_motion_loras, get_available_motion_models, get_motion_lora_path -from .motion_utils import ADKeyframeGroup +from .utils_model import BetaSchedules, get_available_motion_loras, get_available_motion_models, get_motion_lora_path +from .utils_motion import ADKeyframeGroup from .motion_lora import MotionLoraInfo, MotionLoraList from .model_injection import InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelSettings, load_motion_module from .sample_settings import SampleSettings, SeedNoiseGeneration @@ -43,7 +43,7 @@ class AnimateDiffLoaderWithContext: def load_mm_and_inject_params(self, model: ModelPatcher, model_name: str, beta_schedule: str,# apply_mm_groupnorm_hack: bool, - context_options: ContextOptions=None, motion_lora: MotionLoraList=None, motion_model_settings: MotionModelSettings=None, + context_options: ContextOptionsGroup=None, motion_lora: MotionLoraList=None, motion_model_settings: MotionModelSettings=None, sample_settings: SampleSettings=None, motion_scale: float=1.0, apply_v2_models_properly: bool=False, ad_keyframes: ADKeyframeGroup=None, ): # load motion module diff --git a/animatediff/nodes_gen2.py b/animatediff/nodes_gen2.py index fb5208b..ebfb0af 100644 --- a/animatediff/nodes_gen2.py +++ b/animatediff/nodes_gen2.py @@ -4,10 +4,10 @@ import torch import comfy.sample as comfy_sample from comfy.model_patcher import ModelPatcher -from .context import ContextOptions, ContextSchedules +from .context import ContextOptions, ContextOptionsGroup, ContextSchedules from .logger import logger -from .model_utils import BetaSchedules, get_available_motion_loras, get_available_motion_models, get_motion_lora_path -from .motion_utils import ADKeyframeGroup, ADKeyframe +from .utils_model import BetaSchedules, get_available_motion_loras, get_available_motion_models, get_motion_lora_path +from .utils_motion import ADKeyframeGroup, ADKeyframe from .motion_lora import MotionLoraInfo, MotionLoraList from .model_injection import (InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelPatcher, MotionModelSettings, load_motion_module, load_motion_module_gen2, load_motion_lora_as_patches, validate_model_compatibility_gen2) @@ -35,7 +35,7 @@ class UseEvolvedSamplingNode: CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“" FUNCTION = "use_evolved_sampling" - def use_evolved_sampling(self, model: ModelPatcher, beta_schedule: str, m_models: MotionModelGroup=None, context_options: ContextOptions=None, + def use_evolved_sampling(self, model: ModelPatcher, beta_schedule: str, m_models: MotionModelGroup=None, context_options: ContextOptionsGroup=None, sample_settings: SampleSettings=None, beta_schedule_override=None): if m_models is not None: m_models = m_models.clone() diff --git a/animatediff/nodes_multival.py b/animatediff/nodes_multival.py index ed9d8ff..d1193a1 100644 --- a/animatediff/nodes_multival.py +++ b/animatediff/nodes_multival.py @@ -4,7 +4,7 @@ from typing import Union import torch from torch import Tensor -from .motion_utils import linear_conversion, normalize_min_max +from .utils_motion import linear_conversion, normalize_min_max class ScaleType: diff --git a/animatediff/nodes_sample.py b/animatediff/nodes_sample.py index 7fd0475..371874b 100644 --- a/animatediff/nodes_sample.py +++ b/animatediff/nodes_sample.py @@ -2,7 +2,7 @@ from torch import Tensor from .freeinit import FreeInitFilter from .sample_settings import FreeInitOptions, IterationOptions, NoiseLayerAdd, NoiseLayerAddWeighted, NoiseLayerGroup, NoiseLayerReplace, NoiseLayerType, SeedNoiseGeneration, SampleSettings -from .model_utils import BIGMIN, BIGMAX +from .utils_model import BIGMIN, BIGMAX class SampleSettingsNode: diff --git a/animatediff/sample_settings.py b/animatediff/sample_settings.py index 8e361b4..43f25c7 100644 --- a/animatediff/sample_settings.py +++ b/animatediff/sample_settings.py @@ -7,7 +7,7 @@ import comfy.samplers from comfy.model_patcher import ModelPatcher from . import freeinit -from .context import ContextOptions +from .context import ContextOptions, ContextOptionsGroup from .logger import logger @@ -273,8 +273,8 @@ class SeedNoiseGeneration: @staticmethod def _convert_to_repeated_context(noise: Tensor, extra_args: dict, **kwargs): # if no context_length, return unmodified noise - opts: ContextOptions = extra_args["context_options"] - context_length: int = opts.context_length + opts: ContextOptionsGroup = extra_args["context_options"] + context_length: int = opts.context_length if not opts.view_options else opts.view_options.context_length if context_length is None: return noise length = noise.shape[0] @@ -285,9 +285,9 @@ class SeedNoiseGeneration: @staticmethod def _convert_to_freenoise(noise: Tensor, seed: int, extra_args: dict, **kwargs): # if no context_length, return unmodified noise - opts: ContextOptions = extra_args["context_options"] - context_length: int = opts.context_length - context_overlap: int = opts.context_overlap + opts: ContextOptionsGroup = extra_args["context_options"] + context_length: int = opts.context_length if not opts.view_options else opts.view_options.context_length + context_overlap: int = opts.context_overlap if not opts.view_options else opts.view_options.context_overlap video_length: int = noise.shape[0] if context_length is None: return noise diff --git a/animatediff/sampling.py b/animatediff/sampling.py index 895ae59..82cfbaa 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -14,10 +14,10 @@ import comfy.sample import comfy.utils from comfy.controlnet import ControlBase -from .context import ContextFuseMethod, generate_distance_weight, get_context_windows +from .context import ContextFuseMethod, ContextSchedules, get_context_weights, get_context_windows from .sample_settings import IterationOptions, SeedNoiseGeneration, prepare_mask_ad -from .motion_utils import GroupNormAD -from .model_utils import ModelTypeSD, wrap_function_to_inject_xformers_bug_info +from .utils_motion import GroupNormAD +from .utils_model import ModelTypeSD, wrap_function_to_inject_xformers_bug_info from .model_injection import InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelPatcher from .motion_module_ad import AnimateDiffFormat, AnimateDiffInfo, AnimateDiffVersion, VanillaTemporalModule from .logger import logger @@ -115,13 +115,13 @@ def groupnorm_mm_factory(params: InjectionParams): # axes_factor normalizes batch based on total conds and unconds passed in batch; # the conds and unconds per batch can change based on VRAM optimizations that may kick in if not params.is_using_sliding_context(): - axes_factor = input.size(0)//params.full_length + batched_conds = input.size(0)//params.full_length else: - axes_factor = input.size(0)//params.context_options.context_length + batched_conds = input.size(0)//params.context_options.context_length - input = rearrange(input, "(b f) c h w -> b c f h w", b=axes_factor) + input = rearrange(input, "(b f) c h w -> b c f h w", b=batched_conds) input = group_norm(input, self.num_groups, self.weight, self.bias, self.eps) - input = rearrange(input, "b c f h w -> (b f) c h w", b=axes_factor) + input = rearrange(input, "b c f h w -> (b f) c h w", b=batched_conds) return input return groupnorm_mm_forward @@ -141,7 +141,15 @@ def get_additional_models_factory(orig_get_additional_models: Callable, motion_m def apply_params_to_motion_models(motion_models: MotionModelGroup, params: InjectionParams): params = params.clone() - if params.context_options.context_length and params.full_length > params.context_options.context_length: + if params.context_options.context_schedule == ContextSchedules.VIEW_AS_CONTEXT: + params.context_options._current_context.context_length = params.full_length + # check (and message) should be different based on use_on_equal_length setting + if params.context_options.context_length: + pass + + allow_equal = params.context_options.use_on_equal_length + enough_latents = params.full_length >= params.context_options.context_length if allow_equal else params.full_length > params.context_options.context_length + if params.context_options.context_length and enough_latents: logger.info(f"Sliding context window activated - latents passed in ({params.full_length}) greater than context_length {params.context_options.context_length}.") else: logger.info(f"Regular AnimateDiff activated - latents passed in ({params.full_length}) less or equal to context_length {params.context_options.context_length}.") @@ -156,7 +164,9 @@ def apply_params_to_motion_models(motion_models: MotionModelGroup, params: Injec # otherwise, treat context_length as intended AD frame window else: for motion_model in motion_models.models: - if params.context_options.context_length > motion_model.model.encoding_max_len: + view_options = params.context_options.view_options + context_length = view_options.context_length if view_options else params.context_options.context_length + if context_length > motion_model.model.encoding_max_len: raise ValueError(f"AnimateDiff model {motion_model.model.mm_info.mm_name} has upper limit of {motion_model.model.encoding_max_len} frames for a context window, but received context length of {params.context_options.context_length}.") motion_models.set_video_length(params.context_options.context_length, params.full_length) # inject model @@ -422,9 +432,13 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep, return resized_cond # get context windows - context_windows = get_context_windows(ADGS.current_step, ADGS.params.full_length, ADGS.params.context_options) + ADGS.params.context_options.step = ADGS.current_step + context_windows = get_context_windows(ADGS.params.full_length, ADGS.params.context_options) # figure out how input is split - axes_factor = x_in.size(0)//ADGS.params.full_length + batched_conds = x_in.size(0)//ADGS.params.full_length + + if ADGS.motion_models is not None: + ADGS.motion_models.set_view_options(ADGS.params.context_options.view_options) # prepare final cond, uncond, and out_count cond_final = torch.zeros_like(x_in) @@ -442,7 +456,7 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep, model_options["transformer_options"]["ad_params"]["context_length"] = len(ctx_idxs) # account for all portions of input frames full_idxs = [] - for n in range(axes_factor): + for n in range(batched_conds): for ind in ctx_idxs: full_idxs.append((ADGS.params.full_length*n)+ind) # get subsections of x, timestep, cond, uncond, cond_concat @@ -453,17 +467,12 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep, sub_cond_out, sub_uncond_out = comfy.samplers.calc_cond_uncond_batch(model, sub_cond, sub_uncond, sub_x, sub_timestep, model_options) - if ADGS.params.context_options.fuse_method == ContextFuseMethod.FLAT: - # equal weights for idxs - cond_final[full_idxs] += sub_cond_out - uncond_final[full_idxs] += sub_uncond_out - out_count_final[full_idxs] += 1 # increment which indeces were used - elif ADGS.params.context_options.fuse_method == ContextFuseMethod.PYRAMID: - # greater weight towards center of idxs - weights = torch.Tensor(generate_distance_weight(len(ctx_idxs)) * axes_factor).to(device=x_in.device).unsqueeze(-1).unsqueeze(-1).unsqueeze(-1) - cond_final[full_idxs] += sub_cond_out * weights - uncond_final[full_idxs] += sub_uncond_out * weights - out_count_final[full_idxs] += weights + # add conds and counts based on weights of fuse method + weights = get_context_weights(len(ctx_idxs), ADGS.params.context_options.fuse_method) * batched_conds + weights_tensor = torch.Tensor(weights).to(device=x_in.device).unsqueeze(-1).unsqueeze(-1).unsqueeze(-1) + cond_final[full_idxs] += sub_cond_out * weights_tensor + uncond_final[full_idxs] += sub_uncond_out * weights_tensor + out_count_final[full_idxs] += weights_tensor # normalize cond and uncond via division by context usage counts diff --git a/animatediff/model_utils.py b/animatediff/utils_model.py similarity index 100% rename from animatediff/model_utils.py rename to animatediff/utils_model.py diff --git a/animatediff/motion_utils.py b/animatediff/utils_motion.py similarity index 89% rename from animatediff/motion_utils.py rename to animatediff/utils_motion.py index 8c05e74..0f55b82 100644 --- a/animatediff/motion_utils.py +++ b/animatediff/utils_motion.py @@ -27,7 +27,6 @@ else: optimized_attention_mm = attention_sub_quad -# maintain backwards compatibility with the comfy.ops hasattr check (TODO: remove once a non-backwards compatible change happens) class CrossAttentionMM(nn.Module): def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0., dtype=None, device=None, operations=comfy.ops.disable_weight_init): @@ -110,6 +109,27 @@ def extend_to_batch_size(tensor: Tensor, batch_size: int): return tensor +def get_sorted_list_via_attr(objects: list, attr: str) -> list: + if not objects: + return objects + elif len(objects) <= 1: + return [x for x in objects] + # now that we know we have to sort, do it following these rules: + # a) if objects have same value of attribute, maintain their relative order + # b) perform sorting of the groups of objects with same attributes + unique_attrs = {} + for object in objects: + val_attr = getattr(objects, attr) + unique_attrs.get(val_attr, list()).append(object) + # now that we have the unique attr values grouped together in relative order, sort them by key + sorted_attrs = dict(sorted(unique_attrs.items())) + # now flatten out the dict into a list to return + sorted_list = [] + for object_list in sorted_attrs.values(): + sorted_list.extend(object_list) + return sorted_list + + class MotionCompatibilityError(ValueError): pass